diff --git a/nodes.py b/nodes.py index ece8d3e..86a5f7f 100644 --- a/nodes.py +++ b/nodes.py @@ -22,6 +22,30 @@ offload_device = mm.unet_offload_device() VAE_STRIDE = (4, 8, 8) PATCH_SIZE = (1, 2, 2) +class WanVideoAddVideoPromptEmbeds: + @classmethod + def INPUT_TYPES(s): + return {"required": { + "image_embeds": ("WANVIDIMAGE_EMBEDS",), + "video_prompt_embeds": ("WANVIDIMAGE_EMBEDS",), + "video_prompt_latents": ("LATENT", ), + "text_embeds": ("WANVIDEOTEXTEMBEDS", ), + } + } + + RETURN_TYPES = ("WANVIDIMAGE_EMBEDS",) + RETURN_NAMES = ("image_embeds",) + FUNCTION = "add" + CATEGORY = "WanVideoWrapper" + EXPERIMENTAL = True + + def add(self, image_embeds, video_prompt_embeds, video_prompt_latents, text_embeds): + updated = dict(image_embeds) + updated["video_prompt_embeds"] = video_prompt_embeds + updated["video_prompt_embeds"]["video_prompt_latents"] = video_prompt_latents["samples"][0] + updated["video_prompt_embeds"]["text_embeds"] = text_embeds + return (updated,) + class WanVideoEnhanceAVideo: @classmethod @@ -2206,6 +2230,7 @@ NODE_CLASS_MAPPINGS = { "WanVideoAnimateEmbeds": WanVideoAnimateEmbeds, "WanVideoAddLucyEditLatents": WanVideoAddLucyEditLatents, "WanVideoSchedulerSA_ODE": WanVideoSchedulerSA_ODE, + "WanVideoAddVideoPromptEmbeds": WanVideoAddVideoPromptEmbeds, } NODE_DISPLAY_NAME_MAPPINGS = { diff --git a/nodes_model_loading.py b/nodes_model_loading.py index a35fd9b..d6aa107 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -1346,6 +1346,7 @@ class WanVideoModelLoader: "lynx_ip_layers": lynx_ip_layers, "lynx_ref_layers": lynx_ref_layers, "is_longcat": dim == 4096, + "is_VAP": True if "patch_embedding_mot_ref.weight" in sd else False } diff --git a/nodes_sampler.py b/nodes_sampler.py index 6bf72b8..73a2bb6 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -289,6 +289,7 @@ class WanVideoSampler: phantom_latents = fun_ref_image = ATI_tracks = None add_cond = attn_cond = attn_cond_neg = noise_pred_flipped = None humo_audio = humo_audio_neg = None + image_cond_mot_ref = None #I2V image_cond = image_embeds.get("image_embeds", None) @@ -476,6 +477,22 @@ class WanVideoSampler: phantom_start_percent = image_embeds.get("phantom_start_percent", 0.0) phantom_end_percent = image_embeds.get("phantom_end_percent", 1.0) + # Video-as-prompt (VAP) + mot_ref_clip_embeds = mot_ref_context = None + video_prompt_embeds = image_embeds.get("video_prompt_embeds", None) + if video_prompt_embeds is not None: + image_cond_mot_ref = video_prompt_embeds.get("image_embeds", None) + print("image_cond_mot_ref shape:", image_cond_mot_ref.shape) + image_cond_mask_ = video_prompt_embeds.get("mask", None) + if image_cond_mask_ is not None: + image_cond_mot_ref = torch.cat([image_cond_mask_, image_cond_mot_ref]) + latents_mot_ref = video_prompt_embeds.get("video_prompt_latents", None) + print("latents_mot_ref shape:", latents_mot_ref.shape) + x_mot_ref = torch.cat([latents_mot_ref, image_cond_mot_ref], dim=0) + mot_ref_context = video_prompt_embeds.get("text_embeds", None) + mot_ref_clip_embeds = video_prompt_embeds.get("clip_context", None) + print("x_mot_ref shape:", x_mot_ref.shape) + # CLIP image features clip_fea = image_embeds.get("clip_context", None) if clip_fea is not None: @@ -1395,7 +1412,10 @@ class WanVideoSampler: "ovi_negative_text_embeds": ovi_negative_text_embeds, # Audio latent model negative text embeds for Ovi "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 # number of cond latents LongCat to separate attention + "num_cond_latents": len(all_indices) if transformer.is_longcat else None, # number of cond latents LongCat to separate attention + "x_mot_ref": [x_mot_ref.to(z)] if image_cond_mot_ref is not None else None, # motion reference latents for VAP + "mot_ref_context": mot_ref_context if image_cond_mot_ref is not None else None, # motion reference context for VAP + "mot_ref_clip_embeds": mot_ref_clip_embeds, # motion reference clip features for VAP } batch_size = 1 diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index cb327f5..72b1732 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -747,6 +747,34 @@ class WanT2VCrossAttention(WanSelfAttention): return torch.cat([torch.zeros((b, num_cond_latents_thw, x.shape[-1]), dtype=x.dtype, device=x.device), self.o(x)], dim=1).contiguous() return self.o(x) + +class WanT2VCrossAttentionMOTRef(WanSelfAttention): + + def __init__(self, in_features, out_features, num_heads, kv_dim=None, qk_norm=True, eps=1e-6, attention_mode='sdpa', rms_norm_function="default", head_norm=False): + super().__init__(in_features, out_features, num_heads, qk_norm, eps, kv_dim=kv_dim, rms_norm_function=rms_norm_function, head_norm=head_norm) + self.k_img = nn.Linear(in_features, out_features) + self.v_img = nn.Linear(in_features, out_features) + self.norm_k_img = WanRMSNorm(out_features, eps=eps) if qk_norm else nn.Identity() + self.attention_mode = attention_mode + self.ip_adapter = None + self.k_fusion = None + + def forward(self, x, context, grid_sizes=None, clip_embed=None, rope_func="comfy", **kwargs): + b, n, d = x.size(0), self.num_heads, self.head_dim + + q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype), num_chunks=2 if rope_func == "comfy_chunked" else 1).to(x.dtype).view(b, -1, n, d) + k = self.norm_k(self.k(context).to(self.norm_k.weight.dtype)).to(x.dtype).view(b, -1, n, d) + v = self.v(context).view(b, -1, n, d) + + x = attention(q, k, v, attention_mode=self.attention_mode).flatten(2) + + if clip_embed is not None: + k_img = self.norm_k_img(self.k_img(clip_embed)).view(b, -1, n, d) + v_img = self.v_img(clip_embed).view(b, -1, n, d) + img_x = attention(q, k_img, v_img, attention_mode=self.attention_mode).flatten(2) + x = x + img_x + + return self.o(x) class WanI2VCrossAttention(WanSelfAttention): @@ -897,7 +925,7 @@ class WanAttentionBlock(nn.Module): cross_attn_type, 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", use_motion_attn=False, use_humo_audio_attn=False, face_fuser_block=False, lynx_ip_layers=None, lynx_ref_layers=None, - block_idx=0, is_longcat=False): + block_idx=0, mot_ref_block=False, is_longcat=False): super().__init__() self.dim = out_features self.ffn_dim = ffn_dim @@ -913,6 +941,7 @@ class WanAttentionBlock(nn.Module): self.dense_block = False self.dense_attention_mode = "sageattn" self.block_idx = block_idx + self.mot_ref_block = mot_ref_block self.kv_cache = None self.use_motion_attn = use_motion_attn @@ -952,6 +981,16 @@ class WanAttentionBlock(nn.Module): self.seg_idx = None + # video-as-prompt (VAP) + if mot_ref_block: + self.norm1_mot_ref = WanLayerNorm(self.dim, eps) + self.norm2_mot_ref = WanLayerNorm(self.dim, eps) + self.norm3_mot_ref = WanLayerNorm(out_features, eps, elementwise_affine=True) if cross_attn_norm else nn.Identity() + self.self_attn_mot_ref = WanSelfAttention(in_features, out_features, num_heads, qk_norm, eps, self.attention_mode, rms_norm_function=rms_norm_function) + self.cross_attn_mot_ref = WanT2VCrossAttentionMOTRef(in_features, out_features, num_heads, qk_norm, eps, rms_norm_function=rms_norm_function) + self.modulation_mot_ref = nn.Parameter(torch.randn(1, 6, out_features) / in_features**0.5) + self.ffn_mot_ref = nn.Sequential(nn.Linear(in_features, ffn_dim), nn.GELU(approximate='tanh'), nn.Linear(ffn2_dim, out_features)) + # HuMo audio cross-attn if use_humo_audio_attn: self.audio_cross_attn_wrapper = AudioCrossAttentionWrapper(in_features, out_features, num_heads, qk_norm, eps, kv_dim=1536) @@ -1045,7 +1084,8 @@ class WanAttentionBlock(nn.Module): 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, num_cond_latents=None, #longcat image cond amount - ): + # VAP + x_mot_ref=None, context_mot_ref=None, grid_sizes_mot_ref=None, e_mot_ref=None, freqs_mot_ref=None, clip_embed_mot_ref=None, num_mot_ref=1): r""" Args: x(Tensor): Shape [B, L, C] @@ -1081,6 +1121,17 @@ class WanAttentionBlock(nn.Module): input_x = torch.concat([input_x, input_x_ip], dim=1) self.kv_cache = None + # video-as-prompt motion reference + use_mot_ref = x_mot_ref is not None and self.mot_ref_block + if use_mot_ref: + #import einops + # shift_msa_mot_ref, scale_msa_mot_ref, gate_msa_mot_ref, shift_mlp_mot_ref, scale_mlp_mot_ref, gate_mlp_mot_ref = self.get_mod(e_mot_ref.to(x.device), self.modulation_mot_ref) + # norm_x_mot_ref = einops.rearrange(self.norm1_mot_ref(x_mot_ref.to(scale_msa_mot_ref.dtype)), 'b (n t) c -> b n t c', n=num_mot_ref) + # input_x_mot_ref = self.modulate(norm_x_mot_ref, shift_msa_mot_ref, scale_msa_mot_ref).to(input_dtype) + # input_x_mot_ref = einops.rearrange(input_x_mot_ref, 'b n t c -> b (n t) c', n=num_mot_ref) + shift_msa_mot_ref, scale_msa_mot_ref, gate_msa_mot_ref, shift_mlp_mot_ref, scale_mlp_mot_ref, gate_mlp_mot_ref = self.get_mod(e_mot_ref.to(x.device), self.modulation_mot_ref) + input_x_mot_ref = self.modulate(self.norm1_mot_ref(x_mot_ref.to(shift_msa_mot_ref.dtype)), shift_msa_mot_ref, scale_msa_mot_ref).to(input_dtype) + if x_ovi is not None: shift_msa_ovi, scale_msa_ovi, gate_msa_ovi, shift_mlp_ovi, scale_mlp_ovi, gate_mlp_ovi = self.get_mod(e_ovi.to(x.device), self.audio_block.modulation) input_x_ovi = self.modulate(self.audio_block.norm1(x_ovi), shift_msa_ovi, scale_msa_ovi) @@ -1125,8 +1176,13 @@ class WanAttentionBlock(nn.Module): q, k, v = self.self_attn.qkv_fn_longcat(input_x) else: q, k, v = self.self_attn.qkv_fn(input_x) + if use_mot_ref: + q_mot_ref, k_mot_ref, v_mot_ref = self.self_attn_mot_ref.qkv_fn(input_x_mot_ref) + # Apply RoPE if self.rope_func == "comfy": q, k = apply_rope_comfy(q, k, freqs) + if use_mot_ref: + q_mot_ref, k_mot_ref = apply_rope_comfy(q_mot_ref, k_mot_ref, freqs_mot_ref) elif self.rope_func == "comfy_chunked": q, k = apply_rope_comfy_chunked(q, k, freqs) elif self.rope_func == "mocha": @@ -1196,6 +1252,14 @@ class WanAttentionBlock(nn.Module): x_noise = self.self_attn.forward(q[:, num_cond_latents_thw:].contiguous(), k, v, seq_lens) # merge x_cond and x_noise y = torch.cat([x_cond, x_noise], dim=1).contiguous() + elif use_mot_ref: + y_temp = self.self_attn_mot_ref.forward( + torch.cat([q, q_mot_ref], dim=1), + torch.cat([k, k_mot_ref], dim=1), + torch.cat([v, v_mot_ref], dim=1), + seq_lens + ) + y, y_mot_ref = torch.split(y_temp, [q.shape[1], q_mot_ref.shape[1]], dim=1) else: y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale) @@ -1227,8 +1291,13 @@ class WanAttentionBlock(nn.Module): else: if not is_longcat: x = x.addcmul(y, gate_msa) + if use_mot_ref: + #x_mot_ref = x_mot_ref.addcmul(einops.rearrange(y_mot_ref, 'b (n t) c -> b n t c', n=num_mot_ref), gate_msa_mot_ref) + #x_mot_ref = einops.rearrange(x_mot_ref, 'b n t c -> b (n t) c', n=num_mot_ref) + x_mot_ref = x_mot_ref.addcmul(y_mot_ref, gate_msa_mot_ref) else: x = x + (y.view(B, -1, N//T, C).float() * gate_msa).to(input_dtype).view(B, -1, C) + del y, gate_msa # cross-attention & ffn function @@ -1263,6 +1332,10 @@ class WanAttentionBlock(nn.Module): rope_func=self.rope_func, inner_t=inner_t, inner_c=inner_c, cross_freqs=cross_freqs, adapter_proj=adapter_proj, ip_scale=ip_scale, orig_seq_len=original_seq_len, lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale, num_cond_latents=num_cond_latents) x = x.to(input_dtype) + if use_mot_ref: + x_mot_ref = x_mot_ref + self.cross_attn_mot_ref(self.norm3_mot_ref(x_mot_ref.to(self.norm3_mot_ref.weight.dtype)).to(input_dtype), context_mot_ref, grid_sizes_mot_ref, + clip_embed=clip_embed_mot_ref) + x_mot_ref = x_mot_ref.to(input_dtype) # MultiTalk if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock): x_audio = self.audio_cross_attn(self.norm_x(x.to(self.norm_x.weight.dtype)).to(input_dtype), encoder_hidden_states=multitalk_audio_embedding, @@ -1270,7 +1343,7 @@ class WanAttentionBlock(nn.Module): x = x.add(x_audio, alpha=audio_scale) # 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.to(self.norm4.weight.dtype)).to(input_dtype), mtv_motion_tokens, mtv_motion_rotary_emb, grid_sizes, mtv_freqs) x = x.add(x_motion, alpha=mtv_strength) @@ -1317,7 +1390,20 @@ class WanAttentionBlock(nn.Module): x_ip = x_ip.addcmul(y_ip, gate_msa_ip) y_ip = self.ffn(torch.addcmul(shift_mlp_ip, self.norm2(x_ip), 1 + scale_mlp_ip)) x_ip = x_ip.addcmul(y_ip, gate_mlp_ip) - return x, x_ip, lynx_ref_feature, x_ovi + + if use_mot_ref: + # norm2_x_mot_ref = einops.rearrange(self.norm2_mot_ref(x_mot_ref.to(shift_mlp_mot_ref.dtype)), 'b (n t) c -> b n t c', n=num_mot_ref) + # mod_x_mot_ref = torch.addcmul(shift_mlp_mot_ref, norm2_x_mot_ref, 1 + scale_mlp_mot_ref) + # mod_x_mot_ref = einops.rearrange(mod_x_mot_ref, 'b n t c -> b (n t) c', n=num_mot_ref) + # x_ffn_mot_ref = self.ffn_mot_ref(mod_x_mot_ref.to(input_dtype)) + # x_ffn_mot_ref = einops.rearrange(x_ffn_mot_ref, 'b (n t) c -> b n t c', n=num_mot_ref) + # x_mot_ref = x_mot_ref.addcmul(x_ffn_mot_ref, gate_mlp_mot_ref) + # x_mot_ref = einops.rearrange(x_mot_ref, 'b n t c -> b (n t) c', n=num_mot_ref) + norm2_x_mot_ref = self.norm2_mot_ref(x_mot_ref.to(shift_mlp_mot_ref.dtype)) + mod_x_mot_ref = torch.addcmul(shift_mlp_mot_ref, norm2_x_mot_ref, 1 + scale_mlp_mot_ref) + x_ffn_mot_ref = self.ffn_mot_ref(mod_x_mot_ref.to(input_dtype)) + x_mot_ref = x_mot_ref.addcmul(x_ffn_mot_ref, gate_mlp_mot_ref) + return x, x_ip, lynx_ref_feature, x_ovi, x_mot_ref @torch.compiler.disable() def split_cross_attn_ffn(self, x, context, shift_mlp, scale_mlp, gate_mlp, clip_embed=None, grid_sizes=None): @@ -1520,7 +1606,7 @@ class MLPProj(torch.nn.Module): def forward(self, image_embeds): if hasattr(self, 'emb_pos'): image_embeds = image_embeds + self.emb_pos.to(image_embeds.device) - clip_extra_context_tokens = self.proj(image_embeds) + clip_extra_context_tokens = self.proj(image_embeds.to(self.proj[1].weight.dtype)).to(image_embeds.dtype) return clip_extra_context_tokens from .s2v.auxi_blocks import MotionEncoder_tc @@ -1622,47 +1708,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'), + 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, + 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", + 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 humo_audio=False, # WanAnimate is_wananimate=False, @@ -1670,8 +1731,8 @@ class WanModel(torch.nn.Module): # lynx lynx_ip_layers=None, lynx_ref_layers=None, - # ovi - is_ovi_audio_model=False, + # VAP + is_VAP = False, # LongCat is_longcat=False, ): @@ -1808,6 +1869,9 @@ class WanModel(torch.nn.Module): nn.SiLU(), ConvMLP(dim, dim * 4, kernel_size=7, padding=3), ) + + if is_VAP: + self.patch_embedding_mot_ref = nn.Conv3d(in_dim, dim, kernel_size=patch_size, stride=patch_size) self.original_patch_embedding = self.patch_embedding self.expanded_patch_embedding = self.patch_embedding @@ -1826,6 +1890,12 @@ class WanModel(torch.nn.Module): self.time_embedding = TimestepEmbedder(t_embed_dim=adaln_tembed_dim, frequency_embedding_size=freq_dim) + if is_VAP: + self.time_embedding_mot_ref = nn.Sequential(nn.Linear(freq_dim, dim), nn.SiLU(), nn.Linear(dim, dim)) + self.time_projection_mot_ref = nn.Sequential(nn.SiLU(), nn.Linear(dim, dim * 6)) + self.text_embedding_mot_ref = nn.Sequential(nn.Linear(text_dim, dim), nn.GELU(approximate='tanh'), nn.Linear(dim, dim)) + self.img_emb_mot_ref = MLPProj(1280, dim) + if vace_layers is not None: self.vace_layers = [i for i in range(0, self.num_layers, 2)] if vace_layers is None else vace_layers self.vace_in_dim = self.in_dim if vace_in_dim is None else vace_in_dim @@ -1859,13 +1929,15 @@ class WanModel(torch.nn.Module): else: cross_attn_type = 'no_cross_attn' + VAP_layers = [0, 4, 8, 12, 16, 20, 24, 28, 32, 36] + self.blocks = nn.ModuleList([ WanAttentionBlock(cross_attn_type, self.in_features, self.out_features, ffn_dim, ffn2_dim, num_heads, qk_norm, cross_attn_norm, eps, attention_mode=self.attention_mode, rope_func=self.rope_func, rms_norm_function=rms_norm_function, use_motion_attn=(i % 4 == 0 and use_motion_attn), use_humo_audio_attn=self.humo_audio, face_fuser_block = (i % 5 == 0 and is_wananimate), lynx_ip_layers=lynx_ip_layers, lynx_ref_layers=lynx_ref_layers, - block_idx=i, is_longcat=is_longcat) + block_idx=i, is_longcat=is_longcat, mot_ref_block=i in VAP_layers and is_VAP) for i in range(num_layers) ]) #MTV Crafter @@ -2139,12 +2211,15 @@ class WanModel(torch.nn.Module): return x.add(residual_out, alpha=strength) - def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, attn_cond_shape=None, steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None): + def rope_encode_comfy(self, t, h, w, freq_offset=0, t_start=0, attn_cond_shape=None, steps_t=None, steps_h=None, steps_w=None, ntk_alphas=[1,1,1], device=None, dtype=None, mot=False): 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 mot: + t_start = -t_len + if steps_t is None: steps_t = t_len if steps_h is None: @@ -2226,6 +2301,7 @@ class WanModel(torch.nn.Module): x_ovi=None, seq_len_ovi=None, ovi_negative_text_embeds=None, flashvsr_LQ_latent=None, flashvsr_strength=1.0, num_cond_latents=None, + x_mot_ref=None, mot_ref_context=None, mot_ref_clip_embeds=None, ): r""" Forward pass through the diffusion model @@ -2256,6 +2332,7 @@ class WanModel(torch.nn.Module): if mtv_motion_tokens is not None: bs, motion_seq_len = mtv_motion_tokens.shape[0], mtv_motion_tokens.shape[1] mtv_motion_tokens = torch.cat([mtv_motion_tokens, self.pad_motion_tokens.to(mtv_motion_tokens).expand(bs, motion_seq_len, -1)], dim=-1) + mtv_motion_tokens = mtv_motion_tokens.to(self.base_dtype) # Fantasy Portrait adapter_proj = ip_scale = None @@ -2346,6 +2423,18 @@ class WanModel(torch.nn.Module): d = self.dim // self.num_heads 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) + + # video-as-prompt motion ref + if x_mot_ref is not None: + x_mot_ref = [self.patch_embedding_mot_ref(u.unsqueeze(0).to(torch.float32)).to(x[0].dtype) for u in x_mot_ref] + grid_sizes_mot_ref = torch.stack([torch.tensor(u.shape[2:], device=device, dtype=torch.long) for u in x_mot_ref]) + + x_mot_ref = [u.flatten(2).transpose(1, 2) for u in x_mot_ref] + seq_lens_mot_ref = torch.tensor([u.size(1) for u in x_mot_ref], dtype=torch.int32) + x_mot_ref = torch.cat([torch.cat([u, u.new_zeros(1, seq_lens_mot_ref - u.size(1), u.size(2))], dim=1) for u in x_mot_ref]) + + x_mot_ref = x_mot_ref.to(self.main_device, self.base_dtype) + num_mot_ref = 1 # WanAnimate motion_vec = None @@ -2456,13 +2545,16 @@ class WanModel(torch.nn.Module): s2v_ref_latent.shape[4], t_start=max(30, F + 9), device=x.device, dtype=x.dtype) freqs = torch.cat([freqs, freqs_ref], dim=1) - self.cached_freqs = freqs self.cached_shape = current_shape self.cached_cond = has_cond self.cached_rope_k = self.rope_embedder.k self.cached_ntk_alphas = ntk_alphas + if x_mot_ref is not None: + freqs_mot_ref = self.rope_encode_comfy(F, H, W, mot=True, freq_offset=freq_offset, ntk_alphas=ntk_alphas, attn_cond_shape=attn_cond_shape, device=x.device, dtype=x.dtype) + + # Stand-In RoPE frequencies if x_ip is not None: # Generate RoPE frequencies for x_ip @@ -2503,6 +2595,10 @@ class WanModel(torch.nn.Module): time_embed_dtype = self.base_dtype 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 + if x_mot_ref is not None: + t_mod_ref = torch.tensor([1], device=t.device, dtype=t.dtype) + e_mot_ref = self.time_embedding_mot_ref(sinusoidal_embedding_1d(self.freq_dim, t_mod_ref.flatten()).to(time_embed_dtype)) # b, dim + e0_mot_ref = self.time_projection_mot_ref(e_mot_ref).unflatten(1, (6, self.dim)) # b, 6, dim else: time_embed_dtype = self.time_embedding.mlp[0].weight.dtype if time_embed_dtype not in [torch.float16, torch.bfloat16, torch.float32]: @@ -2592,6 +2688,11 @@ class WanModel(torch.nn.Module): context = self.text_embedding( torch.stack([torch.cat([u, u.new_zeros(self.text_len - u.size(0), u.size(1))]) for u in context]).to(text_embed_dtype)) + if mot_ref_context is not None: + context_mot_ref = mot_ref_context["prompt_embeds"] if not is_uncond else mot_ref_context["negative_prompt_embeds"] + context_mot_ref = self.text_embedding_mot_ref( + torch.stack([torch.cat([u, u.new_zeros(self.text_len - u.size(0), u.size(1))]) for u in context_mot_ref]).to(text_embed_dtype)) + if self.is_longcat: context[:, tokens:] = 0 @@ -2609,13 +2710,19 @@ class WanModel(torch.nn.Module): else: context = None - clip_embed = None + clip_embed = clip_embed_mot_ref = 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) + if self.offload_img_emb: + self.img_emb.to(self.offload_device, non_blocking=self.use_non_blocking) + if mot_ref_clip_embeds is not None: + mot_ref_clip_embeds = mot_ref_clip_embeds.to(self.main_device) + if self.offload_img_emb: + self.img_emb.to(self.main_device) + clip_embed_mot_ref = self.img_emb_mot_ref(mot_ref_clip_embeds) # bs x 257 x dim if self.offload_img_emb: self.img_emb.to(self.offload_device, non_blocking=self.use_non_blocking) @@ -2841,6 +2948,13 @@ class WanModel(torch.nn.Module): kwargs['grid_sizes_ovi'] = grid_sizes_ovi kwargs['seq_lens_ovi'] = seq_lens_ovi kwargs['freqs_ovi'] = freqs_ovi + if x_mot_ref is not None: + kwargs['context_mot_ref'] = context_mot_ref + kwargs['freqs_mot_ref'] = freqs_mot_ref + kwargs['grid_sizes_mot_ref'] = grid_sizes_mot_ref + kwargs['e_mot_ref'] = e0_mot_ref.to(self.base_dtype) + kwargs['num_mot_ref'] = num_mot_ref + kwargs['clip_embed_mot_ref'] = clip_embed_mot_ref if vace_data is not None: @@ -2928,7 +3042,9 @@ class WanModel(torch.nn.Module): if b in self.slg_blocks and is_uncond: if self.slg_start_percent <= current_step_percentage <= self.slg_end_percent: 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 + # ====run block start===== + x, x_ip, lynx_ref_feature, x_ovi, x_mot_ref = block(x, x_ip=x_ip, lynx_ref_feature=lynx_ref_feature, x_ovi=x_ovi, x_mot_ref=x_mot_ref, **kwargs) + # ====end run block===== 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 if block.has_face_fuser_block and motion_vec is not None: