fix cpu fallback and torch remap handling

This commit is contained in:
xmarre
2026-04-12 05:54:14 +02:00
parent 752def584f
commit ca8118de03
+7 -1
View File
@@ -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"