diff --git a/src/diffusers/hooks/hooks.py b/src/diffusers/hooks/hooks.py index 50fb00133aab..4c285da3e552 100644 --- a/src/diffusers/hooks/hooks.py +++ b/src/diffusers/hooks/hooks.py @@ -112,7 +112,7 @@ def initialize_hook(self, module: torch.nn.Module) -> torch.nn.Module: """ return module - def deinitalize_hook(self, module: torch.nn.Module) -> torch.nn.Module: + def deinitialize_hook(self, module: torch.nn.Module) -> torch.nn.Module: r""" Hook that is executed when a model is deinitialized. @@ -122,6 +122,20 @@ def deinitalize_hook(self, module: torch.nn.Module) -> torch.nn.Module: """ return module + def __init_subclass__(cls, **kwargs) -> None: + # A subclass may still override the misspelled `deinitalize_hook`. Point the other name at + # that override so `HookRegistry.remove_hook` and existing callers both reach it. + super().__init_subclass__(**kwargs) + defines_new = "deinitialize_hook" in cls.__dict__ + defines_old = "deinitalize_hook" in cls.__dict__ + if defines_old and not defines_new: + cls.deinitialize_hook = cls.__dict__["deinitalize_hook"] + elif defines_new and not defines_old: + cls.deinitalize_hook = cls.__dict__["deinitialize_hook"] + + # Historical misspelling of `deinitialize_hook`. Same function, so existing callers keep working. + deinitalize_hook = deinitialize_hook + def pre_forward(self, module: torch.nn.Module, *args, **kwargs) -> tuple[tuple[Any], dict[str, Any]]: r""" Hook that is executed just before the forward method of the model. @@ -263,7 +277,7 @@ def remove_hook(self, name: str, recurse: bool = True) -> None: else: self._fn_refs[index + 1].forward = old_forward - self._module_ref = hook.deinitalize_hook(self._module_ref) + self._module_ref = hook.deinitialize_hook(self._module_ref) del self.hooks[name] self._hook_order.pop(index) self._fn_refs.pop(index) diff --git a/src/diffusers/hooks/layerwise_casting.py b/src/diffusers/hooks/layerwise_casting.py index bfd6b1281003..2829dca05bd9 100644 --- a/src/diffusers/hooks/layerwise_casting.py +++ b/src/diffusers/hooks/layerwise_casting.py @@ -55,7 +55,7 @@ def initialize_hook(self, module: torch.nn.Module): module.to(dtype=self.storage_dtype, non_blocking=self.non_blocking) return module - def deinitalize_hook(self, module: torch.nn.Module): + def deinitialize_hook(self, module: torch.nn.Module): raise NotImplementedError( "LayerwiseCastingHook does not support deinitialization. A model once enabled with layerwise casting will " "have casted its weights to a lower precision dtype for storage. Casting this back to the original dtype " diff --git a/src/diffusers/hooks/sea_cache.py b/src/diffusers/hooks/sea_cache.py index d228b0772e48..d161812c5a20 100644 --- a/src/diffusers/hooks/sea_cache.py +++ b/src/diffusers/hooks/sea_cache.py @@ -492,7 +492,7 @@ def cached_decoder_stack(und_seq, gen_seq, rotary_emb): unwrapped_module._run_decoder_stack = self._installed_decoder_stack return module - def deinitalize_hook(self, module: torch.nn.Module): + def deinitialize_hook(self, module: torch.nn.Module): if self.use_stack_boundary: unwrapped_module = unwrap_module(module) if unwrapped_module.__dict__.get("_run_decoder_stack") is self._installed_decoder_stack: diff --git a/tests/hooks/test_hooks.py b/tests/hooks/test_hooks.py index 2dce9da5dcee..dcd3f5d893cb 100644 --- a/tests/hooks/test_hooks.py +++ b/tests/hooks/test_hooks.py @@ -19,6 +19,8 @@ import torch from diffusers.hooks import HookRegistry, ModelHook +from diffusers.hooks.layerwise_casting import LayerwiseCastingHook +from diffusers.hooks.sea_cache import SeaCacheRootHook from diffusers.training_utils import free_memory from diffusers.utils.logging import get_logger @@ -436,3 +438,49 @@ def test_register_hook_preserves_forward_signature_torch_export(self): exported_output = exported.module()(hidden_states=hidden_states, timestep=timestep) assert torch.allclose(exported_output, eager_output) + + def test_deinitialize_hook_alias(self): + assert ModelHook.deinitialize_hook is ModelHook.deinitalize_hook + assert LayerwiseCastingHook.deinitialize_hook is LayerwiseCastingHook.deinitalize_hook + assert LayerwiseCastingHook.deinitialize_hook is not ModelHook.deinitialize_hook + assert SeaCacheRootHook.deinitialize_hook is SeaCacheRootHook.deinitalize_hook + assert SeaCacheRootHook.deinitialize_hook is not ModelHook.deinitialize_hook + + class LegacyHook(ModelHook): + def deinitalize_hook(self, module): + module.saw_legacy = True + return module + + class RenamedHook(ModelHook): + def deinitialize_hook(self, module): + module.saw_renamed = True + return module + + class ChildOfRenamed(RenamedHook): + def deinitalize_hook(self, module): + module.saw_child = True + return module + + assert LegacyHook.deinitialize_hook is LegacyHook.deinitalize_hook + assert RenamedHook.deinitialize_hook is RenamedHook.deinitalize_hook + assert ChildOfRenamed.deinitialize_hook is ChildOfRenamed.deinitalize_hook + assert ChildOfRenamed.deinitialize_hook is not RenamedHook.deinitialize_hook + + legacy_module = torch.nn.Linear(2, 2) + registry = HookRegistry.check_if_exists_or_initialize(legacy_module) + registry.register_hook(LegacyHook(), "legacy") + registry.remove_hook("legacy") + assert legacy_module.saw_legacy is True + + renamed_module = torch.nn.Linear(2, 2) + registry = HookRegistry.check_if_exists_or_initialize(renamed_module) + renamed_hook = RenamedHook() + registry.register_hook(renamed_hook, "renamed") + renamed_hook.deinitalize_hook(renamed_module) + assert renamed_module.saw_renamed is True + + child_module = torch.nn.Linear(2, 2) + registry = HookRegistry.check_if_exists_or_initialize(child_module) + registry.register_hook(ChildOfRenamed(), "child") + registry.remove_hook("child") + assert child_module.saw_child is True