MCPcopy Create free account
hub / github.com/OpenDriveLab/ReSim / collect_hooks_

Method collect_hooks_

SwissArmyTransformer/sat/model/base_model.py:140–176  ·  view source on GitHub ↗
(self)

Source from the content-addressed store, hash-verified

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

Callers 3

__init__Method · 0.95
add_mixinMethod · 0.95
del_mixinMethod · 0.95

Calls 1

getMethod · 0.80

Tested by

no test coverage detected