Separate attention for input images like in original
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user