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