I2V cross_attn fixes

This commit is contained in:
kijai
2026-02-16 15:31:20 +02:00
parent 3d7b49e2df
commit 5d36631795
+11 -11
View File
@@ -647,7 +647,7 @@ 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, 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, longcat_num_cond_latents=None, **kwargs):
adapter_proj=None, ip_scale=1.0, orig_seq_len=None, lynx_x_ip=None, lynx_ip_scale=1.0, longcat_num_cond_latents=None, **kwargs):
b, n, d = x.size(0), self.num_heads, self.head_dim
s = x.size(1)
# compute query
@@ -702,7 +702,7 @@ class WanT2VCrossAttention(WanSelfAttention):
# FantasyPortrait adapter attention
if adapter_proj is not None:
if len(adapter_proj.shape) == 4:
q_in = q[:, :orig_seq_len]
q_in = q[:, :orig_seq_len]
adapter_q = q_in.view(b * num_latent_frames, -1, n, d)
ip_key = self.ip_adapter_single_stream_k_proj(adapter_proj).view(b * num_latent_frames, -1, n, d)
ip_value = self.ip_adapter_single_stream_v_proj(adapter_proj).view(b * num_latent_frames, -1, n, d)
@@ -745,7 +745,7 @@ class WanI2VCrossAttention(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, rope_func="comfy",
adapter_proj=None, adapter_attn_mask=None, ip_scale=1.0, orig_seq_len=None, **kwargs):
adapter_proj=None, ip_scale=1.0, orig_seq_len=None, **kwargs):
r"""
Args:
x(Tensor): Shape [B, L1, C]
@@ -757,22 +757,22 @@ class WanI2VCrossAttention(WanSelfAttention):
q = self.norm_q(self.q(x).to(self.norm_q.weight.dtype),num_chunks=2 if rope_func == "comfy_chunked" else 1).view(b, -1, n, d).to(x.dtype)
if nag_context is not None:
x_text = self.normalized_attention_guidance(b, n, d, q, context, nag_context, nag_params)
x_positive, x_negative = self.nag_attention(b, n, d, q, context, nag_context)
x = self.normalized_attention_guidance(x_positive, x_negative, nag_params)
del x_positive, x_negative
else:
# text attention
k = self.norm_k(self.k(context).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(x.dtype)
v = self.v(context).view(b, -1, n, d)
x_text = attention(q, k, v, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2)
x = attention(q, k, v, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2)
del k, v
#img attention
if clip_embed is not None:
k_img = self.norm_k_img(self.k_img(clip_embed).to(self.norm_k_img.weight.dtype)).view(b, -1, n, d).to(x.dtype)
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, heads=self.num_heads).flatten(2)
x_text.add_(img_x)
x = x_text
else:
x = x_text
x.add_(attention(q, k_img, v_img, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2))
del k_img, v_img
# FantasyTalking audio attention
if audio_proj is not None:
@@ -805,7 +805,7 @@ class WanI2VCrossAttention(WanSelfAttention):
adapter_x = attention(q, ip_key, ip_value, attention_mode=self.attention_mode, heads=self.num_heads)
adapter_x = adapter_x.flatten(2)
x = x + adapter_x * ip_scale
del q
return self.o(x)
class WanHuMoCrossAttention(WanSelfAttention):