From 6b3fb6bbc9fd04ac7e3e31058188364ccc726e82 Mon Sep 17 00:00:00 2001 From: "jinli.yl" Date: Wed, 26 Jun 2024 18:32:55 +0800 Subject: [PATCH] [dev] modify base model registry module --- memory_scope/models/base_model.py | 2 +- memory_scope/utils/registry.py | 4 ++-- 2 files changed, 3 insertions(+), 3 deletions(-) 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]