From d504c961741d2cc7965e3ce4598894d5292dc7b9 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 26 Oct 2025 19:14:47 +0200 Subject: [PATCH] Separate attention for input images like in original --- nodes_sampler.py | 1 + wanvideo/modules/model.py | 34 +++++++++++++++++++++++++++++----- 2 files changed, 30 insertions(+), 5 deletions(-) diff --git a/nodes_sampler.py b/nodes_sampler.py index d8bdc4a..2980da3 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -1396,6 +1396,7 @@ 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 } batch_size = 1 diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index dd81240..f180977 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -691,10 +691,15 @@ class WanT2VCrossAttention(WanSelfAttention): def forward(self, x, context, grid_sizes=None, clip_embed=None, audio_proj=None, audio_scale=1.0, num_latent_frames=21, nag_params={}, nag_context=None, is_uncond=False, rope_func="comfy", inner_t=None, inner_c=None, cross_freqs=None, - adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, orig_seq_len=None, lynx_x_ip=None, lynx_ip_scale=1.0, **kwargs): + adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, orig_seq_len=None, lynx_x_ip=None, lynx_ip_scale=1.0, num_cond_latents=None, **kwargs): b, n, d = x.size(0), self.num_heads, self.head_dim + s = x.size(1) # compute query - if x.shape[-1] == 4096: #longcat + is_longcat = x.shape[-1] == 4096 + if is_longcat: + if num_cond_latents is not None and num_cond_latents > 0: + num_cond_latents_thw = num_cond_latents * (s // grid_sizes[0][0]) + x = x[:, num_cond_latents_thw:] q = self.norm_q(self.q(x).view(b, -1, n, d)) else: q = self.norm_q(self.q(x),num_chunks=2 if rope_func == "comfy_chunked" else 1).view(b, -1, n, d) @@ -702,7 +707,7 @@ class WanT2VCrossAttention(WanSelfAttention): if nag_context is not None and not is_uncond: x = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params) else: - if x.shape[-1] == 4096: + if is_longcat: k = self.norm_k(self.k(context).view(b, -1, n, d)) else: k = self.norm_k(self.k(context)).view(b, -1, n, d) @@ -715,6 +720,9 @@ class WanT2VCrossAttention(WanSelfAttention): k = rope_apply_c(k, cross_freqs, inner_c).to(q) x = attention(q, k, v, attention_mode=self.attention_mode).flatten(2) + if is_longcat: + if num_cond_latents is not None and num_cond_latents > 0: + x = torch.cat([torch.zeros((b, num_cond_latents_thw, x.shape[-1]), dtype=x.dtype, device=x.device), x], dim=1).contiguous() if lynx_x_ip is not None and self.ip_adapter is not None and ip_scale !=0: lynx_x_ip = self.ip_adapter(self, q, lynx_x_ip) @@ -1077,7 +1085,8 @@ class WanAttentionBlock(nn.Module): mtv_motion_tokens=None, mtv_motion_rotary_emb=None, mtv_strength=1.0, mtv_freqs=None, #mtv crafter humo_audio_input=None, humo_audio_scale=1.0, #humo audio 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, #ovi + 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 ): r""" Args: @@ -1222,6 +1231,19 @@ class WanAttentionBlock(nn.Module): full_k = torch.cat([k, k_ip], dim=1) full_v = torch.cat([v, v_ip], dim=1) y = self.self_attn.forward(q, full_k, full_v, seq_lens) + elif is_longcat: + if num_cond_latents is not None and num_cond_latents > 0: + num_cond_latents_thw = num_cond_latents * (N // grid_sizes[0][0]) + # process the condition tokens + q_cond = q[:, :num_cond_latents_thw].contiguous() + k_cond = k[:, :num_cond_latents_thw].contiguous() + v_cond = v[:, :num_cond_latents_thw].contiguous() + x_cond = self.self_attn.forward(q_cond, k_cond, v_cond, seq_lens) + # process the noise tokens + q_noise = q[:, num_cond_latents_thw:].contiguous() + x_noise = self.self_attn.forward(q_noise, k, v, seq_lens) + # merge x_cond and x_noise + y = torch.cat([x_cond, x_noise], dim=1).contiguous() else: y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale) @@ -1288,7 +1310,7 @@ class WanAttentionBlock(nn.Module): x = x + self.cross_attn(self.norm3(x), context, grid_sizes, clip_embed=clip_embed, audio_proj=audio_proj, audio_scale=audio_scale, num_latent_frames=num_latent_frames, nag_params=nag_params, nag_context=nag_context, is_uncond=is_uncond, 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) + 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) # MultiTalk if multitalk_audio_embedding is not None and not isinstance(self, VaceWanAttentionBlock): x_audio = self.audio_cross_attn(self.norm_x(x), encoder_hidden_states=multitalk_audio_embedding, @@ -2251,6 +2273,7 @@ class WanModel(torch.nn.Module): lynx_embeds=None, x_ovi=None, seq_len_ovi=None, ovi_negative_text_embeds=None, flashvsr_LQ_latent=None, flashvsr_strength=1.0, + num_cond_latents=None, ): r""" Forward pass through the diffusion model @@ -2854,6 +2877,7 @@ class WanModel(torch.nn.Module): lynx_x_ip=lynx_x_ip, lynx_ip_scale=lynx_ip_scale, lynx_ref_scale=lynx_ref_scale, + num_cond_latents=num_cond_latents ) if self.audio_model is not None: kwargs['e_ovi'] = e0_ovi.to(self.base_dtype)