diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 124e6f8..fc14f1b 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -1090,20 +1090,34 @@ class WanVideoModelLoader: if new_key != key: sd[new_key] = sd.pop(key) + is_scaled_fp8 = False + if quantization == "disabled": for k, v in sd.items(): if isinstance(v, torch.Tensor): if v.dtype == torch.float8_e4m3fn: quantization = "fp8_e4m3fn" if "scaled_fp8" in sd: + is_scaled_fp8 = True quantization = "fp8_e4m3fn_scaled" break elif v.dtype == torch.float8_e5m2: quantization = "fp8_e5m2" if "scaled_fp8" in sd: + is_scaled_fp8 = True quantization = "fp8_e5m2_scaled" break - + + scale_weights = {} + if "fp8" in quantization: + for k, v in sd.items(): + if k.endswith(".scale_weight"): + is_scaled_fp8 = True + break + + if is_scaled_fp8 and "scaled" not in quantization: + quantization = quantization + "_scaled" + if torch.cuda.is_available(): #only warning for now major, minor = torch.cuda.get_device_capability(device) @@ -1111,9 +1125,9 @@ class WanVideoModelLoader: if compile_args is not None and "e4" in quantization and (major, minor) < (8, 9): log.warning("WARNING: Torch.compile with fp8_e4m3fn weights on CUDA compute capability < 8.9 may not be supported. Please use fp8_e5m2, GGUF or higher precision instead, or check the latest triton version that adds support for older architectures https://github.com/woct0rdho/triton-windows/releases/tag/v3.5.0-windows.post21") - if "scaled_fp8" in sd and "scaled" not in quantization: + if is_scaled_fp8 and "scaled" not in quantization: raise ValueError("The model is a scaled fp8 model, please set quantization to '_scaled'") - elif "scaled_fp8" not in sd and "scaled" in quantization: + if not is_scaled_fp8 and "scaled" in quantization: raise ValueError("The model is not a scaled fp8 model, please disable '_scaled' in quantization") if "vace_blocks.0.after_proj.weight" in sd and not "patch_embedding.weight" in sd: diff --git a/nodes_sampler.py b/nodes_sampler.py index d097a46..c506bdd 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -1214,8 +1214,15 @@ class WanVideoSampler: if mocha_embeds is not None: if context_window is not None and mocha_embeds.shape[2] != context_frames: - partial_mocha_embeds = mocha_embeds[:, context_window] - partial_mocha_embeds[:, -mocha_num_refs:] = mocha_embeds[:, -mocha_num_refs:] + latent_frames = len(context_window) + # [latent_frames, 1 mask frame, mocha_num_refs] + latent_end = latent_frames + mask_end = latent_end + 1 + partial_latents = mocha_embeds[:, :, context_window] # windowed latents + mask_frame = mocha_embeds[:, :, latent_end:mask_end] # single mask frame + ref_frames = mocha_embeds[:, :, -mocha_num_refs:] # reference frames + + partial_mocha_embeds = torch.cat([partial_latents, mask_frame, ref_frames], dim=2) z = torch.cat([z, partial_mocha_embeds.to(z)], dim=1) else: z = torch.cat([z, mocha_embeds.to(z)], dim=1)