Add comfy core attention as option

For edge cases that may benefit from the older attention types
This commit is contained in:
kijai
2025-12-10 19:39:05 +02:00
parent eb6dd96049
commit 369043c5d8
3 changed files with 36 additions and 38 deletions
+1
View File
@@ -1011,6 +1011,7 @@ class WanVideoModelLoader:
"radial_sage_attention",
"sageattn_compiled",
"sageattn_ultravico",
"comfy"
], {"default": "sdpa"}),
"compile_args": ("WANCOMPILEARGS", ),
"block_swap_args": ("BLOCKSWAPARGS", ),
+5 -1
View File
@@ -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()
+30 -37
View File
@@ -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