Separate attention for input images like in original

This commit is contained in:
kijai
2025-10-26 19:14:47 +02:00
parent 43acf83adb
commit d504c96174
2 changed files with 30 additions and 5 deletions
+1
View File
@@ -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
+29 -5
View File
@@ -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)