From fbdae295480cbfa374279ce7bdc68ad09d9fd4af Mon Sep 17 00:00:00 2001 From: javierdejesusda Date: Sun, 24 May 2026 21:16:38 +0200 Subject: [PATCH 1/2] [utils] preserve method metadata in apply_forward_hook Add @functools.wraps to the inner wrapper so it inherits __name__, __doc__, __module__ and the signature of the wrapped method. Without this, methods decorated with @apply_forward_hook (e.g. AutoencoderKL's encode and decode) lose their docstring and signature, so doc-builder cannot render them through the [[autodoc]] directive and they show up as "wrapper" with a generic (*args, **kwargs) signature on the live API docs. The same affects VQModel, ConsistencyDecoderVAE and the other autoencoders that use the decorator. Add regression tests asserting that the wrapped method exposes the right metadata (name, docstring, signature including default values, __wrapped__) and that the underlying offload hook still fires through the wrapper. Fixes #6575 --- src/diffusers/utils/accelerate_utils.py | 3 +++ tests/others/test_utils.py | 21 +++++++++++++++++++++ 2 files changed, 24 insertions(+) diff --git a/src/diffusers/utils/accelerate_utils.py b/src/diffusers/utils/accelerate_utils.py index af3b712b5a9e..b1ad56f55227 100644 --- a/src/diffusers/utils/accelerate_utils.py +++ b/src/diffusers/utils/accelerate_utils.py @@ -15,6 +15,8 @@ Accelerate utilities: Utilities related to accelerate """ +import functools + from packaging import version from .import_utils import is_accelerate_available @@ -40,6 +42,7 @@ def apply_forward_hook(method): if version.parse(accelerate_version) < version.parse("0.17.0"): return method + @functools.wraps(method) def wrapper(self, *args, **kwargs): if hasattr(self, "_hf_hook") and hasattr(self._hf_hook, "pre_forward"): self._hf_hook.pre_forward(self) diff --git a/tests/others/test_utils.py b/tests/others/test_utils.py index 4600f5f3710a..646d335d2ec8 100755 --- a/tests/others/test_utils.py +++ b/tests/others/test_utils.py @@ -284,6 +284,27 @@ def _capture(target_device): ) +class ApplyForwardHookTester(unittest.TestCase): + """Tests for :func:`diffusers.utils.accelerate_utils.apply_forward_hook`.""" + + def test_preserves_wrapped_function_metadata(self): + import inspect + + from diffusers.utils.accelerate_utils import apply_forward_hook + + @apply_forward_hook + def example(self, x: int, y: int = 0) -> int: + """Example method docstring.""" + return x + y + + assert example.__name__ == "example" + assert example.__doc__ == "Example method docstring." + assert hasattr(example, "__wrapped__") + signature = inspect.signature(example) + assert list(signature.parameters) == ["self", "x", "y"] + assert signature.parameters["y"].default == 0 + + # Copied from https://github.com/huggingface/transformers/blob/main/tests/utils/test_expectations.py class ExpectationsTester(unittest.TestCase): def test_expectations(self): From 13cb81ee5669c57ca108228f11a06b2e64aa6723 Mon Sep 17 00:00:00 2001 From: Javier de Jesus Date: Wed, 30 Sep 2026 17:14:05 +0000 Subject: [PATCH 2/2] Trim apply_forward_hook test to a single assertion --- tests/others/test_utils.py | 13 ++----------- 1 file changed, 2 insertions(+), 11 deletions(-) diff --git a/tests/others/test_utils.py b/tests/others/test_utils.py index b3b9c7ece0ad..e985f8620d92 100755 --- a/tests/others/test_utils.py +++ b/tests/others/test_utils.py @@ -269,25 +269,16 @@ def _capture(target_device): assert "moved to" in cuda_out, f"Non-MPS target should still emit the CPU-fallback info log, got: {cuda_out}" -class ApplyForwardHookTester(unittest.TestCase): - """Tests for :func:`diffusers.utils.accelerate_utils.apply_forward_hook`.""" - +class TestApplyForwardHook: def test_preserves_wrapped_function_metadata(self): - import inspect - from diffusers.utils.accelerate_utils import apply_forward_hook @apply_forward_hook - def example(self, x: int, y: int = 0) -> int: + def example(self): """Example method docstring.""" - return x + y assert example.__name__ == "example" assert example.__doc__ == "Example method docstring." - assert hasattr(example, "__wrapped__") - signature = inspect.signature(example) - assert list(signature.parameters) == ["self", "x", "y"] - assert signature.parameters["y"].default == 0 # Copied from https://github.com/huggingface/transformers/blob/main/tests/utils/test_expectations.py