From 369043c5d84fd9be057c9f31e97f164bace0740d Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 10 Dec 2025 19:39:05 +0200 Subject: [PATCH] Add comfy core attention as option For edge cases that may benefit from the older attention types --- nodes_model_loading.py | 1 + wanvideo/modules/attention.py | 6 +++- wanvideo/modules/model.py | 67 ++++++++++++++++------------------- 3 files changed, 36 insertions(+), 38 deletions(-) diff --git a/nodes_model_loading.py b/nodes_model_loading.py index d6a1c04..12148d1 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -1011,6 +1011,7 @@ class WanVideoModelLoader: "radial_sage_attention", "sageattn_compiled", "sageattn_ultravico", + "comfy" ], {"default": "sdpa"}), "compile_args": ("WANCOMPILEARGS", ), "block_swap_args": ("BLOCKSWAPARGS", ), diff --git a/wanvideo/modules/attention.py b/wanvideo/modules/attention.py index c74a38a..09caac8 100644 --- a/wanvideo/modules/attention.py +++ b/wanvideo/modules/attention.py @@ -1,6 +1,8 @@ import torch from ...utils import log +from comfy.ldm.modules.attention import optimized_attention + def attention_func_error(*args, **kwargs): raise ImportError("Selected attention mode not available. Please ensure required packages are installed correctly.") @@ -92,7 +94,7 @@ except: def attention(q, k, v, q_lens=None, k_lens=None, max_seqlen_q=None, max_seqlen_k=None, dropout_p=0., softmax_scale=None, q_scale=None, causal=False, window_size=(-1, -1), deterministic=False, dtype=torch.bfloat16, - attention_mode='sdpa', attn_mask=None, multi_factor=0.9): + attention_mode='sdpa', attn_mask=None, multi_factor=0.9, heads=128): if "flash" in attention_mode: return flash_attention(q, k, v, q_lens=q_lens, k_lens=k_lens, dropout_p=dropout_p, softmax_scale=softmax_scale, q_scale=q_scale, causal=causal, window_size=window_size, deterministic=deterministic, dtype=dtype, version=2 if attention_mode == 'flash_attn_2' else 3, @@ -107,6 +109,8 @@ def attention(q, k, v, q_lens=None, k_lens=None, max_seqlen_q=None, max_seqlen_k return sageattn_func(q, k, v, tensor_layout="NHD").contiguous() elif attention_mode == 'sageattn_ultravico': return sageattn_func_ultravico([q, k, v], multi_factor=multi_factor).contiguous() + elif attention_mode == 'comfy': + return optimized_attention(q.transpose(1,2), k.transpose(1,2), v.transpose(1,2), heads=heads, skip_reshape=True) else: # sdpa if not (q.dtype == k.dtype == v.dtype): return torch.nn.functional.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2).to(q.dtype), v.transpose(1, 2).to(q.dtype), attn_mask=attn_mask).transpose(1, 2).contiguous() diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 866757d..baf1cf8 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -480,7 +480,7 @@ class WanSelfAttention(nn.Module): if self.ref_adapter is not None and lynx_ref_feature is not None: ref_x = self.ref_adapter(self, q, lynx_ref_feature) - x = attention(q, k, v, k_lens=seq_lens, attention_mode=attention_mode) + x = attention(q, k, v, k_lens=seq_lens, attention_mode=attention_mode, heads=self.num_heads) if self.ref_adapter is not None and lynx_ref_feature is not None: x = x.add(ref_x, alpha=lynx_ref_scale) @@ -499,35 +499,26 @@ class WanSelfAttention(nn.Module): # Concatenate main and IP keys/values for main attention full_k = torch.cat([k, k_ip], dim=1) full_v = torch.cat([v, v_ip], dim=1) - main_out = attention(q, full_k, full_v, k_lens=seq_lens, attention_mode=attention_mode) - - cond_out = attention(q_ip, k_ip, v_ip, k_lens=seq_lens, attention_mode=attention_mode) + main_out = attention(q, full_k, full_v, k_lens=seq_lens, attention_mode=attention_mode, heads=self.num_heads) + + cond_out = attention(q_ip, k_ip, v_ip, k_lens=seq_lens, attention_mode=attention_mode, heads=self.num_heads) x = torch.cat([main_out, cond_out], dim=1) return self.o(x.flatten(2)) - - + + def forward_radial(self, q, k, v, dense_step=False): if dense_step: x = RadialSpargeSageAttnDense(q, k, v, self.mask_map) else: x = RadialSpargeSageAttn(q, k, v, self.mask_map, decay_factor=self.decay_factor) return self.o(x.flatten(2)) - - + + def forward_multitalk(self, q, k, v, seq_lens, grid_sizes, ref_target_masks): - x = attention( - q, k, v, - k_lens=seq_lens, - attention_mode=self.attention_mode - ) - - # output - x = x.flatten(2) - x = self.o(x) - + x = attention(q, k, v, k_lens=seq_lens, attention_mode=self.attention_mode, heads=self.num_heads) + x = self.o(x.flatten(2)) x_ref_attn_map = get_attn_map_with_target(q.type_as(x), k.type_as(x), grid_sizes[0], ref_target_masks=ref_target_masks) - return x, x_ref_attn_map @@ -556,7 +547,8 @@ class WanSelfAttention(nn.Module): k[:, start_idx:end_idx, :, :], v[:, start_idx:end_idx, :, :], k_lens=seq_lens, - attention_mode=self.attention_mode + attention_mode=self.attention_mode, + heads=self.num_heads ) outputs.append(chunk_out) x = torch.cat(outputs, dim=1) @@ -577,10 +569,10 @@ class WanSelfAttention(nn.Module): k_negative = self.norm_k(self.k(context_negative).to(self.norm_k.weight.dtype)).view(b, -1, n, d).to(q.dtype) v_negative = self.v(context_negative).view(b, -1, n, d) - x_positive = attention(q, k_positive, v_positive, attention_mode=self.attention_mode) + x_positive = attention(q, k_positive, v_positive, attention_mode=self.attention_mode, heads=self.num_heads) x_positive = x_positive.flatten(2) - x_negative = attention(q, k_negative, v_negative, attention_mode=self.attention_mode) + x_negative = attention(q, k_negative, v_negative, attention_mode=self.attention_mode, heads=self.num_heads) x_negative = x_negative.flatten(2) nag_guidance = x_positive * nag_scale - x_negative * (nag_scale - 1) @@ -668,7 +660,7 @@ class WanT2VCrossAttention(WanSelfAttention): q = rope_apply_z(q, grid_sizes, cross_freqs, inner_t).to(q) k = rope_apply_c(k, cross_freqs, inner_c).to(q) - x = attention(q, k, v, attention_mode=self.attention_mode).flatten(2) + x = attention(q, k, v, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2) 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) @@ -680,12 +672,12 @@ class WanT2VCrossAttention(WanSelfAttention): audio_q = q.view(b * num_latent_frames, -1, n, d) ip_key = self.k_proj(audio_proj).view(b * num_latent_frames, -1, n, d) ip_value = self.v_proj(audio_proj).view(b * num_latent_frames, -1, n, d) - audio_x = attention(audio_q, ip_key, ip_value, attention_mode=self.attention_mode) + audio_x = attention(audio_q, ip_key, ip_value, attention_mode=self.attention_mode, heads=self.num_heads) audio_x = audio_x.view(b, q.size(1), n, d).flatten(2) elif len(audio_proj.shape) == 3: ip_key = self.k_proj(audio_proj).view(b, -1, n, d) ip_value = self.v_proj(audio_proj).view(b, -1, n, d) - audio_x = attention(q, ip_key, ip_value, attention_mode=self.attention_mode).flatten(2) + audio_x = attention(q, ip_key, ip_value, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2) x = x + audio_x * audio_scale # FantasyPortrait adapter attention @@ -696,13 +688,13 @@ class WanT2VCrossAttention(WanSelfAttention): 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) - adapter_x = attention(adapter_q, ip_key, ip_value, attention_mode=self.attention_mode) + adapter_x = attention(adapter_q, ip_key, ip_value, attention_mode=self.attention_mode, heads=self.num_heads) adapter_x = adapter_x.view(b, q_in.size(1), n, d) adapter_x = adapter_x.flatten(2) elif len(adapter_proj.shape) == 3: ip_key = self.ip_adapter_single_stream_k_proj(adapter_proj).view(b, -1, n, d) ip_value = self.ip_adapter_single_stream_v_proj(adapter_proj).view(b, -1, n, d) - adapter_x = attention(q_in, ip_key, ip_value, attention_mode=self.attention_mode) + adapter_x = attention(q_in, ip_key, ip_value, attention_mode=self.attention_mode, heads=self.num_heads) adapter_x = adapter_x.flatten(2) x[:, :orig_seq_len] = x[:, :orig_seq_len] + adapter_x * ip_scale @@ -714,7 +706,7 @@ class WanT2VCrossAttention(WanSelfAttention): q = rope_apply(q, grid_sizes, kwargs["src_freqs"]) k_target = rope_apply(k_target, kwargs["target_grid_sizes"], kwargs["target_freqs"]) - target_x = attention(q, k_target, v_target, k_lens=kwargs["target_seq_lens"]).flatten(2) + target_x = attention(q, k_target, v_target, k_lens=kwargs["target_seq_lens"], heads=self.num_heads).flatten(2) x = x.add(target_x) @@ -750,13 +742,13 @@ class WanI2VCrossAttention(WanSelfAttention): # 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).flatten(2) + x_text = attention(q, k, v, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2) #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).flatten(2) + img_x = attention(q, k_img, v_img, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2) x = x_text + img_x else: x = x_text @@ -768,12 +760,12 @@ class WanI2VCrossAttention(WanSelfAttention): ip_key = self.k_proj(audio_proj).view(b * num_latent_frames, -1, n, d) ip_value = self.v_proj(audio_proj).view(b * num_latent_frames, -1, n, d) - audio_x = attention(audio_q, ip_key, ip_value, attention_mode=self.attention_mode) + audio_x = attention(audio_q, ip_key, ip_value, attention_mode=self.attention_mode, heads=self.num_heads) audio_x = audio_x.view(b, q.size(1), n, d).flatten(2) elif len(audio_proj.shape) == 3: ip_key = self.k_proj(audio_proj).view(b, -1, n, d) ip_value = self.v_proj(audio_proj).view(b, -1, n, d) - audio_x = attention(q, ip_key, ip_value, attention_mode=self.attention_mode).flatten(2) + audio_x = attention(q, ip_key, ip_value, attention_mode=self.attention_mode, heads=self.num_heads).flatten(2) x = x + audio_x * audio_scale # FantasyPortrait adapter attention @@ -783,13 +775,13 @@ class WanI2VCrossAttention(WanSelfAttention): 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) - adapter_x = attention(adapter_q, ip_key, ip_value, attention_mode=self.attention_mode) + adapter_x = attention(adapter_q, ip_key, ip_value, attention_mode=self.attention_mode, heads=self.num_heads) adapter_x = adapter_x.view(b, q.size(1), n, d) adapter_x = adapter_x.flatten(2) elif len(adapter_proj.shape) == 3: ip_key = self.ip_adapter_single_stream_k_proj(adapter_proj).view(b, -1, n, d) ip_value = self.ip_adapter_single_stream_v_proj(adapter_proj).view(b, -1, n, d) - adapter_x = attention(q, ip_key, ip_value, attention_mode=self.attention_mode) + 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 @@ -816,7 +808,7 @@ class WanHuMoCrossAttention(WanSelfAttention): k = k.reshape(-1, 16, n, d) v = v.reshape(-1, 16, n, d) - x_text = attention(q, k, v, attention_mode=self.attention_mode) + x_text = attention(q, k, v, attention_mode=self.attention_mode, heads=self.num_heads) x_text = x_text.view(b, -1, n, d).flatten(2) x = x_text @@ -855,7 +847,8 @@ class MTVCrafterMotionAttention(WanSelfAttention): x = attention( q=rope_apply(q, grid_sizes, freqs), k=apply_rotary_emb(k, pe).transpose(1, 2), - v=v + v=v, + heads=self.num_heads, ) return self.o(x.flatten(2)) @@ -1095,7 +1088,7 @@ class WanAttentionBlock(nn.Module): v_ref = self.ref_attn_v_img(x_onetoall_ref).view(b, s, n, d) - onetoall_ref = attention(q_ref, k_ref, v_ref, k_lens=seq_lens, attention_mode=self.attention_mode) + onetoall_ref = attention(q_ref, k_ref, v_ref, k_lens=seq_lens, attention_mode=self.attention_mode, heads=self.num_heads) del q_ref, k_ref, v_ref #RoPE and QKV computation