Update nodes_model_loading.py

This commit is contained in:
kijai
2025-08-29 20:41:40 +03:00
parent 8f6507fa64
commit 3a290cfd15
+2 -1
View File
@@ -822,7 +822,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
if "patch_embedding" in name:
dtype_to_use = torch.float32
load_device = device
load_device = transformer_load_device
if block_swap_args is not None:
if block_idx is not None:
if block_idx >= len(transformer.blocks) - block_swap_args.get("blocks_to_swap", 0):
@@ -1343,6 +1343,7 @@ class WanVideoModelLoader:
scale_weights.clear()
patcher.patches.clear()
transformer.patched_linear = False
sd = None
else:
from .custom_linear import _replace_linear
transformer = _replace_linear(transformer, base_dtype, sd, scale_weights=scale_weights)