From e54fa5d05992924824425a0dca823618fe3cec13 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 28 Nov 2025 20:32:16 +0200 Subject: [PATCH] Init --- __init__.py | 9 +++ nodes_model_loading.py | 18 ++++- nodes_sampler.py | 10 +++ steadydancer/mobilenetv2_dcd.py | 102 +++++++++++++++++++++++ steadydancer/nodes.py | 60 ++++++++++++++ steadydancer/small_archs.py | 138 ++++++++++++++++++++++++++++++++ wanvideo/modules/model.py | 126 +++++++++++++++-------------- 7 files changed, 400 insertions(+), 63 deletions(-) create mode 100644 steadydancer/mobilenetv2_dcd.py create mode 100644 steadydancer/nodes.py create mode 100644 steadydancer/small_archs.py diff --git a/__init__.py b/__init__.py index 41f629c..ebb563c 100644 --- a/__init__.py +++ b/__init__.py @@ -77,6 +77,13 @@ except Exception as e: OVI_NODE_CLASS_MAPPINGS = {} OVI_NODE_DISPLAY_NAME_MAPPINGS = {} +try: + from .steadydancer.nodes import NODE_CLASS_MAPPINGS as STEADYDANCER_NODE_CLASS_MAPPINGS, NODE_DISPLAY_NAME_MAPPINGS as STEADYDANCER_NODE_DISPLAY_NAME_MAPPINGS +except Exception as e: + log.warning(f"WanVideoWrapper WARNING: SteadyDancer nodes not available due to error in importing them: {e}") + STEADYDANCER_NODE_CLASS_MAPPINGS = {} + STEADYDANCER_NODE_DISPLAY_NAME_MAPPINGS = {} + NODE_CLASS_MAPPINGS.update(RECAM_MASTER_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(UNIANIMATE_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(SKYREELS_NODE_CLASS_MAPPINGS) @@ -100,6 +107,7 @@ NODE_CLASS_MAPPINGS.update(LYNX_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(OVI_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(FLASHVSR_NODE_CLASS_MAPPINGS) NODE_CLASS_MAPPINGS.update(MOCHA_NODE_CLASS_MAPPINGS) +NODE_CLASS_MAPPINGS.update(STEADYDANCER_NODE_CLASS_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(RECAM_MASTER_NODE_DISPLAY_NAME_MAPPINGS) NODE_DISPLAY_NAME_MAPPINGS.update(UNIANIMATE_NODE_DISPLAY_NAME_MAPPINGS) @@ -124,5 +132,6 @@ NODE_DISPLAY_NAME_MAPPINGS.update(LYNX_NODE_DISPLAY_NAME_MAPPINGS) 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(MOCHA_NODE_DISPLAY_NAME_MAPPINGS) +NODE_DISPLAY_NAME_MAPPINGS.update(STEADYDANCER_NODE_DISPLAY_NAME_MAPPINGS) __all__ = ["NODE_CLASS_MAPPINGS", "NODE_DISPLAY_NAME_MAPPINGS"] \ No newline at end of file diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 32a2ff4..477bee4 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -872,6 +872,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None, key = name.replace("_orig_mod.", "") value=sd[key] + keep_fp32 = ["patch_embedding", "motion_encoder", "condition_embedding"] if gguf: dtype_to_use = torch.float32 if "patch_embedding" in name or "motion_encoder" in name else base_dtype @@ -883,7 +884,7 @@ def load_weights(transformer, sd=None, weight_dtype=None, base_dtype=None, dtype_to_use = value.dtype if "bias" in name or "img_emb" in name: dtype_to_use = base_dtype - if "patch_embedding" in name or "motion_encoder" in name: + if any(k in name for k in keep_fp32): dtype_to_use = torch.float32 if "modulation" in name or "norm" in name: dtype_to_use = value.dtype if value.dtype == torch.float32 else base_dtype @@ -1499,6 +1500,21 @@ class WanVideoModelLoader: device=device, ) + # SteadyDancer + if "condition_embedding_align.cross_attn.in_proj_bias" in sd: + from .steadydancer.mobilenetv2_dcd import DYModule + from .steadydancer.small_archs import PoseRefNetNoBNV3, FactorConv3d + in_dim_c = 16 + transformer.patch_embedding_fuse = nn.Conv3d(in_channels + in_dim_c + in_dim_c, dim, kernel_size=patch_size, stride=patch_size) # x, fused pose, aligned pose + transformer.patch_embedding_ref_c = nn.Conv3d(in_dim_c, dim, kernel_size=patch_size, stride=patch_size) # ref_c + transformer.condition_embedding_spatial = DYModule(inp=in_dim_c, oup=in_dim_c) # Spatial Structure Adaptive Extractor + transformer.condition_embedding_temporal = nn.Sequential( # Temporal Motion Coherence Module + FactorConv3d(in_channels=in_dim_c, out_channels=in_dim_c, kernel_size=(3, 3, 3), stride=1), nn.SiLU(), + FactorConv3d(in_channels=in_dim_c, out_channels=in_dim_c, kernel_size=(3, 3, 3), stride=1), nn.SiLU(), + FactorConv3d(in_channels=in_dim_c, out_channels=in_dim_c, kernel_size=(3, 3, 3), stride=1), nn.SiLU()) + transformer.condition_embedding_align = PoseRefNetNoBNV3(in_channels_x=16, in_channels_c=16, hidden_dim=128, num_heads=8) # Frame-wise Attention Alignment Unit + + comfy_model.diffusion_model = transformer comfy_model.load_device = transformer_load_device patcher = comfy.model_patcher.ModelPatcher(comfy_model, device, offload_device) diff --git a/nodes_sampler.py b/nodes_sampler.py index 13ae43d..33de8d7 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -1171,6 +1171,15 @@ class WanVideoSampler: latent = add_noise(ttm_reference_latents, noise, timesteps[ttm_start_step].to(noise.device)).to(latent) + # SteadyDancer + sdance_embeds = image_embeds.get("sdance_embeds", None) + sdancer_input = None + if sdance_embeds is not None: + print("Using SteadyDancer embeddings") + print(f"SteadyDancer embeds keys: {list(sdance_embeds.keys())}") + sdancer_input = sdance_embeds.copy() + sdancer_input = dict_to_device(sdancer_input, device, dtype) + #region model pred 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, @@ -1450,6 +1459,7 @@ class WanVideoSampler: "flashvsr_LQ_latent": flashvsr_LQ_latent, # FlashVSR LQ latent for upsampling "flashvsr_strength": flashvsr_strength, # FlashVSR strength "num_cond_latents": len(all_indices) if transformer.is_longcat else None, + "sdancer_input": sdancer_input, # SteadyDancer input } batch_size = 1 diff --git a/steadydancer/mobilenetv2_dcd.py b/steadydancer/mobilenetv2_dcd.py new file mode 100644 index 0000000..4fa3e55 --- /dev/null +++ b/steadydancer/mobilenetv2_dcd.py @@ -0,0 +1,102 @@ +# Modify from https://github.com/liyunsheng13/dcd/blob/main/models/imagenet/mobilenetv2_dcd.py + +import torch +import torch.nn as nn +import torch.nn.functional as F + + +class Hsigmoid(nn.Module): + def __init__(self, inplace=True): + super(Hsigmoid, self).__init__() + self.inplace = inplace + + def forward(self, x): + return F.relu6(x + 3., inplace=self.inplace) / 3. + + +class DYModule(nn.Module): + def __init__(self, inp, oup, fc_squeeze=8): + super(DYModule, self).__init__() + self.conv = nn.Conv2d(inp, oup, 1, 1, 0, bias=False) + if inp < oup: + self.mul = 4 + reduction = 8 + self.avg_pool = nn.AdaptiveAvgPool2d(2) + else: + self.mul = 1 + reduction = 2 + self.avg_pool = nn.AdaptiveAvgPool2d(1) + + self.dim = min((inp * self.mul) // reduction, oup // reduction) + while self.dim ** 2 > inp * self.mul * 2: + reduction *= 2 + self.dim = min((inp * self.mul) // reduction, oup // reduction) + if self.dim < 4: + self.dim = 4 + + squeeze = max(inp * self.mul, self.dim ** 2) // fc_squeeze + if squeeze < 4: + squeeze = 4 + self.conv_q = nn.Conv2d(inp, self.dim, 1, 1, 0, bias=False) + + self.fc = nn.Sequential( + nn.Linear(inp * self.mul, squeeze, bias=False), + SEModule_small(squeeze), + ) + self.fc_phi = nn.Linear(squeeze, self.dim ** 2, bias=False) + self.fc_scale = nn.Linear(squeeze, oup, bias=False) + self.hs = Hsigmoid() + self.conv_p = nn.Conv2d(self.dim, oup, 1, 1, 0, bias=False) + # self.bn1 = nn.BatchNorm2d(self.dim) + self.bn1 = nn.GroupNorm(num_groups=4, num_channels=self.dim) + # self.bn2 = nn.BatchNorm1d(self.dim) + self.bn2 = nn.GroupNorm(num_groups=4, num_channels=self.dim) + + def forward(self, x): + r = self.conv(x) + + b, c, h, w = x.size() + y = self.avg_pool(x).view(b, c * self.mul) + y = self.fc(y) + dy_phi = self.fc_phi(y).view(b, self.dim, self.dim) + dy_scale = self.hs(self.fc_scale(y)).view(b, -1, 1, 1) + r = dy_scale.expand_as(r) * r + + x = self.conv_q(x) + x = self.bn1(x) + x = x.view(b, -1, h * w) + x = self.bn2(torch.matmul(dy_phi, x)) + x + x = x.view(b, -1, h, w) + x = self.conv_p(x) + return x + r + + +class SEModule_small(nn.Module): + def __init__(self, channel): + super(SEModule_small, self).__init__() + self.fc = nn.Sequential( + nn.Linear(channel, channel, bias=False), + Hsigmoid() + ) + + def forward(self, x): + y = self.fc(x) + return x * y + + +class SEModule(nn.Module): + def __init__(self, channel, reduction=4): + super(SEModule, self).__init__() + self.avg_pool = nn.AdaptiveAvgPool2d(1) + self.fc = nn.Sequential( + nn.Linear(channel, channel // reduction, bias=False), + nn.ReLU(inplace=True), + nn.Linear(channel // reduction, channel, bias=False), + Hsigmoid() + ) + + def forward(self, x): + b, c, _, _ = x.size() + y = self.avg_pool(x).view(b, c) + y = self.fc(y).view(b, c, 1, 1) + return x * y.expand_as(x) diff --git a/steadydancer/nodes.py b/steadydancer/nodes.py new file mode 100644 index 0000000..581f5f1 --- /dev/null +++ b/steadydancer/nodes.py @@ -0,0 +1,60 @@ +import os +import torch +import numpy as np +from ..utils import log + +from accelerate import init_empty_weights +from accelerate.utils import set_module_tensor_to_device + +import comfy.model_management as mm +from comfy.utils import load_torch_file, ProgressBar +import folder_paths + +script_directory = os.path.dirname(os.path.abspath(__file__)) +device = mm.get_torch_device() +offload_device = mm.unet_offload_device() + + +class WanVideoAddSteadyDancerEmbeds: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "embeds": ("WANVIDIMAGE_EMBEDS",), + "pose_latents_positive": ("LATENT",), + "strength": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 100.0, "step": 0.01, "tooltip": "Strength of the portrait 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": { + "pose_latents_negative": ("LATENT",), + "clip_vision_embeds": ("WANVIDIMAGE_CLIPEMBEDS",), + + } + } + + RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",) + RETURN_NAMES = ("image_embeds",) + FUNCTION = "add" + CATEGORY = "WanVideoWrapper" + + def add(self, embeds, pose_latents_positive, strength, start_percent=0.0, end_percent=1.0, pose_latents_negative=None, clip_vision_embeds=None): + sdance_embeds = { + "cond_pos": pose_latents_positive["samples"][0], + "cond_neg": pose_latents_negative["samples"][0] if pose_latents_negative else None, + "strength": strength, + "start_percent": start_percent, + "end_percent": end_percent, + "clip_fea": clip_vision_embeds, + } + + updated = dict(embeds) + updated["sdance_embeds"] = sdance_embeds + return (updated,) + + +NODE_CLASS_MAPPINGS = { + "WanVideoAddSteadyDancerEmbeds": WanVideoAddSteadyDancerEmbeds, + } +NODE_DISPLAY_NAME_MAPPINGS = { + "WanVideoAddSteadyDancerEmbeds": "WanVideo Add SteadyDancer Embeds", + } diff --git a/steadydancer/small_archs.py b/steadydancer/small_archs.py new file mode 100644 index 0000000..8787db3 --- /dev/null +++ b/steadydancer/small_archs.py @@ -0,0 +1,138 @@ +import torch +import torch.nn as nn + + +class FactorConv3d(nn.Module): + """ + (2+1)D decomposition of 3D convolution: 1xHxW spatial convolution → Swish → Tx1x1 temporal convolution + """ + def __init__(self, + in_channels: int, + out_channels: int, + kernel_size, + stride: int = 1, + dilation: int = 1): + super().__init__() + + if isinstance(kernel_size, int): + k_t, k_h, k_w = kernel_size, kernel_size, kernel_size + else: + k_t, k_h, k_w = kernel_size + + pad_t = (k_t - 1) * dilation // 2 + pad_hw = (k_h - 1) * dilation // 2 + + self.spatial = nn.Conv3d( + in_channels, in_channels, + kernel_size=(1, k_h, k_w), + stride=(1, stride, stride), + padding=(0, pad_hw, pad_hw), + dilation=(1, dilation, dilation), + groups=in_channels, + bias=False + ) + + self.temporal = nn.Conv3d( + in_channels, out_channels, + kernel_size=(k_t, 1, 1), + stride=(stride, 1, 1), + padding=(pad_t, 0, 0), + dilation=(dilation, 1, 1), + bias=True + ) + + self.act = nn.SiLU() + + def forward(self, x): + x = self.spatial(x) + x = self.act(x) + x = self.temporal(x) + return x + + +class LayerNorm2D(nn.Module): + """ + LayerNorm over C for a 4-D tensor (B, C, H, W) + """ + def __init__(self, num_channels, eps=1e-5, affine=True): + super().__init__() + self.num_channels = num_channels + self.eps = eps + self.affine = affine + if affine: + self.weight = nn.Parameter(torch.ones(1, num_channels, 1, 1)) + self.bias = nn.Parameter(torch.zeros(1, num_channels, 1, 1)) + + def forward(self, x): + # x: (B, C, H, W) + mean = x.mean(dim=1, keepdim=True) # (B, 1, H, W) + var = x.var (dim=1, keepdim=True, unbiased=False) + x = (x - mean) / torch.sqrt(var + self.eps) + if self.affine: + x = x * self.weight + self.bias + return x + + +class PoseRefNetNoBNV3(nn.Module): + def __init__(self, + in_channels_c: int, + in_channels_x: int, + hidden_dim: int = 256, + num_heads: int = 8, + dropout: float = 0.1): + super().__init__() + self.d_model = hidden_dim + self.nhead = num_heads + + self.proj_p = nn.Conv2d(in_channels_c, hidden_dim, kernel_size=1) + self.proj_r = nn.Conv2d(in_channels_x, hidden_dim, kernel_size=1) + + self.proj_p_back = nn.Conv2d(hidden_dim, in_channels_c, kernel_size=1) + + self.cross_attn = nn.MultiheadAttention(hidden_dim, + num_heads=num_heads, + dropout=dropout) + + self.ffn_pose = nn.Sequential( + nn.Conv2d(hidden_dim, hidden_dim, kernel_size=1), + nn.SiLU(), + nn.Conv2d(hidden_dim, hidden_dim, kernel_size=1) + ) + + self.norm1 = LayerNorm2D(hidden_dim) + self.norm2 = LayerNorm2D(hidden_dim) + + def forward(self, pose, ref, mask=None): + """ + pose : (B, C1, T, H, W) + ref : (B, C2, T, H, W) + mask : (B, T*H*W) optional key_padding_mask + return: (B, d_model, T, H, W) + """ + B, _, T, H, W = pose.shape + L = H * W + + p_trans = pose.permute(0, 2, 1, 3, 4).contiguous().flatten(0, 1) + r_trans = ref.permute(0, 2, 1, 3, 4).contiguous().flatten(0, 1) + + p_trans = self.proj_p(p_trans) + r_trans = self.proj_r(r_trans) + + p_trans = p_trans.flatten(2).transpose(1, 2) + r_trans = r_trans.flatten(2).transpose(1, 2) + + out = self.cross_attn(query=r_trans, + key=p_trans, + value=p_trans, + key_padding_mask=mask)[0] + + out = out.transpose(1, 2).contiguous().view(B*T, -1, H, W) + out = self.norm1(out) + + ffn_out = self.ffn_pose(out) + out = out + ffn_out + out = self.norm2(out) + out = self.proj_p_back(out) + out = out.view(B, T, -1, H, W).contiguous().transpose(1, 2) + + return out diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index cdb2341..c10a480 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -1623,55 +1623,22 @@ class AudioInjector_WAN(nn.Module): class WanModel(torch.nn.Module): def __init__(self, model_type='t2v', - patch_size=(1, 2, 2), - text_len=512, - in_dim=16, - dim=2048, - in_features=5120, - out_features=5120, - ffn_dim=8192, - ffn2_dim=8192, - freq_dim=256, - text_dim=4096, - out_dim=16, - num_heads=16, - num_layers=32, - qk_norm=True, - cross_attn_norm=True, - eps=1e-6, - attention_mode='sdpa', - rope_func='comfy', - rms_norm_function='default', - main_device=torch.device('cuda'), - offload_device=torch.device('cpu'), - dtype=torch.float16, - teacache_coefficients=[], - magcache_ratios=[], - vace_layers=None, - vace_in_dim=None, - inject_sample_info=False, - add_ref_conv=False, - in_dim_ref_conv=16, - add_control_adapter=False, - in_dim_control_adapter=24, - use_motion_attn=False, + patch_size=(1, 2, 2), text_len=512, + in_dim=16, dim=2048, in_features=5120, out_features=5120, ffn_dim=8192, ffn2_dim=8192, + freq_dim=256, text_dim=4096, out_dim=16, 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, + teacache_coefficients=[], magcache_ratios=[], vace_layers=None, vace_in_dim=None, + inject_sample_info=False, add_ref_conv=False, in_dim_ref_conv=16, add_control_adapter=False, + in_dim_control_adapter=24, use_motion_attn=False, #s2v - cond_dim=0, - audio_dim=1024, - num_audio_token=4, - enable_adain=False, - adain_mode="attn_norm", - audio_inject_layers=[0, 4, 8, 12, 16, 20, 24, 27, 30, 33, 36, 39], - zero_timestep=False, - humo_audio=False, + cond_dim=0, audio_dim=1024, num_audio_token=4, enable_adain=False, zero_timestep=False, humo_audio=False, + adain_mode="attn_norm", audio_inject_layers=[0, 4, 8, 12, 16, 20, 24, 27, 30, 33, 36, 39], # WanAnimate - is_wananimate=False, - motion_encoder_dim=512, + is_wananimate=False, motion_encoder_dim=512, # lynx - lynx_ip_layers=None, - lynx_ref_layers=None, - # ovi - is_ovi_audio_model=False, + lynx_ip_layers=None, lynx_ref_layers=None, # LongCat is_longcat=False, ): @@ -2190,8 +2157,7 @@ class WanModel(torch.nn.Module): self, x, t, context, seq_len, is_uncond=False, current_step_percentage=0.0, current_step=0, last_step=0, total_steps=50, - clip_fea=None, - y=None, + clip_fea=None, y=None, device=torch.device('cuda'), freqs=None, enhance_enabled=False, @@ -2203,8 +2169,7 @@ class WanModel(torch.nn.Module): fps_embeds=None, fun_ref=None, fun_camera=None, audio_proj=None, audio_scale=1.0, - uni3c_data=None, - controlnet=None, + uni3c_data=None, controlnet=None, add_cond=None, attn_cond=None, nag_params={}, nag_context=None, multitalk_audio=None, @@ -2227,6 +2192,7 @@ class WanModel(torch.nn.Module): flashvsr_LQ_latent=None, flashvsr_strength=1.0, num_cond_latents=None, add_text_emb=None, + sdancer_input=None # SteadyDancer ): r""" Forward pass through the diffusion model @@ -2313,7 +2279,11 @@ class WanModel(torch.nn.Module): freqs = freqs.to(device) _, F, H, W = x[0].shape + + if sdancer_input is not None: + x_noise_clone = torch.stack(x) + # I2V if y is not None: if hasattr(self, "randomref_embedding_pose") and unianim_data is not None: if unianim_data['start_percent'] <= current_step_percentage <= unianim_data['end_percent']: @@ -2321,7 +2291,7 @@ class WanModel(torch.nn.Module): if random_ref_emb is not None: y[0].add_(random_ref_emb, alpha=unianim_data["strength"]) x = [torch.cat([u, v], dim=0) for u, v in zip(x, y)] - + #uni3c controlnet if uni3c_data is not None: render_latent = uni3c_data["render_latent"].to(self.base_dtype) @@ -2330,13 +2300,27 @@ class WanModel(torch.nn.Module): hidden_states = torch.cat([hidden_states, torch.zeros_like(hidden_states[:, :4])], dim=1) render_latent = torch.cat([hidden_states[:, :20], render_latent], dim=1) - # patch embed - if control_lora_enabled: - self.expanded_patch_embedding.to(self.main_device) - x = [self.expanded_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in x] + # SteadyDancer + if sdancer_input is not None: + sdancer_cond = sdancer_input["cond_pos"] if not is_uncond else sdancer_input["cond_neg"] + condition_temporal = [self.condition_embedding_temporal(c.unsqueeze(0).float()).to(self.base_dtype) for c in [sdancer_cond]] # Temporal Motion Coherence Module. + sdancer_cond = sdancer_cond.unsqueeze(0) + bs, _, time_steps, _, _ = sdancer_cond.shape + condition_reshape = rearrange(sdancer_cond, 'b c t h w -> (b t) c h w') + condition_spatial = self.condition_embedding_spatial(condition_reshape.float()).to(self.base_dtype) # Spatial Structure Adaptive Extractor. + condition_spatial = rearrange(condition_spatial, '(b t) c h w -> b c t h w', t=time_steps, b=bs) + condition_fused = sdancer_cond + condition_temporal[0] + condition_spatial # Hierarchical Aggregation (1): condition, temporal condition, spatial condition + condition_aligned = self.condition_embedding_align(condition_fused.float(), x_noise_clone).to(self.base_dtype) # Frame-wise Attention Alignment Unit. else: - 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] + # patch embed + if control_lora_enabled: + self.expanded_patch_embedding.to(self.main_device) + x = [self.expanded_patch_embedding(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in x] + else: + 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] + + orig_frames = x[0].shape[1] # ovi audio model if self.audio_model is not None: @@ -2365,13 +2349,29 @@ class WanModel(torch.nn.Module): fun_camera = self.control_adapter(fun_camera) x = [u + v for u, v in zip(x, fun_camera)] + # SteadyDancer + if sdancer_input is not None: + ref_x = y[0][4:, :1] # reuse I2V input as reference, slice mask off + msk = torch.ones(4, 1, H, W, device=ref_x.device) # new mask goes in middle + ref_x = [torch.concat([ref_x, msk, ref_x])] + ref_c = sdancer_cond[0][:, :1] + ref_c = [torch.concat([ref_c, msk * 0, ref_c])] # zero mask for cond ref + # Condition Fusion/Injection, Hierarchical Aggregation (2): x, fused condition, aligned condition + x = [self.patch_embedding_fuse(torch.cat([u[None], c[None], a[None]], 1)) for u, c, a in zip(x, condition_fused, condition_aligned)] + # Condition Augmentation: x_cond, ref_x, ref_c + ref_x = [self.patch_embedding(r.unsqueeze(0).float()).to(self.base_dtype) for r in ref_x] + ref_c = [self.patch_embedding_ref_c(r[:16].unsqueeze(0).float()).to(self.base_dtype) for r in ref_c] + F += ref_x[0].shape[2] + ref_c[0].shape[2] # update frame count for rope + x = [torch.cat([r, u, v], dim=2) for r, u, v in zip(x, ref_x, ref_c)] + seq_len = torch.tensor([u.flatten(2).transpose(1, 2).size(1) for u in x], dtype=torch.int32).max() # update seq len + # grid sizes and seq len grid_sizes = torch.stack([torch.tensor(u.shape[2:], device=device, dtype=torch.long) for u in x]) original_grid_sizes = grid_sizes.clone() x = [u.flatten(2).transpose(1, 2) for u in x] self.original_seq_len = x[0].shape[1] seq_lens = torch.tensor([u.size(1) for u in x], dtype=torch.int32) - assert seq_lens.max() <= seq_len + assert seq_lens.max() <= seq_len, f"max seq len {seq_lens.max()} exceeds provided seq_len {seq_len}" cond_mask_weight = None if self.trainable_cond_mask is not None: @@ -2570,11 +2570,13 @@ class WanModel(torch.nn.Module): # clip vision embedding clip_embed = None if clip_fea is not None and hasattr(self, "img_emb"): - clip_fea = clip_fea.to(self.main_device) if self.offload_img_emb: self.img_emb.to(self.main_device) - clip_embed = self.img_emb(clip_fea) # bs x 257 x dim - #context = torch.concat([context_clip, context], dim=1) + clip_embed = self.img_emb(clip_fea.to(self.main_device)) # bs x 257 x dim + if sdancer_input is not None: + clip_fea_c = sdancer_input.get("clip_fea_c", None) + if clip_fea_c is not None: + clip_embed += self.img_emb(clip_fea_c.to(self.main_device)) if self.offload_img_emb: self.img_emb.to(self.offload_device, non_blocking=self.use_non_blocking) @@ -3028,7 +3030,7 @@ class WanModel(torch.nn.Module): x_ovi = [u.float() for u in x_ovi] x = self.unpatchify(x, original_grid_sizes) # type: ignore[arg-type] - x = [u.float() for u in x] + x = [u[:, :orig_frames, ...].float() for u in x] return (x, x_ovi, pred_id) if pred_id is not None else (x, x_ovi, None) def unpatchify(self, x, grid_sizes):