Adjust ultravico frame_tokens

This commit is contained in:
kijai
2025-12-29 02:48:59 +02:00
parent 486564060f
commit 2fe4834178
4 changed files with 13 additions and 11 deletions
+1 -1
View File
@@ -1016,7 +1016,7 @@ class WanVideoSetAttentionModeOverride:
"required": { "required": {
"model": ("WANVIDEOMODEL", ), "model": ("WANVIDEOMODEL", ),
"attention_mode": (attention_modes, {"default": "sdpa"}), "attention_mode": (attention_modes, {"default": "sdpa"}),
"start_step": ("INT", {"default": 1, "min": 1, "max": 10000, "step": 1, "tooltip": "Step to start applying the attention mode override"}), "start_step": ("INT", {"default": 0, "min": 0, "max": 10000, "step": 1, "tooltip": "Step to start applying the attention mode override"}),
"end_step": ("INT", {"default": 10000, "min": 1, "max": 10000, "step": 1, "tooltip": "Step to end applying the attention mode override"}), "end_step": ("INT", {"default": 10000, "min": 1, "max": 10000, "step": 1, "tooltip": "Step to end applying the attention mode override"}),
"verbose": ("BOOLEAN", {"default": False, "tooltip": "Print verbose info about attention mode override during generation"}), "verbose": ("BOOLEAN", {"default": False, "tooltip": "Print verbose info about attention mode override during generation"}),
}, },
+2 -2
View File
@@ -38,7 +38,7 @@ def _attn_fwd_inner(acc, l_i, m_i, q, q_scale, kv_len, current_flag,
qk = tl.dot(q, k).to(tl.float32) * q_scale * k_scale qk = tl.dot(q, k).to(tl.float32) * q_scale * k_scale
window_th = 1560 * 21 / 2 window_th = frame_tokens * window_width / 2
dist2 = tl.abs(m - n).to(tl.int32) dist2 = tl.abs(m - n).to(tl.int32)
dist_mask = dist2 <= window_th dist_mask = dist2 <= window_th
@@ -46,7 +46,7 @@ def _attn_fwd_inner(acc, l_i, m_i, q, q_scale, kv_len, current_flag,
qk = tl.where(dist_mask | negative_mask, qk, qk*multi_factor) qk = tl.where(dist_mask | negative_mask, qk, qk*multi_factor)
window3 = (m <= frame_tokens) & (n > 21*frame_tokens) window3 = (m <= frame_tokens) & (n > window_width*frame_tokens)
qk = tl.where(window3, -1e4, qk) qk = tl.where(window3, -1e4, qk)
+4 -4
View File
@@ -80,9 +80,9 @@ except:
try: try:
from ...ultravico.sageattn.core import sage_attention as sageattn_ultravico from ...ultravico.sageattn.core import sage_attention as sageattn_ultravico
@torch.library.custom_op("wanvideo::sageattn_ultravico", mutates_args=()) @torch.library.custom_op("wanvideo::sageattn_ultravico", mutates_args=())
def sageattn_func_ultravico(qkv: List[torch.Tensor], attn_mask: torch.Tensor | None = None, dropout_p: float = 0.0, is_causal: bool = False, multi_factor: float = 0.9 def sageattn_func_ultravico(qkv: List[torch.Tensor], attn_mask: torch.Tensor | None = None, dropout_p: float = 0.0, is_causal: bool = False, multi_factor: float = 0.9, frame_tokens: int = 1536
) -> torch.Tensor: ) -> torch.Tensor:
return sageattn_ultravico(qkv, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, multi_factor=multi_factor) return sageattn_ultravico(qkv, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, multi_factor=multi_factor, frame_tokens=frame_tokens)
@sageattn_func_ultravico.register_fake @sageattn_func_ultravico.register_fake
def _(qkv, attn_mask=None, dropout_p=0.0, is_causal=False, multi_factor=0.9): def _(qkv, attn_mask=None, dropout_p=0.0, is_causal=False, multi_factor=0.9):
@@ -94,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., 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, 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, heads=128): attention_mode='sdpa', attn_mask=None, multi_factor=0.9, frame_tokens=1536, heads=128):
if "flash" in attention_mode: 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, 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, q_scale=q_scale, causal=causal, window_size=window_size, deterministic=deterministic, dtype=dtype, version=2 if attention_mode == 'flash_attn_2' else 3,
@@ -108,7 +108,7 @@ def attention(q, k, v, q_lens=None, k_lens=None, max_seqlen_q=None, max_seqlen_k
elif attention_mode == 'sageattn': elif attention_mode == 'sageattn':
return sageattn_func(q, k, v, tensor_layout="NHD").contiguous() return sageattn_func(q, k, v, tensor_layout="NHD").contiguous()
elif attention_mode == 'sageattn_ultravico': elif attention_mode == 'sageattn_ultravico':
return sageattn_func_ultravico([q, k, v], multi_factor=multi_factor).contiguous() return sageattn_func_ultravico([q, k, v], multi_factor=multi_factor, frame_tokens=frame_tokens).contiguous()
elif attention_mode == 'comfy': elif attention_mode == 'comfy':
return optimized_attention(q.transpose(1,2), k.transpose(1,2), v.transpose(1,2), heads=heads, skip_reshape=True) return optimized_attention(q.transpose(1,2), k.transpose(1,2), v.transpose(1,2), heads=heads, skip_reshape=True)
else: # sdpa else: # sdpa
+6 -4
View File
@@ -467,7 +467,7 @@ class WanSelfAttention(nn.Module):
v = (self.v(x) + self.v_loras(x)).view(b, s, n, d) v = (self.v(x) + self.v_loras(x)).view(b, s, n, d)
return q, k, v return q, k, v
def forward(self, q, k, v, seq_lens, lynx_ref_feature=None, lynx_ref_scale=1.0, attention_mode_override=None, onetoall_ref=None, onetoall_ref_scale=1.0): def forward(self, q, k, v, seq_lens, lynx_ref_feature=None, lynx_ref_scale=1.0, attention_mode_override=None, onetoall_ref=None, onetoall_ref_scale=1.0, frame_tokens=1536):
r""" r"""
Args: Args:
x(Tensor): Shape [B, L, num_heads, C / num_heads] x(Tensor): Shape [B, L, num_heads, C / num_heads]
@@ -477,12 +477,13 @@ class WanSelfAttention(nn.Module):
""" """
attention_mode = self.attention_mode attention_mode = self.attention_mode
if attention_mode_override is not None: if attention_mode_override is not None:
print("Overriding attention mode to:", attention_mode_override)
attention_mode = attention_mode_override attention_mode = attention_mode_override
if self.ref_adapter is not None and lynx_ref_feature is not None: if self.ref_adapter is not None and lynx_ref_feature is not None:
ref_x = self.ref_adapter(self, q, lynx_ref_feature) ref_x = self.ref_adapter(self, q, lynx_ref_feature)
x = attention(q, k, v, k_lens=seq_lens, attention_mode=attention_mode, heads=self.num_heads) x = attention(q, k, v, k_lens=seq_lens, attention_mode=attention_mode, heads=self.num_heads, frame_tokens=frame_tokens)
if self.ref_adapter is not None and lynx_ref_feature is not None: if self.ref_adapter is not None and lynx_ref_feature is not None:
x = x.add(ref_x, alpha=lynx_ref_scale) x = x.add(ref_x, alpha=lynx_ref_scale)
@@ -1006,7 +1007,7 @@ class WanAttentionBlock(nn.Module):
longcat_num_cond_latents=0, longcat_avatar_options=None, #longcat image cond amount longcat_num_cond_latents=0, longcat_avatar_options=None, #longcat image cond amount
x_onetoall_ref=None, onetoall_freqs=None, onetoall_ref=None, onetoall_ref_scale=1.0, #one-to-all x_onetoall_ref=None, onetoall_freqs=None, onetoall_ref=None, onetoall_ref_scale=1.0, #one-to-all
e_tr=None, tr_num=0, tr_start=0, #token replacement e_tr=None, tr_num=0, tr_start=0, #token replacement
attention_mode_override=None, attention_mode_override=None, frame_tokens=None,
): ):
r""" r"""
Args: Args:
@@ -1244,7 +1245,7 @@ class WanAttentionBlock(nn.Module):
y = torch.cat([x_ref, x_cond, x_noise], dim=1).contiguous() y = torch.cat([x_ref, x_cond, x_noise], dim=1).contiguous()
else: else:
y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale, y = self.self_attn.forward(q, k, v, seq_lens, lynx_ref_feature=lynx_ref_feature, lynx_ref_scale=lynx_ref_scale,
onetoall_ref=onetoall_ref, onetoall_ref_scale=onetoall_ref_scale, attention_mode_override=attention_mode_override) onetoall_ref=onetoall_ref, onetoall_ref_scale=onetoall_ref_scale, attention_mode_override=attention_mode_override, frame_tokens=frame_tokens)
del q, k, v del q, k, v
@@ -3041,6 +3042,7 @@ class WanModel(torch.nn.Module):
camera_embed=camera_embed, camera_embed=camera_embed,
audio_proj=audio_proj, audio_proj=audio_proj,
num_latent_frames = F, num_latent_frames = F,
frame_tokens=x.shape[1] // F,
original_seq_len=self.original_seq_len, original_seq_len=self.original_seq_len,
enhance_enabled=enhance_enabled, enhance_enabled=enhance_enabled,
audio_scale=audio_scale, audio_scale=audio_scale,