| 246 | module._diffusers_hook.remove_hook(name, recurse=False) |
| 247 | |
| 248 | def reset_stateful_hooks(self, recurse: bool = True) -> None: |
| 249 | for hook_name in reversed(self._hook_order): |
| 250 | hook = self.hooks[hook_name] |
| 251 | if hook._is_stateful: |
| 252 | hook.reset_state(self._module_ref) |
| 253 | |
| 254 | if recurse: |
| 255 | for module_name, module in unwrap_module(self._module_ref).named_modules(): |
| 256 | if module_name == "": |
| 257 | continue |
| 258 | module = unwrap_module(module) |
| 259 | if hasattr(module, "_diffusers_hook"): |
| 260 | module._diffusers_hook.reset_stateful_hooks(recurse=False) |
| 261 | |
| 262 | @classmethod |
| 263 | def check_if_exists_or_initialize(cls, module: torch.nn.Module) -> "HookRegistry": |