Add comfy core attention as option
For edge cases that may benefit from the older attention types
This commit is contained in:
@@ -1011,6 +1011,7 @@ class WanVideoModelLoader:
|
||||
"radial_sage_attention",
|
||||
"sageattn_compiled",
|
||||
"sageattn_ultravico",
|
||||
"comfy"
|
||||
], {"default": "sdpa"}),
|
||||
"compile_args": ("WANCOMPILEARGS", ),
|
||||
"block_swap_args": ("BLOCKSWAPARGS", ),
|
||||
|
||||
@@ -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
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user