From c9f49720c55cd7bbe2fb620e33418a45e00892db Mon Sep 17 00:00:00 2001 From: Disty0 Date: Wed, 25 Jun 2025 23:37:41 +0300 Subject: [PATCH] Cleanup --- modules/devices.py | 2 +- modules/sd_hijack_accelerate.py | 4 ++-- modules/sd_offload.py | 2 +- 3 files changed, 4 insertions(+), 4 deletions(-) diff --git a/modules/devices.py b/modules/devices.py index 1c35f2683..8d15fd238 100644 --- a/modules/devices.py +++ b/modules/devices.py @@ -665,6 +665,6 @@ def normalize_device(dev): def same_device(d1, d2): - if d1.type != d2.type: + if torch.device(d1).type != torch.device(d2).type: return False return normalize_device(d1) == normalize_device(d2) diff --git a/modules/sd_hijack_accelerate.py b/modules/sd_hijack_accelerate.py index f8cf8983f..7f312a029 100644 --- a/modules/sd_hijack_accelerate.py +++ b/modules/sd_hijack_accelerate.py @@ -36,7 +36,7 @@ def hijack_set_module_tensor( # note: majority of time is spent on .to(old_value.dtype) if tensor_name in module._buffers: # pylint: disable=protected-access module._buffers[tensor_name] = value.to(device, old_value.dtype) # pylint: disable=protected-access - elif value is not None or not devices.same_device(torch.device(device), module._parameters[tensor_name].device): # 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, old_value.dtype) # pylint: disable=protected-access t1 = time.time() @@ -64,7 +64,7 @@ def hijack_set_module_tensor_simple( with devices.inference_context(): if tensor_name in module._buffers: # pylint: disable=protected-access module._buffers[tensor_name] = value.to(device) # pylint: disable=protected-access - elif value is not None or not devices.same_device(torch.device(device), module._parameters[tensor_name].device): # 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) # pylint: disable=protected-access t1 = time.time() diff --git a/modules/sd_offload.py b/modules/sd_offload.py index e36db4332..70777a30a 100644 --- a/modules/sd_offload.py +++ b/modules/sd_offload.py @@ -171,7 +171,7 @@ class OffloadHook(accelerate.hooks.ModelHook): return module def pre_forward(self, module, *args, **kwargs): - if devices.normalize_device(module.device) != devices.normalize_device(devices.device): + if not devices.same_device(module.device, devices.device): device_index = torch.device(devices.device).index if device_index is None: device_index = 0