From e3fed708292945435a8de3117f251ca3a00eccfb Mon Sep 17 00:00:00 2001 From: Vladimir Mandic Date: Mon, 29 Jun 2026 15:26:25 +0200 Subject: [PATCH] fix hijack accelerate Signed-off-by: Vladimir Mandic --- modules/sd_hijack_accelerate.py | 34 +++------------------------------ modules/sd_offload.py | 12 ++++++------ 2 files changed, 9 insertions(+), 37 deletions(-) diff --git a/modules/sd_hijack_accelerate.py b/modules/sd_hijack_accelerate.py index 45e892d1e..801c6e8a7 100644 --- a/modules/sd_hijack_accelerate.py +++ b/modules/sd_hijack_accelerate.py @@ -8,7 +8,7 @@ from modules import devices tensor_to_timer = 0 -orig_set_module = accelerate.utils.set_module_tensor_to_device +orig_set_module = accelerate.utils.modeling.set_module_tensor_to_device orig_torch_conv = torch.nn.modules.conv.Conv2d._conv_forward # pylint: disable=protected-access @@ -42,42 +42,14 @@ def hijack_set_module_tensor( tensor_to_timer += (t1 - t0) -def hijack_set_module_tensor_simple( - module: nn.Module, - tensor_name: str, - device: int | str | torch.device, - value: torch.Tensor | None = None, - dtype: str | torch.dtype | None = None, # pylint: disable=unused-argument - fp16_statistics: torch.HalfTensor | None = None, # pylint: disable=unused-argument -): - global tensor_to_timer # pylint: disable=global-statement - if device == 'cpu': # override to load directly to gpu - device = devices.device - t0 = time.time() - if "." in tensor_name: - splits = tensor_name.split(".") - for split in splits[:-1]: - module = getattr(module, split) - tensor_name = splits[-1] - old_value = getattr(module, tensor_name) - with devices.inference_context(): - if tensor_name in module._buffers: # pylint: disable=protected-access - module._buffers[tensor_name] = value.to(device, non_blocking=False) # pylint: disable=protected-access - elif value is not None or not devices.same_device(device, module._parameters[tensor_name].device): # pylint: disable=protected-access - param_cls = type(module._parameters[tensor_name]) # pylint: disable=protected-access - module._parameters[tensor_name] = param_cls(value, requires_grad=old_value.requires_grad).to(device, non_blocking=False) # pylint: disable=protected-access - t1 = time.time() - tensor_to_timer += (t1 - t0) - - def hijack_accelerate(): - accelerate.utils.set_module_tensor_to_device = hijack_set_module_tensor + accelerate.utils.modeling.set_module_tensor_to_device = hijack_set_module_tensor global tensor_to_timer # pylint: disable=global-statement tensor_to_timer = 0 def restore_accelerate(): - accelerate.utils.set_module_tensor_to_device = orig_set_module + accelerate.utils.modeling.set_module_tensor_to_device = orig_set_module def torch_conv_forward(self, input, weight, bias): # pylint: disable=redefined-builtin diff --git a/modules/sd_offload.py b/modules/sd_offload.py index 0f60fc645..dd143e1ef 100644 --- a/modules/sd_offload.py +++ b/modules/sd_offload.py @@ -280,12 +280,12 @@ class OffloadHook(accelerate.hooks.ModelHook): skip_keys = getattr(module, "_skip_keys", None) try: module = accelerate.dispatch_model(module, - main_device=torch.device(devices.device), - device_map=device_map, - offload_dir=offload_dir, - skip_keys=skip_keys, - force_hooks=True, - ) + main_device=torch.device(devices.device), + device_map=device_map, + offload_dir=offload_dir, + skip_keys=skip_keys, + force_hooks=True, + ) except Exception as e: # reapply hook log.warning(f'Offload: type=balanced op=dispatch module={module.__class__.__name__} {e}') module = accelerate.hooks.remove_hook_from_module(module, recurse=True)