Squashed commit of the following:

commit c3eb0f49faf68ab953f1b08b7e00225e041e5d0b
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 9 12:55:49 2025 +0200

    move workflow

commit e129e25c26f9b55b527dd3e9f15c6e3f215af11f
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 9 11:17:17 2025 +0200

    Fix padding

commit f252f34eff5cc15ec6fc475f929cafa3e5b7f46c
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 9 01:38:17 2025 +0200

    Add long video example

commit 09ceab808b67a3b2fb7d1ee5fc0a1ad667739e2a
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Tue Dec 9 01:31:48 2025 +0200

    Support extension

commit 7ca221874e8a2cabfc766c51bb63774fde3c851b
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 12:28:29 2025 +0200

    Might as well not even do control pass on uncond...

commit b55caf299e4d89148f5885e8e56bf8e411472dc3
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 12:15:59 2025 +0200

    Cfg fixes

commit fd54ba23e6746acb33a8bf124e5bc7de9d947ff1
Merge: 2f97b1b e867e64
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 10:39:55 2025 +0200

    Merge branch 'main' into onetoall

commit 2f97b1bd887367962542b9a6058f9f6e3c4ad4d7
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 09:32:09 2025 +0200

    Add ref_mask input

commit 74cad232fd35347c50f2ed7465ff13e179ef8402
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 03:44:42 2025 +0200

    Update nodes_model_loading.py

commit 01a038eb4a30f29d868fbaef190e6e90da1a058d
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 03:11:08 2025 +0200

    Fix indentation

commit a95f4d6eaa4468e818910fec7ba11e1f92423d9b
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 02:54:47 2025 +0200

    Update model.py

commit ad006985a1bafdf5941c0fa85a47852eb20a818a
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 02:54:19 2025 +0200

    Fix token replace

commit b5f0f44f1720586950756ad142a538e04814270f
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 02:50:52 2025 +0200

    Don't use token replace by default

commit 874174ec2921c528a4373097fd0bebbbb5257606
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 02:24:47 2025 +0200

    Create WanToAllAnimation_test.json

commit 9e6175855618c94c1bcb89c4b89879219410ce53
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 02:23:15 2025 +0200

    Add token replacement

commit 41fd76dfcbf0e70a3a7308a6fa0652fb492ed1f6
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 00:45:33 2025 +0200

    Use correct norm for reference attn

commit 705f5dcc8b6cd5fa6fe453f9bd01ffdf43a23078
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Mon Dec 8 00:11:17 2025 +0200

    cleanup

commit 4f095d97f80da807417d49d9aa7e9ee47145c85f
Merge: 3e4e4db 2369cdb
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Dec 7 18:44:01 2025 +0200

    Merge branch 'main' into onetoall

commit 3e4e4db35d3e266c39d48cd683f60384a737eca5
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sun Dec 7 00:27:23 2025 +0200

    handle controlnet better

commit c5742552a9af4a3ae208f9c2ead6e1105cc2c348
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 6 17:24:45 2025 +0200

    cleanup

commit c06ff9c06651c32953236802bd7fb385b9cf93ab
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 6 03:41:02 2025 +0200

    3D rope for controlnet

commit 948ea6b783f54892515cbc9cfe66484913904ee7
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 6 03:08:04 2025 +0200

    pose input scaling

commit 90c2eff3b2d30d3a92ff5c27e4327a0ac80b642c
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Sat Dec 6 02:37:48 2025 +0200

    Cleanup

commit 9f7683422c1aa8ebe4d3380a86be98d6c589b270
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Fri Dec 5 23:29:05 2025 +0200

    pose control

commit 0f217be4d8742741b0f89db50138214302a58dc3
Author: kijai <40791699+kijai@users.noreply.github.com>
Date:   Fri Dec 5 20:55:10 2025 +0200

    Support reference input
This commit is contained in:
kijai
2025-12-09 12:56:11 +02:00
parent e867e642d4
commit 8b037bce2e
11 changed files with 10605 additions and 166 deletions
+49 -14
View File
@@ -6,7 +6,7 @@ import numpy as np
from tqdm import tqdm
import re
from .wanvideo.modules.model import WanModel, LoRALinearLayer
from .wanvideo.modules.model import WanModel, LoRALinearLayer, WanRMSNorm
from .wanvideo.modules.t5 import T5EncoderModel
from .wanvideo.modules.clip import CLIPModel
from .wanvideo.wan_video_vae import WanVideoVAE, WanVideoVAE38
@@ -852,18 +852,18 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None,
total=param_count,
leave=True):
block_idx = vace_block_idx = None
if "vace_blocks." in name:
if name.startswith("vace_blocks."):
try:
vace_block_idx = int(name.split("vace_blocks.")[1].split(".")[0])
except Exception:
vace_block_idx = None
elif "blocks." in name and "face" not in name:
elif name.startswith("blocks.") and "face" not in name:
try:
block_idx = int(name.split("blocks.")[1].split(".")[0])
except Exception:
block_idx = None
if "loras" in name or "controlnet" in name:
if "loras" in name:
continue
# GGUF: skip GGUFParameter params
@@ -1310,7 +1310,7 @@ class WanVideoModelLoader:
if dim == 1536:
model_variant = "1_3B"
if dim == 3072:
log.info(f"5B model detected, no Teacache or MagCache coefficients available, consider using EasyCache for this model")
log.info("5B model detected, no Teacache or MagCache coefficients available, consider using EasyCache for this model")
if "high" in model.lower() or "low" in model.lower():
if "i2v" in model.lower():
@@ -1375,7 +1375,7 @@ class WanVideoModelLoader:
with init_empty_weights():
transformer.audio_model = WanModel(**TRANSFORMER_CONFIG).eval()
from .wanvideo.modules.model import WanLayerNorm, WanRMSNorm
from .wanvideo.modules.model import WanLayerNorm
for block in transformer.blocks:
block.cross_attn.k_fusion = nn.Linear(block.dim, block.dim)
@@ -1477,16 +1477,12 @@ class WanVideoModelLoader:
# Additional cond latents
if "add_conv_in.weight" in sd:
def zero_module(module):
for p in module.parameters():
torch.nn.init.zeros_(p)
return module
inner_dim = sd["add_conv_in.weight"].shape[0]
add_cond_in_dim = sd["add_conv_in.weight"].shape[1]
attn_cond_in_dim = sd["attn_conv_in.weight"].shape[1]
transformer.add_conv_in = torch.nn.Conv3d(add_cond_in_dim, inner_dim, kernel_size=transformer.patch_size, stride=transformer.patch_size)
transformer.add_proj = zero_module(torch.nn.Linear(inner_dim, inner_dim))
transformer.attn_conv_in = torch.nn.Conv3d(attn_cond_in_dim, inner_dim, kernel_size=transformer.patch_size, stride=transformer.patch_size)
transformer.add_conv_in = nn.Conv3d(add_cond_in_dim, inner_dim, kernel_size=transformer.patch_size, stride=transformer.patch_size)
transformer.add_proj = nn.Linear(inner_dim, inner_dim)
transformer.attn_conv_in = nn.Conv3d(attn_cond_in_dim, inner_dim, kernel_size=transformer.patch_size, stride=transformer.patch_size)
# Bindweave text_projection
if "text_projection.0.weight" in sd:
@@ -1516,6 +1512,45 @@ class WanVideoModelLoader:
transformer.condition_embedding_align = PoseRefNetNoBNV3(in_channels_x=16, in_channels_c=16, hidden_dim=128, num_heads=8) # Frame-wise Attention Alignment Unit
if "image_to_cond.conv_in.bias" in sd:
# One-to-all
from .onetoall.controlnet import MiniHunyuanEncoder, MiniEncoder2D
from .onetoall.refextractor_2d import WanRefextractor, WanAttentionBlock
log.info("One-to-all model detected, patching model...")
with init_empty_weights():
transformer.image_to_cond = MiniEncoder2D(
in_channels = sd["image_to_cond.conv_in.bias"].shape[0],
out_channels = in_channels,
down_block_types= ("DownEncoderBlockInflated","DownEncoderBlockInflated","DownEncoderBlockInflated"),
block_out_channels=(16, 16, 16),
norm_num_groups = 4,
layers_per_block = 1,
spatial_compression_ratio=1
)
transformer.input_hint_block = MiniHunyuanEncoder(
in_channels=3,
out_channels=in_channels,
block_out_channels=(16, 16, 16, 16),
norm_num_groups=4,
layers_per_block=1,
spatial_compression_ratio=16
)
controlnet_layers = 1
transformer.controlnet = nn.Module()
transformer.controlnet.blocks = nn.ModuleList([WanAttentionBlock(in_features, out_features, ffn_dim, ffn2_dim, num_heads) for _ in range(controlnet_layers)])
transformer.controlnet_zero = nn.ModuleList([nn.Linear(in_features, out_features) for _ in range(controlnet_layers)])
transformer.refextractor = WanRefextractor(
patch_size=(1, 2, 2), in_dim=sd["refextractor.patch_embedding.weight"].shape[1],
dim=dim, in_features=in_features, out_features=out_features, ffn_dim=ffn_dim, ffn2_dim=ffn2_dim,
num_heads=num_heads, num_layers=7)
for block in transformer.blocks:
block.ref_attn_k_img = nn.Linear(in_features, out_features)
block.ref_attn_v_img = nn.Linear(in_features, out_features)
block.ref_attn_norm_k_img = WanRMSNorm(out_features, eps=1e-6)
comfy_model.diffusion_model = transformer
comfy_model.load_device = transformer_load_device
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
@@ -1580,7 +1615,7 @@ class WanVideoModelLoader:
if gguf:
raise ValueError("GGUF models don't support vram management")
from .diffsynth.vram_management import enable_vram_management, AutoWrappedModule, AutoWrappedLinear
from .wanvideo.modules.model import WanLayerNorm, WanRMSNorm
from .wanvideo.modules.model import WanLayerNorm
total_params_in_model = sum(p.numel() for p in patcher.model.diffusion_model.parameters())
log.info(f"Total number of parameters in the loaded model: {total_params_in_model}")