diff --git a/src/diffusers/utils/accelerate_utils.py b/src/diffusers/utils/accelerate_utils.py index ae6a8ca747ac..f7deeb6e3931 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 d1a59cec52f1..e985f8620d92 100755 --- a/tests/others/test_utils.py +++ b/tests/others/test_utils.py @@ -269,6 +269,18 @@ 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 TestApplyForwardHook: + def test_preserves_wrapped_function_metadata(self): + from diffusers.utils.accelerate_utils import apply_forward_hook + + @apply_forward_hook + def example(self): + """Example method docstring.""" + + assert example.__name__ == "example" + assert example.__doc__ == "Example method docstring." + + # Copied from https://github.com/huggingface/transformers/blob/main/tests/utils/test_expectations.py class TestExpectations: def test_expectations(self):