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: 2f97b1be867e64Author: 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: 3e4e4db2369cdbAuthor: 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:
+49
-14
@@ -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}")
|
||||
|
||||
Reference in New Issue
Block a user