[registry] Add a strict mode to model registration (#14933)
This commit is contained in:
@@ -19,8 +19,10 @@ class _ModelRegistry:
|
|||||||
# Keyed by model_arch
|
# Keyed by model_arch
|
||||||
models: Dict[str, Union[Type[nn.Module], str]] = field(default_factory=dict)
|
models: Dict[str, Union[Type[nn.Module], str]] = field(default_factory=dict)
|
||||||
|
|
||||||
def register(self, package_name: str, overwrite: bool = False):
|
def register(
|
||||||
new_models = import_model_classes(package_name)
|
self, package_name: str, overwrite: bool = False, strict: bool = False
|
||||||
|
):
|
||||||
|
new_models = import_model_classes(package_name, strict=strict)
|
||||||
if overwrite:
|
if overwrite:
|
||||||
self.models.update(new_models)
|
self.models.update(new_models)
|
||||||
else:
|
else:
|
||||||
@@ -88,7 +90,7 @@ class _ModelRegistry:
|
|||||||
|
|
||||||
|
|
||||||
@lru_cache()
|
@lru_cache()
|
||||||
def import_model_classes(package_name: str):
|
def import_model_classes(package_name: str, strict: bool = False):
|
||||||
model_arch_name_to_cls = {}
|
model_arch_name_to_cls = {}
|
||||||
package = importlib.import_module(package_name)
|
package = importlib.import_module(package_name)
|
||||||
for _, name, ispkg in pkgutil.iter_modules(package.__path__, package_name + "."):
|
for _, name, ispkg in pkgutil.iter_modules(package.__path__, package_name + "."):
|
||||||
@@ -100,6 +102,8 @@ def import_model_classes(package_name: str):
|
|||||||
try:
|
try:
|
||||||
module = importlib.import_module(name)
|
module = importlib.import_module(name)
|
||||||
except Exception as e:
|
except Exception as e:
|
||||||
|
if strict:
|
||||||
|
raise
|
||||||
logger.warning(f"Ignore import error when loading {name}: {e}")
|
logger.warning(f"Ignore import error when loading {name}: {e}")
|
||||||
continue
|
continue
|
||||||
if hasattr(module, "EntryClass"):
|
if hasattr(module, "EntryClass"):
|
||||||
|
|||||||
Reference in New Issue
Block a user