| 3 | |
| 4 | |
| 5 | class Registry(object): |
| 6 | def __init__(self, name): |
| 7 | self._name = name |
| 8 | self._module_dict = dict() |
| 9 | |
| 10 | def __repr__(self): |
| 11 | format_str = self.__class__.__name__ + '(name={}, items={})'.format( |
| 12 | self._name, list(self._module_dict.keys())) |
| 13 | return format_str |
| 14 | |
| 15 | def __len__(self): |
| 16 | return len(self._module_dict) |
| 17 | |
| 18 | @property |
| 19 | def name(self): |
| 20 | return self._name |
| 21 | |
| 22 | @property |
| 23 | def module_dict(self): |
| 24 | return self._module_dict |
| 25 | |
| 26 | def get(self, key): |
| 27 | return self._module_dict.get(key, None) |
| 28 | |
| 29 | def registe_with_name(self, module_name=None, force=False): |
| 30 | return partial(self.register, module_name=module_name, force=force) |
| 31 | |
| 32 | def register(self, module_build_function, module_name=None, force=False): |
| 33 | """Register a module build function. |
| 34 | |
| 35 | Args: |
| 36 | module (:obj:`nn.Module`): Module to be registered. |
| 37 | """ |
| 38 | if not inspect.isfunction(module_build_function): |
| 39 | raise TypeError( |
| 40 | 'module_build_function must be a function, but got {}'.format( |
| 41 | type(module_build_function))) |
| 42 | if module_name is None: |
| 43 | module_name = module_build_function.__name__ |
| 44 | if not force and module_name in self._module_dict: |
| 45 | raise KeyError('{} is already registered in {}'.format( |
| 46 | module_name, self.name)) |
| 47 | self._module_dict[module_name] = module_build_function |
| 48 | |
| 49 | return module_build_function |
| 50 | |
| 51 | |
| 52 | MODULE_BUILD_FUNCS = Registry('model build functions') |
no outgoing calls
no test coverage detected