Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
18 changes: 16 additions & 2 deletions src/diffusers/hooks/hooks.py
Original file line number Diff line number Diff line change
Expand Up @@ -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.

Expand All @@ -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.
Expand Down Expand Up @@ -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)
Expand Down
2 changes: 1 addition & 1 deletion src/diffusers/hooks/layerwise_casting.py
Original file line number Diff line number Diff line change
Expand Up @@ -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 "
Expand Down
2 changes: 1 addition & 1 deletion src/diffusers/hooks/sea_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -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:
Expand Down
48 changes: 48 additions & 0 deletions tests/hooks/test_hooks.py
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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
Loading