diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index ce101d8..60616e8 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -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):