Merge branch 'kijai:main' into custom-sigmas

This commit is contained in:
cmeka
2025-10-21 12:48:02 -04:00
committed by GitHub
2 changed files with 26 additions and 5 deletions
+17 -3
View File
@@ -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:
+9 -2
View File
@@ -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)