From ca8118de031b8dd8b00a56946586dee0536071db Mon Sep 17 00:00:00 2001 From: xmarre Date: Sun, 12 Apr 2026 05:54:14 +0200 Subject: [PATCH] fix cpu fallback and torch remap handling --- patches.py | 8 +++++++- 1 file changed, 7 insertions(+), 1 deletion(-) diff --git a/patches.py b/patches.py index 83e5c32..0381dd1 100644 --- a/patches.py +++ b/patches.py @@ -71,10 +71,14 @@ def _prepare_loaded_tensor( requested_device: torch.device, *, disable_mmap: bool, + move_to_requested_device: bool = False, ) -> torch.Tensor: if _tensor_key_requires_cpu(key): return _copy_tensor_if_needed(tensor, torch.device("cpu")) + if move_to_requested_device and tensor.device != requested_device: + return _copy_tensor_if_needed(tensor, requested_device) + if disable_mmap and tensor.device.type == "cpu": return _copy_tensor_if_needed(tensor, requested_device, force_copy=True) @@ -188,9 +192,10 @@ def _patched_load_torch_file(ckpt, safe_load=False, device=None, return_metadata for key in handle.keys(): sd[key] = _prepare_loaded_tensor( key, - handle.get_tensor(key).to(requested_device), + handle.get_tensor(key), requested_device, disable_mmap=False, + move_to_requested_device=True, ) if return_metadata: metadata = handle.metadata() @@ -261,6 +266,7 @@ def _patched_load_torch_file(ckpt, safe_load=False, device=None, return_metadata value, requested_device, disable_mmap=False, + move_to_requested_device=True, ) method = "torch_load_cpu_first_to_cuda" if requested_device.type == "cuda" else "torch_load_cpu"