I2V cross_attn fixes
This commit is contained in:
+11
-11
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user