File size: 5,997 Bytes
8c9ba62 | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 | import traceback
from typing import Any, Type
from trinity.utils.log import get_logger
class Registry(object):
"""A class for registry."""
def __init__(self, name: str, default_mapping: dict = {}):
"""
Args:
name (`str`): The name of the registry.
default_mapping (`dict`): Default mapping from module names to module paths (strings).
"""
self._name = name
self._modules = {}
self._default_mapping = default_mapping
self.logger = get_logger()
@property
def name(self) -> str:
"""
Get name of current registry.
Returns:
`str`: The name of current registry.
"""
return self._name
@property
def modules(self) -> dict:
"""
Get all modules in current registry.
Returns:
`dict`: A dict storing modules in current registry.
"""
return self._modules
def get(self, module_key) -> Any:
"""
Get module named module_key from in current registry. If not found,
return None.
Args:
module_key (`str`): specified module name
Returns:
`Any`: the module object
"""
module = self._modules.get(module_key, None)
if module is None:
# try to get from default mapping
if module_key in self._default_mapping:
module_path, class_name = self._default_mapping[module_key].rsplit(".", 1)
try:
module = self._dynamic_import(module_path, class_name)
except Exception:
self.logger.error(
f"Failed to dynamically import {class_name} from {module_path}:\n"
+ traceback.format_exc()
)
raise ImportError(f"Cannot dynamically import {class_name} from {module_path}")
# try to get from string path
elif isinstance(module_key, str) and "." in module_key:
module_path, class_name = module_key.rsplit(".", 1)
try:
module = self._dynamic_import(module_path, class_name)
except Exception:
self.logger.error(
f"Failed to dynamically import {class_name} from {module_path}:\n"
+ traceback.format_exc()
)
raise ImportError(f"Cannot dynamically import {class_name} from {module_path}")
self._register_module(module_name=module_key, module_cls=module)
elif module_key is None:
self.logger.info("Empty module key, return None")
return None
else:
raise ValueError(f"Invalid module key: {module_key}")
return module
def _register_module(self, module_name=None, module_cls=None, force=False):
"""
Register module to registry.
"""
if module_name is None:
module_name = module_cls.__name__
if module_name in self._modules and not force:
self.logger.warning(
f"{module_name} is already registered in {self._name}, "
f"if you want to override it, please set force=True."
)
raise KeyError(f"{module_name} is already registered in {self._name}")
self._modules[module_name] = module_cls
module_cls._name = module_name
def register_module(self, module_name: str, module_cls: Type = None, force=False, lazy=False):
"""
Register module class object to registry with the specified module name.
Args:
module_name (`str`): The module name.
module_cls (`Type`): module class object
force (`bool`): Whether to override an existing class with
the same name. Default: False.
lazy (`bool`): Whether to register the module class object lazily.
Default: False.
Example:
.. code-block:: python
WORKFLOWS = Registry("workflows")
# register a module using decorator
@WORKFLOWS.register_module(name="workflow_name")
class MyWorkflow(Workflow):
pass
# or register a module directly
WORKFLOWS.register_module(
name="workflow_name",
module_cls=MyWorkflow,
force=True,
)
"""
if not (module_name is None or isinstance(module_name, str)):
raise TypeError(f"module_name must be either of None, str," f"got {type(module_name)}")
if module_cls is not None:
self._register_module(module_name=module_name, module_cls=module_cls, force=force)
return module_cls
# if module_cls is None, should return a decorator function
def _register(module_cls):
"""
Register module class object to registry.
Args:
module_cls (`Type`): module class object
Returns:
`Type`: Decorated module class object.
"""
self._register_module(module_name=module_name, module_cls=module_cls, force=force)
return module_cls
return _register
def _dynamic_import(self, module_path: str, class_name: str) -> Type:
"""
Dynamically import a module class object from the specified module path.
Args:
module_path (`str`): The module path. For example, "my_package.my_module".
class_name (`str`): The class name. For example, "MyWorkflow".
Returns:
`Type`: The imported module class object.
"""
import importlib
module = importlib.import_module(module_path)
module_cls = getattr(module, class_name)
return module_cls
|