| 138 | return self.transformer(*args, **kwargs) |
| 139 | |
| 140 | def collect_hooks_(self): |
| 141 | names = list(HOOKS_DEFAULT.keys()) |
| 142 | hooks = {} |
| 143 | hook_origins = {} |
| 144 | for name in names: |
| 145 | if hasattr(self, name): |
| 146 | hooks[name] = getattr(self, name) |
| 147 | hook_origins[name] = 'model' |
| 148 | |
| 149 | for mixin_name, m in self.mixins.items(): |
| 150 | if hasattr(m, name): |
| 151 | if hasattr(getattr(m, name), 'non_conflict'): |
| 152 | # check getattr(m, name), who must accept old_impl as an argument |
| 153 | signature = inspect.signature(getattr(m, name)) |
| 154 | if 'old_impl' not in signature.parameters: |
| 155 | raise ValueError(f'Hook {name} at {mixin_name} must accept old_impl as an argument.') |
| 156 | # ------------- |
| 157 | if name in hooks: |
| 158 | old_impl = hooks[name] |
| 159 | elif name == 'attention_fn': # the only hook without self |
| 160 | old_impl = HOOKS_DEFAULT[name] |
| 161 | else: |
| 162 | old_impl = partial(HOOKS_DEFAULT[name], self) # relax! `partial` does not affect the signature |
| 163 | old_origin = hook_origins.get(name, 'default') |
| 164 | hooks[name] = partial(getattr(m, name), old_impl=old_impl) |
| 165 | hook_origins[name] = mixin_name + ' -> ' + old_origin |
| 166 | elif name in hooks and not hasattr(hooks[name], 'replacable'): # if this hook name is already registered |
| 167 | raise ValueError(f'Hook {name} conflicts at {mixin_name} and {hook_origins[name]}.') |
| 168 | else: # new hook |
| 169 | if name in hooks and hasattr(hooks[name], 'replacable'): |
| 170 | warnings.warn(f'Hook {name} at {mixin_name} replaces {hook_origins[name]}.') |
| 171 | hooks[name] = getattr(m, name) |
| 172 | hook_origins[name] = mixin_name |
| 173 | |
| 174 | self.hooks = hooks |
| 175 | self.hook_origins = hook_origins |
| 176 | return hooks |
| 177 | |
| 178 | def disable_untrainable_params(self): |
| 179 | pass |