diff --git a/memory_scope/models/base_model.py b/memory_scope/models/base_model.py index 038cde81..e2e8767e 100644 --- a/memory_scope/models/base_model.py +++ b/memory_scope/models/base_model.py @@ -33,7 +33,7 @@ class BaseModel(metaclass=ABCMeta): self.data = {} self.logger = Logger.get_logger() - obj_cls = MODEL_REGISTRY.get(self.method_type) + obj_cls = MODEL_REGISTRY[self.method_type] if not obj_cls: raise RuntimeError(f"method_type={self.method_type} is not supported!") diff --git a/memory_scope/utils/registry.py b/memory_scope/utils/registry.py index 88807387..9d7996ea 100644 --- a/memory_scope/utils/registry.py +++ b/memory_scope/utils/registry.py @@ -28,6 +28,6 @@ class Registry(object): raise NotImplementedError self.module_dict.update(module_name_dict) - def get(self, module_name: str): - assert module_name in self.module_dict, f'{module_name} not found in {self.name}' + def __getitem__(self, module_name: str): + assert module_name in self.module_dict, f"{module_name} not found in {self.name}" return self.module_dict[module_name]