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
+10
View File
@@ -84,6 +84,13 @@ except Exception as e:
STEADYDANCER_NODE_CLASS_MAPPINGS = {} STEADYDANCER_NODE_CLASS_MAPPINGS = {}
STEADYDANCER_NODE_DISPLAY_NAME_MAPPINGS = {} STEADYDANCER_NODE_DISPLAY_NAME_MAPPINGS = {}
try:
from .onetoall.nodes import NODE_CLASS_MAPPINGS as ONETOALL_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as ONETOALL_NODE_DISPLAY_NAME_MAPPINGS
except Exception as e:
log.warning(f"WanVideoWrapper WARNING: OneToAll nodes not available due to error in importing them: {e}")
ONETOALL_NODE_CLASS_MAPPINGS = {}
ONETOALL_NODE_DISPLAY_NAME_MAPPINGS = {}
NODE_CLASS_MAPPINGS.update(RECAM_MASTER_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(RECAM_MASTER_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(SKYREELS_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(SKYREELS_NODE_CLASS_MAPPINGS)
@@ -108,6 +115,8 @@ NODE_CLASS_MAPPINGS.update(OVI_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(FLASHVSR_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(FLASHVSR_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(MOCHA_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(MOCHA_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(STEADYDANCER_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(STEADYDANCER_NODE_CLASS_MAPPINGS)
NODE_CLASS_MAPPINGS.update(ONETOALL_NODE_CLASS_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS)
@@ -133,5 +142,6 @@ NODE_DISPLAY_NAME_MAPPINGS.update(OVI_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(FLASHVSR_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(FLASHVSR_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(MOCHA_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(MOCHA_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(STEADYDANCER_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(STEADYDANCER_NODE_DISPLAY_NAME_MAPPINGS)
NODE_DISPLAY_NAME_MAPPINGS.update(ONETOALL_NODE_DISPLAY_NAME_MAPPINGS)
__all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"]
+12 -57
View File
@@ -10,18 +10,13 @@ from diffusers.utils import USE_PEFT_BACKEND, logging, scale_lora_layers, unscal
from diffusers.models.modeling_outputs import Transformer2DModelOutput from diffusers.models.modeling_outputs import Transformer2DModelOutput
from diffusers.models.modeling_utils import ModelMixin from diffusers.models.modeling_utils import ModelMixin
from diffusers.models.transformers.transformer_wan import ( from diffusers.models.transformers.transformer_wan import (
WanTimeTextImageEmbedding, WanTimeTextImageEmbedding,
WanRotaryPosEmbed, WanRotaryPosEmbed,
WanTransformerBlock WanTransformerBlock
) )
logger = logging.get_logger(__name__) # pylint: disable=invalid-name logger = logging.get_logger(__name__) # pylint: disable=invalid-name
def zero_module(module):
for p in module.parameters():
nn.init.zeros_(p)
return module
class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin): class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModelMixin):
r""" r"""
@@ -69,7 +64,7 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
_no_split_modules = ["WanTransformerBlock"] _no_split_modules = ["WanTransformerBlock"]
_keep_in_fp32_modules = ["time_embedder", "scale_shift_table", "norm1", "norm2", "norm3"] _keep_in_fp32_modules = ["time_embedder", "scale_shift_table", "norm1", "norm2", "norm3"]
_keys_to_ignore_on_load_unexpected = ["norm_added_q"] _keys_to_ignore_on_load_unexpected = ["norm_added_q"]
@register_to_config @register_to_config
def __init__( def __init__(
self, self,
@@ -100,10 +95,10 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
## Spatial compression with time awareness ## Spatial compression with time awareness
nn.Sequential( nn.Sequential(
nn.Conv3d( nn.Conv3d(
in_channels, in_channels,
input_channels[0], input_channels[0],
kernel_size=(3, downscale_coef + 1, downscale_coef + 1), kernel_size=(3, downscale_coef + 1, downscale_coef + 1),
stride=(1, downscale_coef, downscale_coef), stride=(1, downscale_coef, downscale_coef),
padding=(1, downscale_coef // 2, downscale_coef // 2) padding=(1, downscale_coef // 2, downscale_coef // 2)
), ),
nn.GELU(approximate="tanh"), nn.GELU(approximate="tanh"),
@@ -122,9 +117,9 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
nn.GroupNorm(2, input_channels[2]), nn.GroupNorm(2, input_channels[2]),
) )
]) ])
inner_dim = num_attention_heads * attention_head_dim inner_dim = num_attention_heads * attention_head_dim
# 1. Patch & position embedding # 1. Patch & position embedding
self.rope = WanRotaryPosEmbed(attention_head_dim, patch_size, rope_max_seq_len) self.rope = WanRotaryPosEmbed(attention_head_dim, patch_size, rope_max_seq_len)
self.patch_embedding = nn.Conv3d(vae_channels + input_channels[2], inner_dim, kernel_size=patch_size, stride=patch_size) self.patch_embedding = nn.Conv3d(vae_channels + input_channels[2], inner_dim, kernel_size=patch_size, stride=patch_size)
@@ -153,11 +148,10 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
for _ in range(len(self.blocks)): for _ in range(len(self.blocks)):
controlnet_block = nn.Linear(inner_dim, out_proj_dim) controlnet_block = nn.Linear(inner_dim, out_proj_dim)
controlnet_block = zero_module(controlnet_block)
self.controlnet_blocks.append(controlnet_block) self.controlnet_blocks.append(controlnet_block)
self.gradient_checkpointing = False self.gradient_checkpointing = False
def forward( def forward(
self, self,
hidden_states: torch.Tensor, hidden_states: torch.Tensor,
@@ -187,7 +181,7 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
# 0. Controlnet encoder # 0. Controlnet encoder
for control_encoder_block in self.control_encoder: for control_encoder_block in self.control_encoder:
controlnet_states = control_encoder_block(controlnet_states) controlnet_states = control_encoder_block(controlnet_states)
hidden_states = torch.cat([hidden_states, controlnet_states], dim=1) hidden_states = torch.cat([hidden_states, controlnet_states], dim=1)
## 1. Patch embedding and stack ## 1. Patch embedding and stack
@@ -216,7 +210,7 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
if encoder_hidden_states_image is not None: if encoder_hidden_states_image is not None:
encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1) encoder_hidden_states = torch.concat([encoder_hidden_states_image, encoder_hidden_states], dim=1)
# 4. Transformer blocks # 4. Transformer blocks
controlnet_hidden_states = () controlnet_hidden_states = ()
if torch.is_grad_enabled() and self.gradient_checkpointing: if torch.is_grad_enabled() and self.gradient_checkpointing:
@@ -239,43 +233,4 @@ class WanControlnet(ModelMixin, ConfigMixin, PeftAdapterMixin, FromOriginalModel
return (controlnet_hidden_states,) return (controlnet_hidden_states,)
return Transformer2DModelOutput(sample=controlnet_hidden_states) return Transformer2DModelOutput(sample=controlnet_hidden_states)
if __name__ == "__main__":
parameters = {
"added_kv_proj_dim": None,
"attention_head_dim": 128,
"cross_attn_norm": True,
"eps": 1e-06,
"ffn_dim": 8960,
"freq_dim": 256,
"image_dim": None,
"in_channels": 3,
"num_attention_heads": 12,
"num_layers": 2,
"patch_size": [1, 2, 2],
"qk_norm": "rms_norm_across_heads",
"rope_max_seq_len": 1024,
"text_dim": 4096,
"downscale_coef": 8,
"out_proj_dim": 12 * 128,
"vae_channels": 16
}
controlnet = WanControlnet(**parameters)
hidden_states = torch.rand(1, 16, 13, 60, 90)
timestep = torch.tensor([1000]).repeat(17550).unsqueeze(0) #torch.randint(low=0, high=1000, size=(1,), dtype=torch.long)
encoder_hidden_states = torch.rand(1, 512, 4096)
controlnet_states = torch.rand(1, 3, 49, 480, 720)
controlnet_hidden_states = controlnet(
hidden_states=hidden_states,
timestep=timestep,
encoder_hidden_states=encoder_hidden_states,
controlnet_states=controlnet_states,
return_dict=False
)
print("Output states count", len(controlnet_hidden_states[0]))
for out_hidden_states in controlnet_hidden_states[0]:
print(out_hidden_states.shape)
File diff suppressed because it is too large Load Diff
+49 -14
View File
@@ -6,7 +6,7 @@ import numpy as np
from tqdm import tqdm from tqdm import tqdm
import re 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.t5 import T5EncoderModel
from .wanvideo.modules.clip import CLIPModel from .wanvideo.modules.clip import CLIPModel
from .wanvideo.wan_video_vae import WanVideoVAE, WanVideoVAE38 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, total=param_count,
leave=True): leave=True):
block_idx = vace_block_idx = None block_idx = vace_block_idx = None
if "vace_blocks." in name: if name.startswith("vace_blocks."):
try: try:
vace_block_idx = int(name.split("vace_blocks.")[1].split(".")[0]) vace_block_idx = int(name.split("vace_blocks.")[1].split(".")[0])
except Exception: except Exception:
vace_block_idx = None vace_block_idx = None
elif "blocks." in name and "face" not in name: elif name.startswith("blocks.") and "face" not in name:
try: try:
block_idx = int(name.split("blocks.")[1].split(".")[0]) block_idx = int(name.split("blocks.")[1].split(".")[0])
except Exception: except Exception:
block_idx = None block_idx = None
if "loras" in name or "controlnet" in name: if "loras" in name:
continue continue
# GGUF: skip GGUFParameter params # GGUF: skip GGUFParameter params
@@ -1310,7 +1310,7 @@ class WanVideoModelLoader:
if dim == 1536: if dim == 1536:
model_variant = "1_3B" model_variant = "1_3B"
if dim == 3072: 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 "high" in model.lower() or "low" in model.lower():
if "i2v" in model.lower(): if "i2v" in model.lower():
@@ -1375,7 +1375,7 @@ class WanVideoModelLoader:
with init_empty_weights(): with init_empty_weights():
transformer.audio_model = WanModel(**TRANSFORMER_CONFIG).eval() 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: for block in transformer.blocks:
block.cross_attn.k_fusion = nn.Linear(block.dim, block.dim) block.cross_attn.k_fusion = nn.Linear(block.dim, block.dim)
@@ -1477,16 +1477,12 @@ class WanVideoModelLoader:
# Additional cond latents # Additional cond latents
if "add_conv_in.weight" in sd: 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] inner_dim = sd["add_conv_in.weight"].shape[0]
add_cond_in_dim = sd["add_conv_in.weight"].shape[1] add_cond_in_dim = sd["add_conv_in.weight"].shape[1]
attn_cond_in_dim = sd["attn_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_conv_in = 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.add_proj = 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.attn_conv_in = nn.Conv3d(attn_cond_in_dim, inner_dim, kernel_size=transformer.patch_size, stride=transformer.patch_size)
# Bindweave text_projection # Bindweave text_projection
if "text_projection.0.weight" in sd: 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 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.diffusion_model = transformer
comfy_model.load_device = transformer_load_device comfy_model.load_device = transformer_load_device
patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device) patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device)
@@ -1580,7 +1615,7 @@ class WanVideoModelLoader:
if gguf: if gguf:
raise ValueError("GGUF models don't support vram management") raise ValueError("GGUF models don't support vram management")
from .diffsynth.vram_management import enable_vram_management, AutoWrappedModule, AutoWrappedLinear 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()) 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}") log.info(f"Total number of parameters in the loaded model: {total_params_in_model}")
+38 -5
View File
@@ -1182,6 +1182,31 @@ class WanVideoSampler:
sdancer_data = sdancer_embeds.copy() sdancer_data = sdancer_embeds.copy()
sdancer_data = dict_to_device(sdancer_data, device, dtype) sdancer_data = dict_to_device(sdancer_data, device, dtype)
# One-to-all-Animation
one_to_all_embeds = image_embeds.get("one_to_all_embeds", None)
one_to_all_data = prev_latents = None
latents_to_not_step = 0
if one_to_all_embeds is not None:
log.info("Using One-to-All embeddings:")
for k, v in one_to_all_embeds.items():
log.info(f" {k}: {v.shape if isinstance(v, torch.Tensor) else v}")
one_to_all_data = one_to_all_embeds.copy()
one_to_all_data = dict_to_device(one_to_all_data, device, dtype)
if one_to_all_embeds.get("pose_images") is not None:
pose_images_in = one_to_all_data.pop("pose_images")
pose_images = transformer.input_hint_block(pose_images_in)
if one_to_all_embeds.get("ref_latent_pos") is not None:
pose_prefix_image = transformer.input_hint_block(one_to_all_data.pop("pose_prefix_image"))
pose_images = torch.cat([pose_prefix_image, pose_images],dim=2)
one_to_all_data["controlnet_tokens"] = pose_images.flatten(2).transpose(1, 2)
prev_latents = one_to_all_data.get("prev_latents", None)
if prev_latents is not None:
log.info(f"Using previous latents for One-to-All Animation with shape: {prev_latents.shape}")
latent[:, :prev_latents.shape[1]] = prev_latents.to(latent)
one_to_all_data["token_replace"] = True
latents_to_not_step = prev_latents.shape[1]
one_to_all_data["num_latent_frames_to_replace"] = latents_to_not_step
#region model pred #region model pred
def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None, def predict_with_cfg(z, cfg_scale, positive_embeds, negative_embeds, timestep, idx, image_cond=None, clip_fea=None,
control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None, control_latents=None, vace_data=None, unianim_data=None, audio_proj=None, control_camera_latents=None,
@@ -1477,6 +1502,7 @@ class WanVideoSampler:
"flashvsr_strength": flashvsr_strength, # FlashVSR strength "flashvsr_strength": flashvsr_strength, # FlashVSR strength
"num_cond_latents": len(all_indices) if transformer.is_longcat else None, "num_cond_latents": len(all_indices) if transformer.is_longcat else None,
"sdancer_input": sdancer_input, # SteadyDancer input "sdancer_input": sdancer_input, # SteadyDancer input
"one_to_all_input": one_to_all_data, # One-to-All input
} }
batch_size = 1 batch_size = 1
@@ -3070,11 +3096,16 @@ class WanVideoSampler:
new_latent.append(latent_slice[:, j:j+1]) new_latent.append(latent_slice[:, j:j+1])
latent = torch.cat(new_latent, dim=1) latent = torch.cat(new_latent, dim=1)
else: else:
latent = sample_scheduler.step( if latents_to_not_step > 0:
noise_pred[:, :orig_noise_len].unsqueeze(0) if recammaster is not None or mocha_embeds is not None else noise_pred.unsqueeze(0), raw_latent = latent[:, :latents_to_not_step]
timestep, noise_pred_in = noise_pred[:, latents_to_not_step:]
latent[:, :orig_noise_len].unsqueeze(0) if recammaster is not None or mocha_embeds is not None else latent.unsqueeze(0), latent = latent[:, latents_to_not_step:]
**scheduler_step_args)[0].squeeze(0) elif recammaster is not None or mocha_embeds is not None:
noise_pred_in = noise_pred[:, :orig_noise_len]
latent = latent[:, :orig_noise_len]
else:
noise_pred_in = noise_pred
latent = sample_scheduler.step(noise_pred_in.unsqueeze(0), timestep, latent.unsqueeze(0), **scheduler_step_args)[0].squeeze(0)
if noise_pred_flipped is not None: if noise_pred_flipped is not None:
latent_backwards = sample_scheduler_flipped.step( latent_backwards = sample_scheduler_flipped.step(
noise_pred_flipped.unsqueeze(0), noise_pred_flipped.unsqueeze(0),
@@ -3083,6 +3114,8 @@ class WanVideoSampler:
**scheduler_step_args)[0].squeeze(0) **scheduler_step_args)[0].squeeze(0)
latent_backwards = torch.flip(latent_backwards, dims=[1]) latent_backwards = torch.flip(latent_backwards, dims=[1])
latent = latent * 0.5 + latent_backwards * 0.5 latent = latent * 0.5 + latent_backwards * 0.5
if latents_to_not_step > 0:
latent = torch.cat([raw_latent, latent], dim=1)
if latent_ovi is not None: if latent_ovi is not None:
latent_ovi = sample_scheduler_ovi.step(noise_pred_ovi.unsqueeze(0), t, latent_ovi.to(device).unsqueeze(0), **scheduler_step_args)[0].squeeze(0) latent_ovi = sample_scheduler_ovi.step(noise_pred_ovi.unsqueeze(0), t, latent_ovi.to(device).unsqueeze(0), **scheduler_step_args)[0].squeeze(0)
+440
View File
@@ -0,0 +1,440 @@
import torch
from torch import nn
from torch.nn import functional as F
from einops import rearrange
import numpy as np
from typing import Tuple
from .unet_causal_3d_blocks import get_down_block3d, CausalConv3d
class ControlNetCausalConditioningEmbedding(nn.Module):
def __init__(self, conditioning_embedding_channels: int, conditioning_channels: int = 3, block_out_channels: Tuple[int, ...] = (16, 32, 96, 256)):
super().__init__()
self.conv_in = CausalConv3d(conditioning_channels, block_out_channels[0], kernel_size=3, padding=1)
self.blocks = nn.ModuleList([])
for i in range(len(block_out_channels) - 1):
channel_in = block_out_channels[i]
channel_out = block_out_channels[i + 1]
self.blocks.append(nn.Conv2d(channel_in, channel_in, kernel_size=3, padding=1))
self.blocks.append(nn.Conv2d(channel_in, channel_out, kernel_size=3, padding=1, stride=2))
self.conv_out = nn.Conv2d(block_out_channels[-1], conditioning_embedding_channels, kernel_size=3, padding=1)
def forward(self, conditioning):
embedding = self.conv_in(conditioning)
embedding = F.silu(embedding)
for block in self.blocks:
embedding = block(embedding)
embedding = F.silu(embedding)
embedding = self.conv_out(embedding)
return embedding
class MiniHunyuanEncoder(nn.Module):
'''
a direct copy of hunyuan encoder
'''
def __init__(
self,
in_channels = 3,
out_channels = 3,
down_block_types = ['DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D'],
block_out_channels = [128, 256, 512, 512],
layers_per_block = 2,
norm_num_groups = 32,
act_fn: str = "silu",
time_compression_ratio: int = 4,
spatial_compression_ratio: int = 8,
):
super().__init__()
self.layers_per_block = layers_per_block
self.conv_in = CausalConv3d(
in_channels, block_out_channels[0], kernel_size=3, stride=1)
self.mid_block = None
self.down_blocks = nn.ModuleList([])
# down
output_channel = block_out_channels[0]
for i, down_block_type in enumerate(down_block_types):
input_channel = output_channel
output_channel = block_out_channels[i]
is_final_block = i == len(block_out_channels) - 1
num_spatial_downsample_layers = int(
np.log2(spatial_compression_ratio))
num_time_downsample_layers = int(np.log2(time_compression_ratio))
if time_compression_ratio == 4:
add_spatial_downsample = bool(
i < num_spatial_downsample_layers)
add_time_downsample = bool(i >= (
len(block_out_channels) - 1 - num_time_downsample_layers) and not is_final_block)
elif time_compression_ratio == 8:
add_spatial_downsample = bool(
i < num_spatial_downsample_layers)
add_time_downsample = bool(i < num_time_downsample_layers)
else:
raise ValueError(
f"Unsupported time_compression_ratio: {time_compression_ratio}")
downsample_stride_HW = (2, 2) if add_spatial_downsample else (1, 1)
downsample_stride_T = (2, ) if add_time_downsample else (1, )
downsample_stride = tuple(
downsample_stride_T + downsample_stride_HW)
down_block = get_down_block3d(
down_block_type,
num_layers=self.layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
add_downsample=bool(
add_spatial_downsample or add_time_downsample),
downsample_stride=downsample_stride,
resnet_eps=1e-6,
downsample_padding=0,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
)
self.down_blocks.append(down_block)
self.conv_out = CausalConv3d(block_out_channels[-1], out_channels, kernel_size=3)
def forward(self, sample):
assert len(sample.shape) == 5, "The input tensor should have 5 dimensions"
sample = self.conv_in(sample)
# down
for down_block in self.down_blocks:
sample = down_block(sample)
sample = self.conv_out(sample)
return sample
class ControlNetConditioningEmbedding(nn.Module):
"""
Quoting from https://arxiv.org/abs/2302.05543: "Stable Diffusion uses a pre-processing method similar to VQ-GAN
[11] to convert the entire dataset of 512 × 512 images into smaller 64 × 64 “latent images” for stabilized
training. This requires ControlNets to convert image-based conditions to 64 × 64 feature space to match the
convolution size. We use a tiny network E(·) of four convolution layers with 4 × 4 kernels and 2 × 2 strides
(activated by ReLU, channels are 16, 32, 64, 128, initialized with Gaussian weights, trained jointly with the full
model) to encode image-space conditions ... into feature maps ..."
"""
def __init__(
self,
conditioning_embedding_channels: int,
conditioning_channels: int = 3,
block_out_channels: Tuple[int, ...] = (16, 32, 96, 256),
):
super().__init__()
self.conv_in = nn.Conv2d(conditioning_channels, block_out_channels[0], kernel_size=3, padding=1)
self.blocks = nn.ModuleList([])
for i in range(len(block_out_channels) - 1):
channel_in = block_out_channels[i]
channel_out = block_out_channels[i + 1]
self.blocks.append(nn.Conv2d(channel_in, channel_in, kernel_size=3, padding=1))
self.blocks.append(nn.Conv2d(channel_in, channel_out, kernel_size=3, padding=1, stride=2))
self.conv_out = nn.Conv2d(block_out_channels[-1], conditioning_embedding_channels, kernel_size=3, padding=1)
def forward(self, conditioning):
embedding = self.conv_in(conditioning)
embedding = F.silu(embedding)
for block in self.blocks:
embedding = block(embedding)
embedding = F.silu(embedding)
embedding = self.conv_out(embedding)
return embedding
class InflatedGroupNorm(nn.GroupNorm):
def forward(self, x):
video_length = x.shape[2]
x = rearrange(x, "b c f h w -> (b f) c h w")
x = super().forward(x)
x = rearrange(x, "(b f) c h w -> b c f h w", f=video_length)
return x
class InflatedConv3d(nn.Conv2d):
def forward(self, x):
video_length = x.shape[2]
x = rearrange(x, "b c f h w -> (b f) c h w")
x = super().forward(x)
x = rearrange(x, "(b f) c h w -> b c f h w", f=video_length)
return x
class ResnetBlockInflated(nn.Module):
def __init__(self, *, in_channels, out_channels=None, dropout=0.0, groups=32, groups_out=None, pre_norm=True, eps=1e-6, non_linearity="swish", output_scale_factor=1.0):
super().__init__()
self.pre_norm = pre_norm
self.pre_norm = True
self.in_channels = in_channels
out_channels = in_channels if out_channels is None else out_channels
self.out_channels = out_channels
self.output_scale_factor = output_scale_factor
if groups_out is None:
groups_out = groups
self.norm1 = InflatedGroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
self.conv1 = InflatedConv3d(in_channels, out_channels, kernel_size=3, stride=1, padding=1)
self.norm2 = InflatedGroupNorm(num_groups=groups_out, num_channels=out_channels, eps=eps, affine=True)
self.dropout = torch.nn.Dropout(dropout)
self.conv2 = InflatedConv3d(out_channels, out_channels, kernel_size=3, stride=1, padding=1)
if non_linearity == "swish":
self.nonlinearity = lambda x: F.silu(x)
elif non_linearity == "silu":
self.nonlinearity = nn.SiLU()
def forward(self, input_tensor, temb):
if temb is not None:
print("Warning: temb is None in ResnetBlockInflated")
hidden_states = input_tensor
hidden_states = self.norm1(hidden_states)
hidden_states = self.nonlinearity(hidden_states)
hidden_states = self.conv1(hidden_states)
if temb is not None:
hidden_states = hidden_states + temb
hidden_states = self.norm2(hidden_states)
hidden_states = self.nonlinearity(hidden_states)
hidden_states = self.dropout(hidden_states)
hidden_states = self.conv2(hidden_states)
output_tensor = (input_tensor + hidden_states) / self.output_scale_factor
return output_tensor
class DownEncoderBlockInflated(nn.Module):
def __init__(self, *, num_layers: int, in_channels: int, out_channels: int, add_downsample: bool, downsample_stride: tuple = (1, 2, 2),
resnet_eps: float = 1e-6, resnet_act_fn: str = "silu", resnet_groups: int = 32):
super().__init__()
self.resnets = nn.ModuleList([ResnetBlockInflated(
in_channels=in_channels if i == 0 else out_channels,
out_channels=out_channels,
eps=resnet_eps,
non_linearity=resnet_act_fn,
groups=resnet_groups,
) for i in range(num_layers)])
self.downsamplers = nn.ModuleList()
if add_downsample:
self.downsamplers.append(
InflatedConv3d(
out_channels,
out_channels,
kernel_size=3,
stride=2,
padding=1,
)
)
self.down_stride = downsample_stride
else:
self.down_stride = (1, 1, 1)
def forward(self, x, temb=None):
for resnet in self.resnets:
x = resnet(x, temb)
for down in self.downsamplers:
x = down(x)
return x
class SFT(nn.Module): # 2D SFT
def __init__(
self, in_channels, out_channels, intermediate_channels=128, groups=32, eps=1e-6):
super().__init__()
self.out_channels = out_channels
self.norm = InflatedGroupNorm(groups, out_channels, eps, affine=True)
self.mlp_shared = nn.Sequential(InflatedConv3d(in_channels, intermediate_channels, kernel_size=3, stride=1, padding=1), nn.SiLU())
self.mlp_gamma = InflatedConv3d(intermediate_channels, out_channels, kernel_size=3, stride=1, padding=1)
self.mlp_beta = InflatedConv3d(intermediate_channels, out_channels, kernel_size=3, stride=1, padding=1)
def forward(self, hidden_state, condition):
"""
hidden_state : (B, Cout, T, H, W)
condition : (B, Cin, 1, H, W)
"""
hidden_state = self.norm(hidden_state) #2D SFT 2D Norm
actv = self.mlp_shared(condition)
gamma = self.mlp_gamma(actv)
beta = self.mlp_beta(actv)
return torch.addcmul(beta, hidden_state, 1 + gamma)
class MiniEncoder2D(nn.Module):
def __init__(
self,
in_channels: int = 3,
out_channels: int = 3,
down_block_types: list = (
"DownEncoderBlockInflated",
"DownEncoderBlockInflated",
"DownEncoderBlockInflated",
"DownEncoderBlockInflated",
),
block_out_channels: list = (128, 256, 512, 512),
layers_per_block: int = 2,
norm_num_groups: int = 32,
act_fn: str = "silu",
spatial_compression_ratio: int = 8,
):
super().__init__()
# -------------------------------------------------------------------
# conv in
# -------------------------------------------------------------------
self.conv_in = InflatedConv3d(in_channels, block_out_channels[0], kernel_size=3, stride=1, padding=1)
self.down_blocks = nn.ModuleList()
output_channel = block_out_channels[0]
num_spatial_down_layers = int(np.log2(spatial_compression_ratio))
for i, block_type in enumerate(down_block_types):
input_channel = output_channel
output_channel = block_out_channels[i]
# is_final_block = i == len(block_out_channels) - 1
add_spatial_downsample = bool(i < num_spatial_down_layers)
downsample_stride = (1, 2, 2) if add_spatial_downsample else (1, 1, 1)
down_block = DownEncoderBlockInflated(
num_layers=layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
add_downsample=add_spatial_downsample,
downsample_stride=downsample_stride,
resnet_eps=1e-6,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
)
self.down_blocks.append(down_block)
self.conv_out = InflatedConv3d(output_channel, out_channels, kernel_size=3, stride=1, padding=1)
def forward(self, x):
# (B,C,1,H,W)
x = self.conv_in(x)
for block in self.down_blocks:
x = block(x)
return self.conv_out(x)
class Driven_Ref_PoseEncoder(nn.Module):
def __init__(
self, in_channels = 3, out_channels = 3,
down_block_types = ['DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D', 'DownEncoderBlockCausal3D'],
block_out_channels = [128, 256, 512, 512], layers_per_block = 2, norm_num_groups = 32,
act_fn: str = "silu", time_compression_ratio: int = 4, spatial_compression_ratio: int = 8,
):
super().__init__()
self.layers_per_block = layers_per_block
self.conv_in = CausalConv3d(in_channels, block_out_channels[0], kernel_size=3, stride=1)
self.mid_block = None
self.down_blocks = nn.ModuleList([])
# down
output_channel = block_out_channels[0]
for i, down_block_type in enumerate(down_block_types):
input_channel = output_channel
output_channel = block_out_channels[i]
is_final_block = i == len(block_out_channels) - 1
num_spatial_downsample_layers = int(
np.log2(spatial_compression_ratio))
num_time_downsample_layers = int(np.log2(time_compression_ratio))
if time_compression_ratio == 4:
add_spatial_downsample = bool(
i < num_spatial_downsample_layers)
add_time_downsample = bool(i >= (
len(block_out_channels) - 1 - num_time_downsample_layers) and not is_final_block)
elif time_compression_ratio == 8:
add_spatial_downsample = bool(
i < num_spatial_downsample_layers)
add_time_downsample = bool(i < num_time_downsample_layers)
else:
raise ValueError(
f"Unsupported time_compression_ratio: {time_compression_ratio}")
downsample_stride_HW = (2, 2) if add_spatial_downsample else (1, 1)
downsample_stride_T = (2, ) if add_time_downsample else (1, )
downsample_stride = tuple(
downsample_stride_T + downsample_stride_HW)
down_block = get_down_block3d(
down_block_type,
num_layers=self.layers_per_block,
in_channels=input_channel,
out_channels=output_channel,
add_downsample=bool(
add_spatial_downsample or add_time_downsample),
downsample_stride=downsample_stride,
resnet_eps=1e-6,
downsample_padding=0,
resnet_act_fn=act_fn,
resnet_groups=norm_num_groups,
attention_head_dim=output_channel,
)
self.down_blocks.append(down_block)
self.conv_out = CausalConv3d(block_out_channels[-1], out_channels, kernel_size=3)
self.ref_pose_encoder = MiniEncoder2D(
in_channels = in_channels,
out_channels = out_channels,
block_out_channels = block_out_channels,
norm_num_groups = norm_num_groups,
layers_per_block = layers_per_block,
spatial_compression_ratio = spatial_compression_ratio,
)
self.sft_layers = nn.ModuleList()
for i, ch in enumerate(block_out_channels):
if i == 0: # 0 层 (H/2,W/2) 不做 SFT
self.sft_layers.append(None)
else: # H/4、H/8、H/16 做 SFT
self.sft_layers.append(
SFT(
in_channels=ch,
out_channels=ch,
intermediate_channels=max(8, ch // 2),
groups=norm_num_groups,
)
)
def forward(self, driven_pose, ref_pose):
# driven_pose b c t h w
# ref_pose b c 1 h w
ref_pose_cond, ref_feats = self.ref_pose_encoder(ref_pose)
x = self.conv_in(driven_pose)
for i, down_block in enumerate(self.down_blocks):
x = down_block(x)
if self.sft_layers[i] is not None:
cond_feat = ref_feats[i]
x = self.sft_layers[i](x, cond_feat)
driven_pose_cond = self.conv_out(x)
return driven_pose_cond, ref_pose_cond
+150
View File
@@ -0,0 +1,150 @@
import torch
from ..utils import log
import comfy.model_management as mm
device = mm.get_torch_device()
offload_device = mm.unet_offload_device()
class WanVideoAddOneToAllReferenceEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"vae": ("WANVAE", {"tooltip": "VAE model"}),
"ref_image": ("IMAGE",),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the reference embedding"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the embedding application"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the embedding application"}),
},
"optional": {
"ref_mask": ("MASK",),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, vae, ref_image, strength, start_percent, end_percent, ref_mask=None):
updated = dict(embeds)
ref_latent = ref_latent_empty = None
vae.to(device)
ref_image_in = (ref_image[..., :3].permute(3, 0, 1, 2) * 2 - 1).to(device, vae.dtype)
ref_latent = vae.encode([ref_image_in], device, tiled=False)
ref_mask_in = None
if ref_mask is not None:
ref_mask_in = (ref_mask.unsqueeze(0).repeat(3, 1, 1, 1) * 2 - 1.).to(device, vae.dtype)
else:
ref_mask_in = torch.zeros_like(ref_image_in)-1
ref_mask_latent = vae.encode([ref_mask_in], device, tiled=False)
if ref_mask is not None and not torch.all(ref_mask == 0):
ref_latent_empty = vae.encode([torch.zeros_like(ref_image_in)-1], device, tiled=False)
else:
ref_latent_empty = ref_mask_latent
vae.to(offload_device)
updated.setdefault("one_to_all_embeds", {})
updated["one_to_all_embeds"]["ref_latent_pos"] = torch.cat([ref_latent, ref_latent_empty], dim=1)
updated["one_to_all_embeds"]["ref_latent_neg"] = torch.cat([ref_latent_empty, ref_latent_empty], dim=1)
updated["one_to_all_embeds"]["ref_strength"] = strength
updated["one_to_all_embeds"]["ref_start_percent"] = start_percent
updated["one_to_all_embeds"]["ref_end_percent"] = end_percent
return (updated,)
class WanVideoAddOneToAllPoseEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"pose_images": ("IMAGE", {"tooltip": "Pose images for the entire video"}),
"strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Strength of the pose control"}),
"start_percent": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "Start percentage of the pose control application"}),
"end_percent": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "End percentage of the pose control application"}),
},
"optional": {
"pose_prefix_image": ("IMAGE",),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",)
RETURN_NAMES = ("image_embeds",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, pose_images, strength, pose_prefix_image=None, start_percent=0.0, end_percent=1.0):
updated = dict(embeds)
updated.setdefault("one_to_all_embeds", {})
pose_images_in = pose_images[..., :3].unsqueeze(0).permute(0, 4, 1, 2, 3) * 2 - 1 # 1 B H W C -> B C 1 H W
updated["one_to_all_embeds"]["pose_images"] = pose_images_in
if pose_prefix_image is not None:
updated["one_to_all_embeds"]["pose_prefix_image"] = pose_prefix_image.unsqueeze(0).permute(0, 4, 1, 2, 3) * 2 - 1 # 1 B H W C -> B C 1 H W
else:
updated["one_to_all_embeds"]["pose_prefix_image"] = pose_images_in[:, :, :1]
updated["one_to_all_embeds"]["controlnet_strength"] = strength
updated["one_to_all_embeds"]["controlnet_start_percent"] = start_percent
updated["one_to_all_embeds"]["controlnet_end_percent"] = end_percent
return (updated,)
class WanVideoAddOneToAllExtendEmbeds:
@classmethod
def INPUT_TYPES(s):
return {"required": {
"embeds": ("WANVIDIMAGE_EMBEDS",),
"prev_latents": ("LATENT", {"tooltip": "Previous latents to be used to continue generation"}),
"window_size": ("INT", {"default": 81, "min": 1, "max": 256, "step": 1, "tooltip": "Number of new frames to generate" }),
"overlap": ("INT", {"default": 5, "min": 0, "max": 64, "step": 1, "tooltip": "Number of overlapping frames between previous and new frames" }),
"frames_processed": ("INT", {"default": 0, "min": 0, "max": 10000, "step": 1, "tooltip": "Number of frames already processed in the video" }),
"if_not_enough_frames": (["pad_with_last", "error"], {"default": "pad_with_last", "tooltip": "What to do if there are not enough frames in pose_images for the window"}),
},
"optional": {
"pose_images": ("IMAGE", {"tooltip": "Pose images for the entire video"}),
}
}
RETURN_TYPES = ("WANVIDIMAGE_EMBEDS", "IMAGE",)
RETURN_NAMES = ("image_embeds", "pose_slice",)
FUNCTION = "add"
CATEGORY = "WanVideoWrapper"
def add(self, embeds, prev_latents, if_not_enough_frames, window_size=81, overlap=5, frames_processed=0, pose_images=None):
updated = dict(embeds)
updated.setdefault("one_to_all_embeds", {})
updated["one_to_all_embeds"]["prev_latents"] = prev_latents["samples"][0]
if pose_images is not None:
pose_images_in = pose_images.clone()[..., :3]
start = max(0, frames_processed - overlap)
end = start + window_size
log.info(f"Extracting pose images from {start} to {end}")
if start >= pose_images_in.shape[0]:
raise ValueError(f"start index {start} exceeds pose images length {pose_images_in.shape[0]}")
if end > pose_images_in.shape[0]:
if if_not_enough_frames == "pad_with_last":
padding_needed = end - pose_images_in.shape[0]
pose_images_in = torch.cat([pose_images_in, pose_images_in[-1:].repeat(padding_needed, 1, 1, 1)], dim=0)
log.info(f"Not enough frames, padding with {padding_needed} frames to reach {end} total frames")
else:
raise ValueError(f"end index {end} exceeds pose images length {pose_images.shape[0]}")
pose_slice = pose_images_in[start:end]
else:
pose_slice = torch.zeros((1, 64, 64, 3))
return (updated, pose_slice)
NODE_CLASS_MAPPINGS = {
"WanVideoAddOneToAllReferenceEmbeds": WanVideoAddOneToAllReferenceEmbeds,
"WanVideoAddOneToAllPoseEmbeds": WanVideoAddOneToAllPoseEmbeds,
"WanVideoAddOneToAllExtendEmbeds": WanVideoAddOneToAllExtendEmbeds,
}
NODE_DISPLAY_NAME_MAPPINGS = {
"WanVideoAddOneToAllReferenceEmbeds": "WanVideo Add OneToAll Reference Embeds",
"WanVideoAddOneToAllPoseEmbeds": "WanVideo Add OneToAll Pose Embeds",
"WanVideoAddOneToAllExtendEmbeds": "WanVideo Add OneToAll Extend Embeds",
}
+220
View File
@@ -0,0 +1,220 @@
from typing import Dict, Union
import torch
import torch.nn as nn
from ..wanvideo.modules.model import WanLayerNorm, WanSelfAttention, EmbedND_RifleX, sinusoidal_embedding_1d, apply_rotary_emb_split, apply_rope_comfy1
class WanAttentionBlock(nn.Module):
def __init__(self, in_features, out_features, ffn_dim, ffn2_dim, num_heads, qk_norm=True, cross_attn_norm=False, eps=1e-6, attention_mode="sdpa", rope_func="comfy", rms_norm_function="default"):
super().__init__()
self.dim = out_features
self.ffn_dim = ffn_dim
self.num_heads = num_heads
self.head_dim = out_features // num_heads
self.qk_norm = qk_norm
self.cross_attn_norm = cross_attn_norm
self.eps = eps
self.attention_mode = attention_mode
self.rope_func = rope_func
# layers
self.norm1 = WanLayerNorm(self.dim, eps)
self.self_attn = WanSelfAttention(in_features, out_features, num_heads, qk_norm, eps, self.attention_mode, rms_norm_function=rms_norm_function, head_norm=False)
self.norm2 = WanLayerNorm(self.dim, eps)
self.ffn = nn.Sequential(nn.Linear(in_features, ffn_dim), nn.GELU(approximate='tanh'), nn.Linear(ffn2_dim, out_features))
self.modulation = nn.Parameter(torch.randn(1, 6, out_features) / in_features**0.5)
def get_mod(self, e, modulation):
if e.dim() == 3:
if e.shape[-1] == 512:
e = self.modulation(e)
return e.unsqueeze(2).chunk(6, dim=-1)
return (modulation + e).chunk(6, dim=1) # 1, 6, dim
elif e.dim() == 4:
e_mod = modulation.unsqueeze(2) + e
return [ei.squeeze(1) for ei in e_mod.unbind(dim=1)]
def modulate(self, norm_x, shift_msa, scale_msa):
return torch.addcmul(shift_msa, norm_x, 1 + scale_msa)
def ffn_chunked(self, mod_x, num_chunks=4):
seq_len = mod_x.shape[1]
if seq_len <= 8192 or num_chunks <= 1:
return self.ffn(mod_x)
return torch.cat([self.ffn(chunk.contiguous()) for chunk in mod_x.chunk(num_chunks, dim=1)], dim=1)
#region attention forward
def forward(self, x, e, seq_lens, freqs, split_rope=True, e_tr=None, tr_start=0, tr_num=0):
use_token_replace = False
if e_tr is not None and tr_num > 0:
tr_shift_msa, tr_scale_msa, tr_gate_msa, tr_shift_mlp, tr_scale_mlp, tr_gate_mlp = self.get_mod(e_tr.to(x.device), self.modulation)
use_token_replace = True
tr_start = tr_start or 0
tr_end = tr_start + (tr_num or 0)
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.get_mod(e.to(x.device), self.modulation)
del e
input_dtype = x.dtype
if use_token_replace:
norm_x = self.norm1(x.to(shift_msa.dtype))
input_x = torch.cat([
torch.addcmul(shift_msa, norm_x[:, :tr_start], 1 + scale_msa), # before replace → T
torch.addcmul(tr_shift_msa, norm_x[:, tr_start:tr_end], 1 + tr_scale_msa), # replace segment → t=0
torch.addcmul(shift_msa, norm_x[:, tr_end:], 1 + scale_msa) # after replace → T
], dim=1).to(input_dtype)
else:
input_x = self.modulate(self.norm1(x.to(shift_msa.dtype)), shift_msa, scale_msa).to(input_dtype)
del shift_msa, scale_msa
b, s, n, d = *x.shape[:2], self.self_attn.num_heads, self.self_attn.head_dim
h_dim = w_dim = 2 * (self.head_dim // 6)
t_dim = self.head_dim - h_dim - w_dim
q = self.self_attn.norm_q(self.self_attn.q(input_x)).to(self.self_attn.norm_q.weight.dtype).view(b, s, n, d)
if split_rope:
q = apply_rotary_emb_split(q, freqs, t_dim) # Apply split rotary embedding (only to H/W dimensions, leaving T unchanged)
else:
q = apply_rope_comfy1(q, freqs)
k = self.self_attn.norm_k(self.self_attn.k(input_x).to(self.self_attn.norm_k.weight.dtype)).to(input_x.dtype).view(b, s, n, d)
if split_rope:
k = apply_rotary_emb_split(k, freqs, t_dim)
else:
k = apply_rope_comfy1(k, freqs)
v = self.self_attn.v(input_x).view(b, s, n, d)
del input_x
y = self.self_attn.forward(q, k, v, seq_lens)
del q, k, v
if use_token_replace:
x = x + torch.cat([
y[:, :tr_start] * gate_msa,
y[:, tr_start:tr_end] * tr_gate_msa,
y[:, tr_end:] * gate_msa
], dim=1).to(input_dtype)
else:
x = x.addcmul(y, gate_msa)
del y, gate_msa
# ffn
if use_token_replace:
norm2_x = self.norm2(x.to(shift_mlp.dtype))
mod_x = torch.cat([
torch.addcmul(shift_mlp, norm2_x[:, :tr_start], 1 + scale_mlp),
torch.addcmul(tr_shift_mlp, norm2_x[:, tr_start:tr_end], 1 + tr_scale_mlp),
torch.addcmul(shift_mlp, norm2_x[:, tr_end:], 1 + scale_mlp)
], dim=1)
else:
mod_x = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp)
del shift_mlp, scale_mlp
x_ffn = self.ffn_chunked(mod_x.to(input_dtype), num_chunks=1)
del mod_x
# gate_mlp
if use_token_replace:
x = x + torch.cat([
x_ffn[:, :tr_start] * gate_mlp,
x_ffn[:, tr_start:tr_end] * tr_gate_mlp,
x_ffn[:, tr_end:] * gate_mlp
], dim=1).to(input_dtype)
else:
x = x.addcmul(x_ffn.to(gate_mlp.dtype), gate_mlp).to(input_dtype)
del gate_mlp
return x
class WanRefextractor(nn.Module):
def __init__(self, patch_size=(1, 2, 2), in_dim=16, dim=5120, in_features=5120, out_features=5120, ffn_dim=8192, ffn2_dim=8192,
freq_dim=256, num_heads=16, num_layers=32, eps=1e-6,
qk_norm=True, cross_attn_norm=True,
attention_mode='sdpa', rope_func='comfy', rms_norm_function='default',
main_device=torch.device('cuda'), offload_device=torch.device('cpu'), dtype=torch.float16):
super().__init__()
self.patch_size = patch_size
self.freq_dim = freq_dim
self.dim = dim
self.main_device = main_device
self.base_dtype = dtype
self.attention_mode = attention_mode
self.patch_embedding = nn.Conv3d(in_dim, dim, kernel_size=patch_size, stride=patch_size)
self.time_embedding = nn.Sequential(nn.Linear(freq_dim, dim), nn.SiLU(), nn.Linear(dim, dim))
self.time_projection = nn.Sequential(nn.SiLU(), nn.Linear(dim, dim * 6))
self.blocks = nn.ModuleList([
WanAttentionBlock(in_features, out_features, ffn_dim, ffn2_dim, num_heads,
qk_norm, cross_attn_norm, eps, attention_mode="sdpa", rope_func=rope_func, rms_norm_function=rms_norm_function)
for i in range(num_layers)
])
self.ref_blocks = nn.ModuleList([])
for _ in range(len(self.blocks)+1):
self.ref_blocks.append(nn.Linear(in_features, out_features))
d = dim // num_heads
self.rope_embedder = EmbedND_RifleX(d,10000.0, [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)], num_frames=1, k=0)
def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None):
patch_size = self.patch_size
t_len = ((t + (patch_size[0] // 2)) // patch_size[0])
h_len = ((h + (patch_size[1] // 2)) // patch_size[1])
w_len = ((w + (patch_size[2] // 2)) // patch_size[2])
if steps_t is None:
steps_t = t_len
if steps_h is None:
steps_h = h_len
if steps_w is None:
steps_w = w_len
img_ids = torch.zeros((steps_t, steps_h, steps_w, 3), device=device, dtype=dtype)
img_ids[:, :, :, 0] = img_ids[:, :, :, 0] + torch.linspace(t_start+freq_offset, t_start + (t_len - 1), steps=steps_t, device=device, dtype=dtype).reshape(-1, 1, 1)
img_ids[:, :, :, 1] = img_ids[:, :, :, 1] + torch.linspace(freq_offset, h_len - 1, steps=steps_h, device=device, dtype=dtype).reshape(1, -1, 1)
img_ids[:, :, :, 2] = img_ids[:, :, :, 2] + torch.linspace(freq_offset, w_len - 1, steps=steps_w, device=device, dtype=dtype).reshape(1, 1, -1)
img_ids = img_ids.reshape(1, -1, img_ids.shape[-1])
freqs = self.rope_embedder(img_ids, ntk_alphas).movedim(1, 2)
return freqs
def forward(
self,
x: torch.Tensor,
timestep: torch.LongTensor,
) -> Union[torch.Tensor, Dict[str, torch.Tensor]]:
B, C, F, H, W = x.shape
freqs = self.rope_encode_comfy(F, H, W, device=x.device, dtype=x.dtype)
self.patch_embedding.to(self.main_device)
x = self.patch_embedding(x.float()).to(x.dtype).flatten(2).transpose(1, 2).to(self.base_dtype)
seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.int32)
time_embed_dtype = self.time_embedding[0].weight.dtype
if time_embed_dtype not in [torch.float16, torch.bfloat16, torch.float32]:
time_embed_dtype = self.base_dtype
e = self.time_embedding(sinusoidal_embedding_1d(self.freq_dim, timestep.flatten()).to(time_embed_dtype)) # b, dim
e0 = self.time_projection(e).unflatten(1, (6, self.dim)).to(self.base_dtype) # b, 6, dim
del e
# 4. Transformer blocks
block_samples = ()
for block in self.blocks:
block_samples = block_samples + (x, )
x = block(x, e0, seq_lens, freqs)
block_samples = block_samples + (x, )
ref_block_samples = ()
for block_sample, ref_block in zip(block_samples, self.ref_blocks):
block_sample = ref_block(block_sample)
ref_block_samples = ref_block_samples + (block_sample, )
return ref_block_samples, freqs
+144
View File
@@ -0,0 +1,144 @@
# Copyright 2024 The HuggingFace Team. All rights reserved.
#
# Licensed under the Apache License, Version 2.0 (the "License");
# you may not use this file except in compliance with the License.
# You may obtain a copy of the License at
#
# http://www.apache.org/licenses/LICENSE-2.0
#
# Unless required by applicable law or agreed to in writing, software
# distributed under the License is distributed on an "AS IS" BASIS,
# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied.
# See the License for the specific language governing permissions and
# limitations under the License.
# ==============================================================================
#
# Modified from diffusers==0.29.2
#
# ==============================================================================
from typing import Optional
import torch
import torch.nn.functional as F
from torch import nn
import comfy.ops
ops = comfy.ops.disable_weight_init
def prepare_causal_attention_mask(n_frame: int, n_hw: int, dtype, device, batch_size: int = None):
seq_len = n_frame * n_hw
mask = torch.full((seq_len, seq_len), float(
"-inf"), dtype=dtype, device=device)
for i in range(seq_len):
i_frame = i // n_hw
mask[i, : (i_frame + 1) * n_hw] = 0
if batch_size is not None:
mask = mask.unsqueeze(0).expand(batch_size, -1, -1)
return mask
class CausalConv3d(nn.Module):
def __init__(self, chan_in, chan_out, kernel_size, stride = 1, dilation = 1, pad_mode='replicate', **kwargs):
super().__init__()
self.pad_mode = pad_mode
padding = (kernel_size // 2, kernel_size // 2, kernel_size // 2, kernel_size // 2, kernel_size - 1, 0) # W, H, T
self.time_causal_padding = padding
self.conv = ops.Conv3d(chan_in, chan_out, kernel_size, stride=stride, dilation=dilation, **kwargs)
def forward(self, x):
x = F.pad(x, self.time_causal_padding, mode=self.pad_mode)
return self.conv(x)
class DownsampleCausal3D(nn.Module):
def __init__(self, channels, use_conv=False, out_channels=None, padding=1, name="conv", kernel_size=3, bias=True, stride=2):
super().__init__()
self.channels, self.out_channels, self.use_conv, self.padding, self.name = channels, out_channels or channels, use_conv, padding, name
self.conv = CausalConv3d(self.channels, self.out_channels, kernel_size=kernel_size, stride=stride, bias=bias)
def forward(self, x, scale=1.0):
return self.conv(x)
class ResnetBlockCausal3D(nn.Module):
def __init__(self, *, in_channels: int, out_channels: Optional[int] = None, groups: int = 32, eps: float = 1e-6, conv_3d_out_channels: Optional[int] = None):
super().__init__()
self.in_channels = in_channels
out_channels = in_channels if out_channels is None else out_channels
self.out_channels = out_channels
self.norm1 = torch.nn.GroupNorm(num_groups=groups, num_channels=in_channels, eps=eps, affine=True)
self.norm2 = torch.nn.GroupNorm(num_groups=groups, num_channels=out_channels, eps=eps, affine=True)
self.conv1 = CausalConv3d(in_channels, out_channels, kernel_size=3, stride=1)
conv_3d_out_channels = conv_3d_out_channels or out_channels
self.conv2 = CausalConv3d(out_channels, conv_3d_out_channels, kernel_size=3, stride=1)
def forward(self, input_tensor: torch.FloatTensor, temb: torch.FloatTensor, scale: float = 1.0) -> torch.FloatTensor:
hidden_states = input_tensor
hidden_states = self.conv1(nn.SiLU()(self.norm1(hidden_states)))
if temb is not None:
hidden_states = hidden_states + temb
hidden_states = self.conv2(nn.SiLU()(self.norm2(hidden_states)))
return input_tensor + hidden_states
def get_down_block3d(down_block_type: str, num_layers: int, in_channels: int, out_channels: int,
add_downsample: bool, downsample_stride: int, resnet_eps: float, resnet_act_fn: str, resnet_groups: Optional[int] = None,
downsample_padding: Optional[int] = None, **kwargs):
down_block_type = down_block_type[7:] if down_block_type.startswith(
"UNetRes") else down_block_type
if down_block_type == "DownEncoderBlockCausal3D":
return DownEncoderBlockCausal3D(
num_layers=num_layers,
in_channels=in_channels,
out_channels=out_channels,
add_downsample=add_downsample,
downsample_stride=downsample_stride,
resnet_eps=resnet_eps,
resnet_act_fn=resnet_act_fn,
resnet_groups=resnet_groups,
downsample_padding=downsample_padding,
)
raise ValueError(f"{down_block_type} does not exist.")
class DownEncoderBlockCausal3D(nn.Module):
def __init__(self, in_channels: int, out_channels: int, num_layers: int = 1, resnet_eps: float = 1e-6,
resnet_groups: int = 32, add_downsample: bool = True, downsample_stride: int = 2, downsample_padding: int = 1, **kwargs):
super().__init__()
resnets = []
for i in range(num_layers):
in_channels = in_channels if i == 0 else out_channels
resnets.append(
ResnetBlockCausal3D(
in_channels=in_channels,
out_channels=out_channels,
eps=resnet_eps,
groups=resnet_groups,
)
)
self.resnets = nn.ModuleList(resnets)
if add_downsample:
self.downsamplers = nn.ModuleList([DownsampleCausal3D(
out_channels,
use_conv=True,
out_channels=out_channels,
padding=downsample_padding,
name="op",
stride=downsample_stride,
)])
else:
self.downsamplers = None
def forward(self, hidden_states: torch.FloatTensor, scale: float = 1.0) -> torch.FloatTensor:
for resnet in self.resnets:
hidden_states = resnet(hidden_states, temb=None, scale=scale)
if self.downsamplers is not None:
for downsampler in self.downsamplers:
hidden_states = downsampler(hidden_states, scale)
return hidden_states
+3 -10
View File
@@ -115,16 +115,9 @@ class WanRotaryPosEmbed(nn.Module):
freqs_w = freqs[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1) freqs_w = freqs[2][:ppw].view(1, 1, ppw, -1).expand(ppf, pph, ppw, -1)
freqs = torch.cat([freqs_f, freqs_h, freqs_w], dim=-1).reshape(1, 1, ppf * pph * ppw, -1) freqs = torch.cat([freqs_f, freqs_h, freqs_w], dim=-1).reshape(1, 1, ppf * pph * ppw, -1)
return freqs return freqs
from ..wanvideo.modules.attention import sageattn_func from ..wanvideo.modules.attention import sageattn_func
def zero_module(module):
# Zero out the parameters of a module and return it.
for p in module.parameters():
p.detach().zero_()
return module
class SimpleAttnProcessor2_0: class SimpleAttnProcessor2_0:
def __init__(self, attention_mode): def __init__(self, attention_mode):
self.attention_mode = attention_mode self.attention_mode = attention_mode
@@ -278,7 +271,7 @@ class MaskCamEmbed(nn.Module):
mid_channels = controlnet_cfg.get("mid_channels", 64) mid_channels = controlnet_cfg.get("mid_channels", 64)
self.mask_proj = nn.Sequential(nn.Conv3d(add_channels, mid_channels, kernel_size=(4, 8, 8), stride=(4, 8, 8)), self.mask_proj = nn.Sequential(nn.Conv3d(add_channels, mid_channels, kernel_size=(4, 8, 8), stride=(4, 8, 8)),
nn.GroupNorm(mid_channels // 8, mid_channels), nn.SiLU()) nn.GroupNorm(mid_channels // 8, mid_channels), nn.SiLU())
self.mask_zero_proj = zero_module(nn.Conv3d(mid_channels, controlnet_cfg["conv_out_dim"], kernel_size=(1, 2, 2), stride=(1, 2, 2))) self.mask_zero_proj = nn.Conv3d(mid_channels, controlnet_cfg["conv_out_dim"], kernel_size=(1, 2, 2), stride=(1, 2, 2))
def forward(self, add_inputs: torch.Tensor): def forward(self, add_inputs: torch.Tensor):
# render_mask.shape [b,c,f,h,w] # render_mask.shape [b,c,f,h,w]
@@ -321,7 +314,7 @@ class WanControlNet(ModelMixin):
) )
self.proj_out = nn.ModuleList( self.proj_out = nn.ModuleList(
[ [
zero_module(nn.Linear(self.dim, 5120)) nn.Linear(self.dim, 5120)
for _ in range(controlnet_cfg["num_layers"]) for _ in range(controlnet_cfg["num_layers"])
] ]
) )
+224 -80
View File
@@ -29,6 +29,18 @@ from comfy import model_management as mm
__all__ = ['WanModel'] __all__ = ['WanModel']
def apply_rotary_emb_split(hidden_states, freqs_cis, t_dim):
"""Apply rotary embedding only to the spatial (H/W) dimensions, leaving temporal (T) unchanged."""
t_part, hw_part = torch.split(hidden_states, [t_dim, hidden_states.shape[-1] - t_dim], dim=-1)
hw_freqs = freqs_cis[..., t_dim//2:, :, :]
x_ = hw_part.to(dtype=hw_freqs.dtype).reshape(*hw_part.shape[:-1], -1, 1, 2)
x_out = hw_freqs[..., 0] * x_[..., 0]
x_out.addcmul_(hw_freqs[..., 1], x_[..., 1])
out_hw = x_out.reshape(*hw_part.shape).type_as(hidden_states)
return torch.cat([t_part, out_hw], dim=-1)
class AdaLayerNorm(nn.Module): class AdaLayerNorm(nn.Module):
def __init__(self, embedding_dim, output_dim=None, norm_elementwise_affine=False, norm_eps=1e-5): def __init__(self, embedding_dim, output_dim=None, norm_elementwise_affine=False, norm_eps=1e-5):
super().__init__() super().__init__()
@@ -100,11 +112,6 @@ class FramePackMotioner(nn.Module):#from comfy.ldm.wan.model
rope = torch.cat([rope_post, rope_2x, rope_4x], dim=1) rope = torch.cat([rope_post, rope_2x, rope_4x], dim=1)
return motion_lat, rope return motion_lat, rope
def zero_module(module):
for p in module.parameters():
p.detach().zero_()
return module
def torch_dfs(model: nn.Module, parent_name='root'): def torch_dfs(model: nn.Module, parent_name='root'):
module_names, modules = [], [] module_names, modules = [], []
current_name = parent_name if parent_name else 'root' current_name = parent_name if parent_name else 'root'
@@ -450,7 +457,7 @@ class WanSelfAttention(nn.Module):
def qkv_fn_v(self, x): def qkv_fn_v(self, x):
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
return self.v(x).view(b, s, n, d) return self.v(x).view(b, s, n, d)
def qkv_fn_ip(self, x): def qkv_fn_ip(self, x):
b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim b, s, n, d = *x.shape[:2], self.num_heads, self.head_dim
q = self.norm_q(self.q(x) + self.q_loras(x).to(self.norm_q.weight.dtype)).to(x.dtype).view(b, s, n, d) q = self.norm_q(self.q(x) + self.q_loras(x).to(self.norm_q.weight.dtype)).to(x.dtype).view(b, s, n, d)
@@ -458,7 +465,7 @@ class WanSelfAttention(nn.Module):
v = (self.v(x) + self.v_loras(x)).view(b, s, n, d) v = (self.v(x) + self.v_loras(x)).view(b, s, n, d)
return q, k, v return q, k, v
def forward(self, q, k, v, seq_lens, lynx_ref_feature=None, lynx_ref_scale=1.0, attention_mode_override=None): def forward(self, q, k, v, seq_lens, lynx_ref_feature=None, lynx_ref_scale=1.0, attention_mode_override=None, onetoall_ref=None, onetoall_ref_scale=1.0):
r""" r"""
Args: Args:
x(Tensor): Shape [B, L, num_heads, C / num_heads] x(Tensor): Shape [B, L, num_heads, C / num_heads]
@@ -478,9 +485,12 @@ class WanSelfAttention(nn.Module):
if self.ref_adapter is not None and lynx_ref_feature is not None: if self.ref_adapter is not None and lynx_ref_feature is not None:
x = x.add(ref_x, alpha=lynx_ref_scale) x = x.add(ref_x, alpha=lynx_ref_scale)
if onetoall_ref is not None:
x = x.add(onetoall_ref, alpha=onetoall_ref_scale)
# output # output
return self.o(x.flatten(2)) return self.o(x.flatten(2))
def forward_ip(self, q, k, v, q_ip, k_ip, v_ip, seq_lens, attention_mode_override=None): def forward_ip(self, q, k, v, q_ip, k_ip, v_ip, seq_lens, attention_mode_override=None):
attention_mode = self.attention_mode attention_mode = self.attention_mode
if attention_mode_override is not None: if attention_mode_override is not None:
@@ -883,6 +893,8 @@ class WanAttentionBlock(nn.Module):
self.kv_cache = None self.kv_cache = None
self.use_motion_attn = use_motion_attn self.use_motion_attn = use_motion_attn
self.has_face_fuser_block = face_fuser_block self.has_face_fuser_block = face_fuser_block
self.ref_attn_k_img = None
self.ref_attn_v_img = None
# layers # layers
self.norm1 = WanLayerNorm(self.dim, eps) self.norm1 = WanLayerNorm(self.dim, eps)
@@ -965,7 +977,7 @@ class WanAttentionBlock(nn.Module):
return norm_x return norm_x
else: else:
return torch.addcmul(shift_msa, norm_x, 1 + scale_msa) return torch.addcmul(shift_msa, norm_x, 1 + scale_msa)
def ffn_chunked(self, mod_x, num_chunks=4): def ffn_chunked(self, mod_x, num_chunks=4):
seq_len = mod_x.shape[1] seq_len = mod_x.shape[1]
if seq_len <= 8192 or num_chunks <= 1: if seq_len <= 8192 or num_chunks <= 1:
@@ -996,6 +1008,8 @@ class WanAttentionBlock(nn.Module):
lynx_x_ip=None, lynx_ref_feature=None, lynx_ip_scale=1.0, lynx_ref_scale=1.0, #lynx lynx_x_ip=None, lynx_ref_feature=None, lynx_ip_scale=1.0, lynx_ref_scale=1.0, #lynx
x_ovi=None, e_ovi=None, freqs_ovi=None, context_ovi=None, seq_lens_ovi=None, grid_sizes_ovi=None, x_ovi=None, e_ovi=None, freqs_ovi=None, context_ovi=None, seq_lens_ovi=None, grid_sizes_ovi=None,
num_cond_latents=None, #longcat image cond amount num_cond_latents=None, #longcat image cond amount
x_onetoall_ref=None, onetoall_freqs=None, onetoall_ref=None, onetoall_ref_scale=1.0, #one-to-all
e_tr=None, tr_num=0, tr_start=0, #token replacement
): ):
r""" r"""
Args: Args:
@@ -1012,6 +1026,13 @@ class WanAttentionBlock(nn.Module):
self.seg_idx = [0, self.seg_idx, x.size(1)] self.seg_idx = [0, self.seg_idx, x.size(1)]
e = e[0] e = e[0]
use_token_replace = False
if e_tr is not None and tr_num > 0:
tr_shift_msa, tr_scale_msa, tr_gate_msa, tr_shift_mlp, tr_scale_mlp, tr_gate_mlp = self.get_mod(e_tr.to(x.device), self.modulation)
use_token_replace = True
tr_start = tr_start or 0
tr_end = tr_start + (tr_num or 0)
shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.get_mod(e.to(x.device), self.modulation) shift_msa, scale_msa, gate_msa, shift_mlp, scale_mlp, gate_mlp = self.get_mod(e.to(x.device), self.modulation)
del e del e
input_dtype = x.dtype input_dtype = x.dtype
@@ -1020,6 +1041,13 @@ class WanAttentionBlock(nn.Module):
is_longcat = C == 4096 is_longcat = C == 4096
if is_longcat: if is_longcat:
input_x = self.modulate(self.norm1(x.view(B, T, -1, C).to(shift_msa.dtype)), shift_msa, scale_msa, seg_idx=self.seg_idx).to(input_dtype).view(B, N, C) input_x = self.modulate(self.norm1(x.view(B, T, -1, C).to(shift_msa.dtype)), shift_msa, scale_msa, seg_idx=self.seg_idx).to(input_dtype).view(B, N, C)
elif use_token_replace:
norm_x = self.norm1(x.to(shift_msa.dtype))
input_x = torch.cat([
torch.addcmul(shift_msa, norm_x[:, :tr_start], 1 + scale_msa), # before replace → T
torch.addcmul(tr_shift_msa, norm_x[:, tr_start:tr_end], 1 + tr_scale_msa), # replace segment → t=0
torch.addcmul(shift_msa, norm_x[:, tr_end:], 1 + scale_msa) # after replace → T
], dim=1).to(input_dtype)
else: else:
input_x = self.modulate(self.norm1(x.to(shift_msa.dtype)), shift_msa, scale_msa, seg_idx=self.seg_idx).to(input_dtype) input_x = self.modulate(self.norm1(x.to(shift_msa.dtype)), shift_msa, scale_msa, seg_idx=self.seg_idx).to(input_dtype)
@@ -1053,6 +1081,23 @@ class WanAttentionBlock(nn.Module):
if lynx_ref_feature is None and self.self_attn.ref_adapter is not None: if lynx_ref_feature is None and self.self_attn.ref_adapter is not None:
lynx_ref_feature = input_x lynx_ref_feature = input_x
onetoall_ref = None
if x_onetoall_ref is not None:
b, s, n, d = *x_onetoall_ref.shape[:2], self.self_attn.num_heads, self.self_attn.head_dim
h_dim = w_dim = 2 * (self.head_dim // 6)
t_dim = self.head_dim - h_dim - w_dim
q_ref = self.self_attn.norm_q(self.self_attn.q(input_x)).to(input_x.dtype).view(b, N, n, d)
q_ref = apply_rotary_emb_split(q_ref, freqs, t_dim) # Apply split rotary embedding (only to H/W dimensions, leaving T unchanged)
k_ref = self.ref_attn_norm_k_img(self.ref_attn_k_img(x_onetoall_ref).to(self.ref_attn_norm_k_img.weight.dtype)).to(x_onetoall_ref.dtype).view(b, s, n, d)
k_ref = apply_rotary_emb_split(k_ref, onetoall_freqs, t_dim)
v_ref = self.ref_attn_v_img(x_onetoall_ref).view(b, s, n, d)
onetoall_ref = attention(q_ref, k_ref, v_ref, k_lens=seq_lens, attention_mode=self.attention_mode)
del q_ref, k_ref, v_ref
#RoPE and QKV computation #RoPE and QKV computation
if inner_t is not None: if inner_t is not None:
#query, key, value #query, key, value
@@ -1104,10 +1149,10 @@ class WanAttentionBlock(nn.Module):
# FETA # FETA
if enhance_enabled: if enhance_enabled:
feta_scores = get_feta_scores(q, k) feta_scores = get_feta_scores(q, k)
#self-attention #self-attention
split_attn = (context is not None split_attn = (context is not None
and (context.shape[0] > 1 or (clip_embed is not None and clip_embed.shape[0] > 1)) and (context.shape[0] > 1 or (clip_embed is not None and clip_embed.shape[0] > 1))
and x.shape[0] == 1 and x.shape[0] == 1
and inner_t is None and inner_t is None
and x_ip is None # Don't split when using IP-Adapter and x_ip is None # Don't split when using IP-Adapter
@@ -1144,18 +1189,18 @@ class WanAttentionBlock(nn.Module):
num_cond_latents_thw = num_cond_latents * (N // num_latent_frames) num_cond_latents_thw = num_cond_latents * (N // num_latent_frames)
# process the condition tokens # process the condition tokens
x_cond = self.self_attn.forward( x_cond = self.self_attn.forward(
q[:, :num_cond_latents_thw].contiguous(), q[:, :num_cond_latents_thw].contiguous(),
k[:, :num_cond_latents_thw].contiguous(), k[:, :num_cond_latents_thw].contiguous(),
v[:, :num_cond_latents_thw].contiguous(), v[:, :num_cond_latents_thw].contiguous(),
seq_lens) seq_lens)
# process the noise tokens # process the noise tokens
x_noise = self.self_attn.forward(q[:, num_cond_latents_thw:].contiguous(), k, v, seq_lens) x_noise = self.self_attn.forward(q[:, num_cond_latents_thw:].contiguous(), k, v, seq_lens)
# merge x_cond and x_noise # merge x_cond and x_noise
y = torch.cat([x_cond, x_noise], dim=1).contiguous() y = torch.cat([x_cond, x_noise], dim=1).contiguous()
else: else:
y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale) y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale, onetoall_ref=onetoall_ref, onetoall_ref_scale=onetoall_ref_scale)
del q, k, v del q, k, v,
# FETA # FETA
if enhance_enabled: if enhance_enabled:
@@ -1173,17 +1218,23 @@ class WanAttentionBlock(nn.Module):
) )
# S2V # S2V
if zero_timestep: if zero_timestep:
z = [] z = []
for i in range(2): for i in range(2):
z.append(y[:, self.seg_idx[i]:self.seg_idx[i + 1]] * gate_msa[:, i:i + 1]) z.append(y[:, self.seg_idx[i]:self.seg_idx[i + 1]] * gate_msa[:, i:i + 1])
y = torch.cat(z, dim=1) y = torch.cat(z, dim=1)
x = x.add(y) x = x.add(y)
else: else:
if not is_longcat: if is_longcat:
x = x.addcmul(y, gate_msa)
else:
x = x + (y.view(B, -1, N//T, C).float() * gate_msa).to(input_dtype).view(B, -1, C) x = x + (y.view(B, -1, N//T, C).float() * gate_msa).to(input_dtype).view(B, -1, C)
elif use_token_replace:
x = x + torch.cat([
y[:, :tr_start] * gate_msa,
y[:, tr_start:tr_end] * tr_gate_msa,
y[:, tr_end:] * gate_msa
], dim=1).to(input_dtype)
else:
x = x.addcmul(y, gate_msa)
del y, gate_msa del y, gate_msa
# cross-attention & ffn function # cross-attention & ffn function
@@ -1225,7 +1276,7 @@ class WanAttentionBlock(nn.Module):
x = x.add(x_audio, alpha=audio_scale) x = x.add(x_audio, alpha=audio_scale)
# MTV-Crafter Motion Attention # MTV-Crafter Motion Attention
if self.use_motion_attn and mtv_motion_tokens is not None and mtv_motion_rotary_emb is not None: if self.use_motion_attn and mtv_motion_tokens is not None and mtv_motion_rotary_emb is not None:
x_motion = self.motion_attn(self.norm4(x), mtv_motion_tokens, mtv_motion_rotary_emb, grid_sizes, mtv_freqs) x_motion = self.motion_attn(self.norm4(x), mtv_motion_tokens, mtv_motion_rotary_emb, grid_sizes, mtv_freqs)
x = x.add(x_motion, alpha=mtv_strength) x = x.add(x_motion, alpha=mtv_strength)
@@ -1248,14 +1299,22 @@ class WanAttentionBlock(nn.Module):
norm2_x = torch.cat(parts, dim=1) norm2_x = torch.cat(parts, dim=1)
x_ffn = self.ffn(norm2_x) x_ffn = self.ffn(norm2_x)
else: else:
if not is_longcat: if is_longcat:
mod_x = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp)
else:
mod_x = torch.addcmul(shift_mlp, self.norm2(x.view(B, -1, N//T, C).float()), 1 + scale_mlp).view(B, -1, C) mod_x = torch.addcmul(shift_mlp, self.norm2(x.view(B, -1, N//T, C).float()), 1 + scale_mlp).view(B, -1, C)
elif use_token_replace:
norm2_x = self.norm2(x.to(shift_mlp.dtype))
mod_x = torch.cat([
torch.addcmul(shift_mlp, norm2_x[:, :tr_start], 1 + scale_mlp),
torch.addcmul(tr_shift_mlp, norm2_x[:, tr_start:tr_end], 1 + tr_scale_mlp),
torch.addcmul(shift_mlp, norm2_x[:, tr_end:], 1 + scale_mlp)
], dim=1)
else:
mod_x = torch.addcmul(shift_mlp, self.norm2(x.to(shift_mlp.dtype)), 1 + scale_mlp)
del shift_mlp, scale_mlp del shift_mlp, scale_mlp
x_ffn = self.ffn_chunked(mod_x.to(input_dtype), num_chunks=1) x_ffn = self.ffn_chunked(mod_x.to(input_dtype), num_chunks=1)
del mod_x del mod_x
# gate_mlp # gate_mlp
if zero_timestep: if zero_timestep:
z = [] z = []
@@ -1264,10 +1323,16 @@ class WanAttentionBlock(nn.Module):
x_ffn = torch.cat(z, dim=1) x_ffn = torch.cat(z, dim=1)
x = x.add(x_ffn) x = x.add(x_ffn)
else: else:
if not is_longcat: if is_longcat:
x = x.addcmul(x_ffn.to(gate_mlp.dtype), gate_mlp).to(input_dtype)
else:
x = x + (gate_mlp * x_ffn.view(B, -1, N//T, C).float()).to(input_dtype).view(B, -1, C) x = x + (gate_mlp * x_ffn.view(B, -1, N//T, C).float()).to(input_dtype).view(B, -1, C)
elif use_token_replace:
x = x + torch.cat([
x_ffn[:, :tr_start] * gate_mlp,
x_ffn[:, tr_start:tr_end] * tr_gate_mlp,
x_ffn[:, tr_end:] * gate_mlp
], dim=1).to(input_dtype)
else:
x = x.addcmul(x_ffn.to(gate_mlp.dtype), gate_mlp).to(input_dtype)
del gate_mlp del gate_mlp
if x_ip is not None: #stand-in if x_ip is not None: #stand-in
@@ -1389,7 +1454,7 @@ class BaseWanAttentionBlock(WanAttentionBlock):
x, x_ip, lynx_ref_feature, x_ovi = super().forward(x, **kwargs) x, x_ip, lynx_ref_feature, x_ovi = super().forward(x, **kwargs)
if vace_hints is None: if vace_hints is None:
return x, x_ip, lynx_ref_feature, x_ovi return x, x_ip, lynx_ref_feature, x_ovi
if self.block_id is not None: if self.block_id is not None:
for i in range(len(vace_hints)): for i in range(len(vace_hints)):
x.add_(vace_hints[i][self.block_id].to(x.device), alpha=vace_context_scale[i]) x.add_(vace_hints[i][self.block_id].to(x.device), alpha=vace_context_scale[i])
@@ -1419,17 +1484,26 @@ class Head(nn.Module):
e = (self.modulation.unsqueeze(2) + e.unsqueeze(1)).chunk(2, dim=1) e = (self.modulation.unsqueeze(2) + e.unsqueeze(1)).chunk(2, dim=1)
return [ei.squeeze(1) for ei in e] return [ei.squeeze(1) for ei in e]
def forward(self, x, e, **kwargs): def forward(self, x, e, e_tr=None, tr_start=0, tr_num=0, **kwargs):
r""" r"""
Args: Args:
x(Tensor): Shape [B, L1, C] x(Tensor): Shape [B, L1, C]
e(Tensor): Shape [B, C] e(Tensor): Shape [B, C]
""" """
e = self.get_mod(e.to(x.device)) e = self.get_mod(e.to(x.device))
x = self.head(self.norm(x.float()).to(x.dtype).mul_(1 + e[1]).add_(e[0])) if tr_num > 0 and e_tr is not None:
e_tr = self.get_mod(e_tr.to(x.device))
tr_end = tr_start + tr_num
norm_x = self.norm(x.float()).to(x.dtype)
x = self.head(torch.cat([
norm_x[:, :tr_start].mul(1 + e[1]).add(e[0]),
norm_x[:, tr_start:tr_end].mul(1 + e_tr[1]).add(e_tr[0]),
norm_x[:, tr_end:].mul(1 + e[1]).add(e[0])
], dim=1))
else:
x = self.head(self.norm(x.float()).to(x.dtype).mul_(1 + e[1]).add_(e[0]))
return x return x
class Head_adaLN(nn.Module): class Head_adaLN(nn.Module):
def __init__(self, dim, out_dim, patch_size, eps=1e-6, adaln_tembed_dim=512): def __init__(self, dim, out_dim, patch_size, eps=1e-6, adaln_tembed_dim=512):
@@ -1459,7 +1533,7 @@ class Head_adaLN(nn.Module):
self.modulation.to(torch.float32) self.modulation.to(torch.float32)
shift, scale = self.modulation(e).unsqueeze(2).chunk(2, dim=-1) # [B, T, 1, C] shift, scale = self.modulation(e).unsqueeze(2).chunk(2, dim=-1) # [B, T, 1, C]
return self.head(self.norm(x.view(B, T, -1, C).float()).mul_(1 + scale).add_(shift).view(B, N, C).to(x.dtype)) return self.head(self.norm(x.view(B, T, -1, C).float()).mul_(1 + scale).add_(shift).view(B, N, C).to(x.dtype))
class MLPProj(torch.nn.Module): class MLPProj(torch.nn.Module):
@@ -1732,7 +1806,7 @@ class WanModel(torch.nn.Module):
nn.SiLU(), nn.SiLU(),
ConvMLP(dim, dim * 4, kernel_size=7, padding=3), ConvMLP(dim, dim * 4, kernel_size=7, padding=3),
) )
self.original_patch_embedding = self.patch_embedding self.original_patch_embedding = self.patch_embedding
self.expanded_patch_embedding = self.patch_embedding self.expanded_patch_embedding = self.patch_embedding
@@ -1803,18 +1877,12 @@ class WanModel(torch.nn.Module):
self.head = Head_adaLN(dim, out_dim, patch_size, eps, adaln_tembed_dim=512) self.head = Head_adaLN(dim, out_dim, patch_size, eps, adaln_tembed_dim=512)
d = self.dim // self.num_heads d = self.dim // self.num_heads
self.rope_embedder = EmbedND_RifleX( self.rope_embedder = EmbedND_RifleX(d, 10000.0, [d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)], num_frames=None, k=None)
d,
10000.0,
[d - 4 * (d // 6), 2 * (d // 6), 2 * (d // 6)],
num_frames=None,
k=None,
)
self.cached_freqs = self.cached_shape = self.cached_cond = None self.cached_freqs = self.cached_shape = self.cached_cond = None
# buffers (don't use register_buffer otherwise dtype will be changed in to()) # buffers (don't use register_buffer otherwise dtype will be changed in to())
assert (dim % num_heads) == 0 and (dim // num_heads) % 2 == 0 assert (dim % num_heads) == 0 and (dim // num_heads) % 2 == 0
if model_type == 'i2v' or model_type == 'fl2v': if model_type == 'i2v' or model_type == 'fl2v':
self.img_emb = MLPProj(1280, dim, fl_pos_emb=model_type == 'fl2v') self.img_emb = MLPProj(1280, dim, fl_pos_emb=model_type == 'fl2v')
@@ -2149,7 +2217,8 @@ class WanModel(torch.nn.Module):
flashvsr_LQ_latent=None, flashvsr_strength=1.0, flashvsr_LQ_latent=None, flashvsr_strength=1.0,
num_cond_latents=None, num_cond_latents=None,
add_text_emb=None, add_text_emb=None,
sdancer_input=None # SteadyDancer sdancer_input=None, # SteadyDancer
one_to_all_input=None, # One-to-All
): ):
r""" r"""
Forward pass through the diffusion model Forward pass through the diffusion model
@@ -2173,9 +2242,9 @@ class WanModel(torch.nn.Module):
List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8] List of denoised video tensors with original input shapes [C_out, F, H / 8, W / 8]
""" """
# Stand-In only used on first positive pass, then cached in kv_cache # Stand-In only used on first positive pass, then cached in kv_cache
if is_uncond or current_step > 0: if is_uncond or current_step > 0:
standin_input = None standin_input = None
# MTV Crafter motion projection # MTV Crafter motion projection
if mtv_motion_tokens is not None: if mtv_motion_tokens is not None:
bs, motion_seq_len = mtv_motion_tokens.shape[0], mtv_motion_tokens.shape[1] bs, motion_seq_len = mtv_motion_tokens.shape[0], mtv_motion_tokens.shape[1]
@@ -2213,14 +2282,14 @@ class WanModel(torch.nn.Module):
lynx_ip_scale = lynx_embeds.get("ip_scale", 1.0) lynx_ip_scale = lynx_embeds.get("ip_scale", 1.0)
lynx_ref_scale = lynx_embeds.get("ref_scale", 1.0) lynx_ref_scale = lynx_embeds.get("ref_scale", 1.0)
#s2v #s2v
if self.model_type == 's2v' and s2v_audio_input is not None: if self.model_type == 's2v' and s2v_audio_input is not None:
if is_uncond: if is_uncond:
s2v_audio_input = s2v_audio_input * 0 # to match original code s2v_audio_input = s2v_audio_input * 0 # to match original code
s2v_audio_input = torch.cat([s2v_audio_input[..., 0:1].repeat(1, 1, 1, s2v_motion_frames[0]), s2v_audio_input], dim=-1) s2v_audio_input = torch.cat([s2v_audio_input[..., 0:1].repeat(1, 1, 1, s2v_motion_frames[0]), s2v_audio_input], dim=-1)
audio_emb_res = self.casual_audio_encoder(s2v_audio_input) audio_emb_res = self.casual_audio_encoder(s2v_audio_input)
if self.enable_adain: if self.enable_adain:
audio_emb_global, audio_emb = audio_emb_res audio_emb_global, audio_emb = audio_emb_res
@@ -2241,7 +2310,7 @@ class WanModel(torch.nn.Module):
if sdancer_input is not None and sdancer_input['start_percent'] <= current_step_percentage <= sdancer_input['end_percent']: if sdancer_input is not None and sdancer_input['start_percent'] <= current_step_percentage <= sdancer_input['end_percent']:
sdancer_enabled = True sdancer_enabled = True
x_noise_clone = torch.stack(x) x_noise_clone = torch.stack(x)
# I2V # I2V
if y is not None: if y is not None:
if hasattr(self, "randomref_embedding_pose") and unianim_data is not None: if hasattr(self, "randomref_embedding_pose") and unianim_data is not None:
@@ -2251,6 +2320,40 @@ class WanModel(torch.nn.Module):
y[0].add_(random_ref_emb, alpha=unianim_data["strength"]) y[0].add_(random_ref_emb, alpha=unianim_data["strength"])
x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)] x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)]
suffix_frames = x[0].shape[1]
prefix_frames = 0
# One-to-all-Animation
onetoall_ref_block_samples = onetoall_freqs = prev_x = prev_control = None
onetoall_ref_scale = 1.0
onetoall_control_enabled = use_token_replace = False
e0_token_replace = token_replace_start = None
replace_token_num = token_replace_start = 0
if one_to_all_input is not None:
# reference condition
ref_cond_latent = one_to_all_input.get("ref_latent_pos", None) if not is_uncond else one_to_all_input.get("ref_latent_neg", None)
if ref_cond_latent is not None and one_to_all_input['ref_start_percent'] <= current_step_percentage <= one_to_all_input['ref_end_percent']:
onetoall_ref_scale = one_to_all_input.get("ref_strength", 1.0)
image_cond = self.image_to_cond(ref_cond_latent.to(self.main_device, self.base_dtype))[0]
x = [torch.cat([v, u], dim=1) for v, u in zip([image_cond], x)]
seq_len += math.ceil((image_cond.shape[-1] * image_cond.shape[-2]) / 4 * image_cond.shape[-3])
F += 1
prefix_frames = 1
suffix_frames += 1
onetoall_ref_block_samples, onetoall_freqs = self.refextractor(ref_cond_latent, timestep=t)
# pose controlnet
controlnet_tokens = one_to_all_input.get("controlnet_tokens", None)
if not is_uncond and controlnet_tokens is not None and one_to_all_input['controlnet_start_percent'] <= current_step_percentage <= one_to_all_input['controlnet_end_percent']:
onetoall_control_enabled = True
onetoall_control_strength = one_to_all_input.get("controlnet_strength", 1.0)
# token replace
if one_to_all_input.get("token_replace", False):
use_token_replace = True
num_latent_frames_to_replace = one_to_all_input.get("num_latent_frames_to_replace", 2)
t_token_replace = torch.zeros_like(t)
token_replace_start = (H // self.patch_size[1]) * (W // self.patch_size[2]) # skip first (ref) frame
replace_token_num = num_latent_frames_to_replace * token_replace_start # zero next frames
#uni3c controlnet #uni3c controlnet
if uni3c_data is not None: if uni3c_data is not None:
render_latent = uni3c_data["render_latent"].to(self.base_dtype) render_latent = uni3c_data["render_latent"].to(self.base_dtype)
@@ -2278,15 +2381,13 @@ class WanModel(torch.nn.Module):
else: else:
self.original_patch_embedding.to(self.main_device) self.original_patch_embedding.to(self.main_device)
x = [self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in x] x = [self.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in x]
orig_frames = x[0].shape[1]
# ovi audio model # ovi audio model
if self.audio_model is not None: if self.audio_model is not None:
x_ovi = [self.audio_model.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x_ovi[0].dtype) for u in x_ovi] x_ovi = [self.audio_model.original_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x_ovi[0].dtype) for u in x_ovi]
grid_sizes_ovi = torch.stack([torch.tensor(u.shape[1:2], dtype=torch.long) for u in x_ovi]) grid_sizes_ovi = torch.stack([torch.tensor(u.shape[1:2], dtype=torch.long) for u in x_ovi])
seq_lens_ovi = torch.tensor([u.size(1) for u in x_ovi], dtype=torch.int32) seq_lens_ovi = torch.tensor([u.size(1) for u in x_ovi], dtype=torch.int32)
x_ovi = torch.cat([torch.cat([u, u.new_zeros(1, seq_len_ovi - u.size(1), u.size(2))], dim=1) for u in x_ovi]) x_ovi = torch.cat([torch.cat([u, u.new_zeros(1, seq_len_ovi - u.size(1), u.size(2))], dim=1) for u in x_ovi])
d = self.dim // self.num_heads d = self.dim // self.num_heads
freqs_ovi = rope_params(1024, d - 4 * (d // 6), freqs_scaling=0.19676).to(self.main_device) freqs_ovi = rope_params(1024, d - 4 * (d // 6), freqs_scaling=0.19676).to(self.main_device)
x_ovi = x_ovi.to(self.main_device, self.base_dtype) x_ovi = x_ovi.to(self.main_device, self.base_dtype)
@@ -2375,7 +2476,7 @@ class WanModel(torch.nn.Module):
seq_len += end_ref_latent_seq_len seq_len += end_ref_latent_seq_len
x = [torch.cat([u, end_ref_latent.unsqueeze(0)], dim=1) for end_ref_latent, u in zip(end_ref_latent, x)] x = [torch.cat([u, end_ref_latent.unsqueeze(0)], dim=1) for end_ref_latent, u in zip(end_ref_latent, x)]
x = torch.cat([torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))], dim=1) for u in x]) x = torch.cat([torch.cat([u, u.new_zeros(1, seq_len - u.size(1), u.size(2))], dim=1) for u in x])
if self.trainable_cond_mask is not None: if self.trainable_cond_mask is not None:
@@ -2397,11 +2498,11 @@ class WanModel(torch.nn.Module):
if freqs is None and "comfy" in self.rope_func: #comfy rope if freqs is None and "comfy" in self.rope_func: #comfy rope
current_shape = (F, H, W) current_shape = (F, H, W)
has_cond = attn_cond is not None has_cond = attn_cond is not None
if (self.cached_freqs is not None and if (self.cached_freqs is not None and
self.cached_shape == current_shape and self.cached_shape == current_shape and
self.cached_cond == has_cond and self.cached_cond == has_cond and
self.cached_rope_k == self.rope_embedder.k and self.cached_rope_k == self.rope_embedder.k and
self.cached_ntk_alphas == ntk_alphas self.cached_ntk_alphas == ntk_alphas
@@ -2411,9 +2512,9 @@ class WanModel(torch.nn.Module):
freqs = self.rope_encode_comfy(F, H, W, freq_offset=freq_offset, ntk_alphas=ntk_alphas, attn_cond_shape=attn_cond_shape, device=x.device, dtype=x.dtype) freqs = self.rope_encode_comfy(F, H, W, freq_offset=freq_offset, ntk_alphas=ntk_alphas, attn_cond_shape=attn_cond_shape, device=x.device, dtype=x.dtype)
if s2v_ref_latent is not None: if s2v_ref_latent is not None:
freqs_ref = self.rope_encode_comfy( freqs_ref = self.rope_encode_comfy(
s2v_ref_latent.shape[2], s2v_ref_latent.shape[2],
s2v_ref_latent.shape[3], s2v_ref_latent.shape[3],
s2v_ref_latent.shape[4], s2v_ref_latent.shape[4],
t_start=max(30, F + 9), device=x.device, dtype=x.dtype) t_start=max(30, F + 9), device=x.device, dtype=x.dtype)
freqs = torch.cat([freqs, freqs_ref], dim=1) freqs = torch.cat([freqs, freqs_ref], dim=1)
@@ -2463,6 +2564,9 @@ class WanModel(torch.nn.Module):
time_embed_dtype = self.base_dtype time_embed_dtype = self.base_dtype
e = self.time_embedding(sinusoidal_embedding_1d(self.freq_dim, t.flatten()).to(time_embed_dtype)) # b, dim e = self.time_embedding(sinusoidal_embedding_1d(self.freq_dim, t.flatten()).to(time_embed_dtype)) # b, dim
e0 = self.time_projection(e).unflatten(1, (6, self.dim)) # b, 6, dim e0 = self.time_projection(e).unflatten(1, (6, self.dim)) # b, 6, dim
if use_token_replace:
e_token_replace = self.time_embedding(sinusoidal_embedding_1d(self.freq_dim, t_token_replace.flatten()).to(time_embed_dtype)) # b, dim
e0_token_replace = self.time_projection(e_token_replace).unflatten(1, (6, self.dim)) # b, 6, dim
else: else:
time_embed_dtype = self.time_embedding.mlp[0].weight.dtype time_embed_dtype = self.time_embedding.mlp[0].weight.dtype
if time_embed_dtype not in [torch.float16, torch.bfloat16, torch.float32]: if time_embed_dtype not in [torch.float16, torch.bfloat16, torch.float32]:
@@ -2480,7 +2584,7 @@ class WanModel(torch.nn.Module):
last_timestep = t[:, -1:] last_timestep = t[:, -1:]
padding = last_timestep.expand(t.size(0), seq_len_ovi - t.size(1)) padding = last_timestep.expand(t.size(0), seq_len_ovi - t.size(1))
t_ovi = torch.cat([t, padding], dim=1) t_ovi = torch.cat([t, padding], dim=1)
e_ovi = self.audio_model.time_embedding(sinusoidal_embedding_1d(self.audio_model.freq_dim, t_ovi.flatten()).to(time_embed_dtype)).unsqueeze(0) # b, dim e_ovi = self.audio_model.time_embedding(sinusoidal_embedding_1d(self.audio_model.freq_dim, t_ovi.flatten()).to(time_embed_dtype)).unsqueeze(0) # b, dim
e0_ovi = self.audio_model.time_projection(e_ovi).unflatten(2, (6, self.dim)).movedim(1, 2) # B, seq_len, 6, dim e0_ovi = self.audio_model.time_projection(e_ovi).unflatten(2, (6, self.dim)).movedim(1, 2) # B, seq_len, 6, dim
else: else:
@@ -2516,14 +2620,14 @@ class WanModel(torch.nn.Module):
if expanded_timesteps: if expanded_timesteps:
e = e.view(b, f, 1, 1, self.dim).expand(b, f, grid_sizes[0][1], grid_sizes[0][2], self.dim) e = e.view(b, f, 1, 1, self.dim).expand(b, f, grid_sizes[0][1], grid_sizes[0][2], self.dim)
e0 = e0.view(b, f, 1, 1, 6, self.dim).expand(b, f, grid_sizes[0][1], grid_sizes[0][2], 6, self.dim) e0 = e0.view(b, f, 1, 1, 6, self.dim).expand(b, f, grid_sizes[0][1], grid_sizes[0][2], 6, self.dim)
e = e.flatten(1, 3) e = e.flatten(1, 3)
e0 = e0.flatten(1, 3) e0 = e0.flatten(1, 3)
e0 = e0.transpose(1, 2) e0 = e0.transpose(1, 2)
if not e0.is_contiguous(): if not e0.is_contiguous():
e0 = e0.contiguous() e0 = e0.contiguous()
e = e.to(self.offload_device, non_blocking=self.use_non_blocking) e = e.to(self.offload_device, non_blocking=self.use_non_blocking)
# clip vision embedding # clip vision embedding
@@ -2580,7 +2684,7 @@ class WanModel(torch.nn.Module):
[u, u.new_zeros(self.text_len - u.size(0), u.size(1))]) [u, u.new_zeros(self.text_len - u.size(0), u.size(1))])
for u in nag_context for u in nag_context
]).to(text_embed_dtype)) ]).to(text_embed_dtype))
if self.offload_txt_emb: if self.offload_txt_emb:
self.text_embedding.to(self.offload_device, non_blocking=self.use_non_blocking) self.text_embedding.to(self.offload_device, non_blocking=self.use_non_blocking)
@@ -2637,7 +2741,7 @@ class WanModel(torch.nn.Module):
accumulated_rel_l1_distance = torch.tensor(0.0, dtype=torch.float32, device=device) accumulated_rel_l1_distance = torch.tensor(0.0, dtype=torch.float32, device=device)
if pred_id is None: if pred_id is None:
pred_id = self.teacache_state.new_prediction(cache_device=self.cache_device) pred_id = self.teacache_state.new_prediction(cache_device=self.cache_device)
should_calc = True should_calc = True
else: else:
previous_modulated_input = self.teacache_state.get(pred_id)['previous_modulated_input'] previous_modulated_input = self.teacache_state.get(pred_id)['previous_modulated_input']
previous_modulated_input = previous_modulated_input.to(device) previous_modulated_input = previous_modulated_input.to(device)
@@ -2665,7 +2769,7 @@ class WanModel(torch.nn.Module):
accumulated_rel_l1_distance = accumulated_rel_l1_distance.to(self.cache_device) accumulated_rel_l1_distance = accumulated_rel_l1_distance.to(self.cache_device)
previous_modulated_input = e.to(self.cache_device).clone() if (self.teacache_use_coefficients and self.teacache_mode == 'e') else e0.to(self.cache_device).clone() previous_modulated_input = e.to(self.cache_device).clone() if (self.teacache_use_coefficients and self.teacache_mode == 'e') else e0.to(self.cache_device).clone()
if not should_calc: if not should_calc:
x = x.to(previous_residual.dtype) + previous_residual.to(x.device) x = x.to(previous_residual.dtype) + previous_residual.to(x.device)
self.teacache_state.update( self.teacache_state.update(
@@ -2803,7 +2907,11 @@ class WanModel(torch.nn.Module):
lynx_x_ip=lynx_x_ip, lynx_x_ip=lynx_x_ip,
lynx_ip_scale=lynx_ip_scale, lynx_ip_scale=lynx_ip_scale,
lynx_ref_scale=lynx_ref_scale, lynx_ref_scale=lynx_ref_scale,
num_cond_latents=num_cond_latents num_cond_latents=num_cond_latents,
onetoall_ref_scale=onetoall_ref_scale,
e_tr=e0_token_replace if use_token_replace else None,
tr_start=token_replace_start,
tr_num=replace_token_num,
) )
if self.audio_model is not None: if self.audio_model is not None:
kwargs['e_ovi'] = e0_ovi.to(self.base_dtype) kwargs['e_ovi'] = e0_ovi.to(self.base_dtype)
@@ -2812,7 +2920,7 @@ class WanModel(torch.nn.Module):
kwargs['seq_lens_ovi'] = seq_lens_ovi kwargs['seq_lens_ovi'] = seq_lens_ovi
kwargs['freqs_ovi'] = freqs_ovi kwargs['freqs_ovi'] = freqs_ovi
if vace_data is not None: if vace_data is not None:
vace_hint_list = [] vace_hint_list = []
vace_scale_list = [] vace_scale_list = []
@@ -2828,7 +2936,7 @@ class WanModel(torch.nn.Module):
vace_hints = self.forward_vace(x, vace_data, seq_len, kwargs) vace_hints = self.forward_vace(x, vace_data, seq_len, kwargs)
vace_hint_list.append(vace_hints) vace_hint_list.append(vace_hints)
vace_scale_list.append(1.0) vace_scale_list.append(1.0)
kwargs['vace_hints'] = vace_hint_list kwargs['vace_hints'] = vace_hint_list
kwargs['vace_context_scale'] = vace_scale_list kwargs['vace_context_scale'] = vace_scale_list
@@ -2898,7 +3006,17 @@ class WanModel(torch.nn.Module):
if b in self.slg_blocks and is_uncond: if b in self.slg_blocks and is_uncond:
if self.slg_start_percent <= current_step_percentage <= self.slg_end_percent: if self.slg_start_percent <= current_step_percentage <= self.slg_end_percent:
continue continue
x, x_ip, lynx_ref_feature, x_ovi = block(x, x_ip=x_ip, lynx_ref_feature=lynx_ref_feature, x_ovi=x_ovi, **kwargs) #run block
x_onetoall_ref = None
if onetoall_ref_block_samples is not None:
interval_ref = len(self.blocks) / len(onetoall_ref_block_samples)
interval_ref = int(np.ceil(interval_ref))
x_onetoall_ref = onetoall_ref_block_samples[b // interval_ref]
# ---run block----#
x, x_ip, lynx_ref_feature, x_ovi = block(x, x_ip=x_ip, lynx_ref_feature=lynx_ref_feature, x_ovi=x_ovi, x_onetoall_ref=x_onetoall_ref, onetoall_freqs=onetoall_freqs, **kwargs)
# ---post block----#
if self.audio_injector is not None and s2v_audio_input is not None: if self.audio_injector is not None and s2v_audio_input is not None:
x = self.audio_injector_forward(b, x, merged_audio_emb, scale=s2v_audio_scale) #s2v x = self.audio_injector_forward(b, x, merged_audio_emb, scale=s2v_audio_scale) #s2v
if block.has_face_fuser_block and motion_vec is not None: if block.has_face_fuser_block and motion_vec is not None:
@@ -2924,6 +3042,31 @@ class WanModel(torch.nn.Module):
#controlnet #controlnet
if (controlnet is not None) and (b % controlnet["controlnet_stride"] == 0) and (b // controlnet["controlnet_stride"] < len(controlnet["controlnet_states"])): if (controlnet is not None) and (b % controlnet["controlnet_stride"] == 0) and (b // controlnet["controlnet_stride"] < len(controlnet["controlnet_states"])):
x[:, :self.original_seq_len] += controlnet["controlnet_states"][b // controlnet["controlnet_stride"]].to(x) * controlnet["controlnet_weight"] x[:, :self.original_seq_len] += controlnet["controlnet_states"][b // controlnet["controlnet_stride"]].to(x) * controlnet["controlnet_weight"]
# One-to-All-Animation controlnet
if onetoall_control_enabled:
if prev_x is not None and (b - 1) < len(self.controlnet.blocks):
tqdm.write(f"Applying One-to-All ControlNet at block {b}")
if b == 1:
ctrl_in = prev_x + controlnet_tokens
elif prev_control is not None:
ctrl_in = prev_control
self.controlnet.blocks[b - 1].to(self.main_device)
control_out = self.controlnet.blocks[b - 1](ctrl_in, e0, seq_lens, freqs, e_tr=e0_token_replace, tr_num=replace_token_num,tr_start=token_replace_start, split_rope=False)
self.controlnet.blocks[b - 1].to(self.offload_device, non_blocking=self.use_non_blocking)
prev_control = control_out
control_out_proj = self.controlnet_zero[b - 1](control_out)
x = x + control_out_proj * onetoall_control_strength
if b < len(self.controlnet.blocks): # Store prev_x only while controlnet is active
prev_x = x
elif b == len(self.controlnet.blocks): # Controlnet done, free memory
prev_x = None
prev_control = None
if controlnet_tokens is not None:
del controlnet_tokens
controlnet_tokens = None
mm.soft_empty_cache()
if lynx_ref_feature_extractor: if lynx_ref_feature_extractor:
return lynx_ref_buffer return lynx_ref_buffer
@@ -2953,20 +3096,20 @@ class WanModel(torch.nn.Module):
accumulated_error = 0.0, accumulated_error = 0.0,
cache_ovi = x_ovi.clone().to(original_x.device) - original_x_ovi if x_ovi is not None else None cache_ovi = x_ovi.clone().to(original_x.device) - original_x_ovi if x_ovi is not None else None
) )
if self.enable_easycache and (self.easycache_start_step <= current_step <= self.easycache_end_step) and pred_id is not None: if self.enable_easycache and (self.easycache_start_step <= current_step <= self.easycache_end_step) and pred_id is not None:
self.easycache_state.update( self.easycache_state.update(
pred_id, pred_id,
previous_raw_output=x.clone(), previous_raw_output=x.clone(),
) )
if self.ref_conv is not None and fun_ref is not None: if self.ref_conv is not None and fun_ref is not None:
fun_ref_length = fun_ref.size(1) fun_ref_length = fun_ref.size(1)
x = x[:, fun_ref_length:] x = x[:, fun_ref_length:]
grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device) grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
if end_ref_latent is not None: if end_ref_latent is not None:
end_ref_latent_length = end_ref_latent.size(1) end_ref_latent_length = end_ref_latent.size(1)
x = x[:, :-end_ref_latent_length] x = x[:, :-end_ref_latent_length]
@@ -2976,10 +3119,11 @@ class WanModel(torch.nn.Module):
x = x[:, :self.original_seq_len] x = x[:, :self.original_seq_len]
grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device) grid_sizes = torch.stack([torch.tensor([u[0] - 1, u[1], u[2]]) for u in grid_sizes]).to(grid_sizes.device)
x = x[:, :self.original_seq_len] x = x[:, :self.original_seq_len]
x = self.head(x, e.to(x.device), temp_length=F) x = self.head(x, e.to(x.device), temp_length=F,
e_tr=e_token_replace.to(x.device) if use_token_replace else None, tr_start=token_replace_start, tr_num=replace_token_num)
if x_ovi is not None: if x_ovi is not None:
x_ovi = self.audio_model.head(x_ovi, e_ovi.to(x_ovi.device)) x_ovi = self.audio_model.head(x_ovi, e_ovi.to(x_ovi.device))
@@ -2987,9 +3131,9 @@ class WanModel(torch.nn.Module):
assert len(x) == len(grid_sizes_ovi) assert len(x) == len(grid_sizes_ovi)
x_ovi = [u[:gs] for u, gs in zip(x_ovi, grid_sizes_ovi)] x_ovi = [u[:gs] for u, gs in zip(x_ovi, grid_sizes_ovi)]
x_ovi = [u.float() for u in x_ovi] x_ovi = [u.float() for u in x_ovi]
x = self.unpatchify(x, original_grid_sizes) # type: ignore[arg-type] x = self.unpatchify(x, original_grid_sizes)
x = [u[:, :orig_frames, ...].float() for u in x] x = [u[:, prefix_frames:suffix_frames, ...].float() for u in x]
return (x, x_ovi, pred_id) if pred_id is not None else (x, x_ovi, None) return (x, x_ovi, pred_id) if pred_id is not None else (x, x_ovi, None)
def unpatchify(self, x, grid_sizes): def unpatchify(self, x, grid_sizes):