Precision adjustments

This commit is contained in:
kijai
2025-10-27 00:23:32 +02:00
parent a0bdf20817
commit c59e52ca44
2 changed files with 14 additions and 10 deletions
+4 -2
View File
@@ -278,7 +278,7 @@ class WanVideoBlockSwap:
def INPUT_TYPES(s):
return {
"required": {
"blocks_to_swap": ("INT", {"default": 20, "min": 0, "max": 48, "step": 1, "tooltip": "Number of transformer blocks to swap, the 14B model has 40, while the 1.3B and 5B models have 30 blocks. LongCat-video has 48 blocks"}),
"blocks_to_swap": ("INT", {"default": 20, "min": 0, "max": 48, "step": 1, "tooltip": "Number of transformer blocks to swap, the 14B model has 40, while the 1.3B and 5B models have 30 blocks. LongCat-video has 48"}),
"offload_img_emb": ("BOOLEAN", {"default": False, "tooltip": "Offload img_emb to offload_device"}),
"offload_txt_emb": ("BOOLEAN", {"default": False, "tooltip": "Offload time_emb to offload_device"}),
},
@@ -878,10 +878,12 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
scale_key = key.replace(".weight", ".scale_weight")
if scale_key in sd:
dtype_to_use = value.dtype
if "modulation" in name or "norm" in name or "bias" in name or "img_emb" in name:
if "bias" in name or "img_emb" in name:
dtype_to_use = base_dtype
if "patch_embedding" in name or "motion_encoder" in name:
dtype_to_use = torch.float32
if "modulation" in name or "norm" in name:
dtype_to_use = value.dtype if value.dtype == torch.float32 else base_dtype
load_device = transformer_load_device
if block_swap_args is not None: