From 2fe483417849fe14ce422c3d616358c8031b5a89 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Mon, 29 Dec 2025 02:48:59 +0200 Subject: [PATCH] Adjust ultravico frame_tokens --- nodes_model_loading.py | 2 +- ultravico/sageattn/attn_qk_int8_per_block.py | 4 ++-- wanvideo/modules/attention.py | 8 ++++---- wanvideo/modules/model.py | 10 ++++++---- 4 files changed, 13 insertions(+), 11 deletions(-) diff --git a/nodes_model_loading.py b/nodes_model_loading.py index 01977f4..9b32b99 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -1016,7 +1016,7 @@ class WanVideoSetAttentionModeOverride: "required": { "model": ("WANVIDEOMODEL", ), "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"}), "verbose": ("BOOLEAN", {"default": False, "tooltip": "Print verbose info about attention mode override during generation"}), }, diff --git a/ultravico/sageattn/attn_qk_int8_per_block.py b/ultravico/sageattn/attn_qk_int8_per_block.py index 3d85856..645d5a4 100644 --- a/ultravico/sageattn/attn_qk_int8_per_block.py +++ b/ultravico/sageattn/attn_qk_int8_per_block.py @@ -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 - window_th = 1560 * 21 / 2 + window_th = frame_tokens * window_width / 2 dist2 = tl.abs(m - n).to(tl.int32) 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) - window3 = (m <= frame_tokens) & (n > 21*frame_tokens) + window3 = (m <= frame_tokens) & (n > window_width*frame_tokens) qk = tl.where(window3, -1e4, qk) diff --git a/wanvideo/modules/attention.py b/wanvideo/modules/attention.py index 09caac8..3b891d3 100644 --- a/wanvideo/modules/attention.py +++ b/wanvideo/modules/attention.py @@ -80,9 +80,9 @@ except: try: from ...ultravico.sageattn.core import sage_attention as sageattn_ultravico @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: - 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 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., 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: 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, @@ -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': 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() + return sageattn_func_ultravico([q, k, v], multi_factor=multi_factor, frame_tokens=frame_tokens).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 diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 6a31cb9..7db9377 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -467,7 +467,7 @@ class WanSelfAttention(nn.Module): v = (self.v(x) + self.v_loras(x)).view(b, s, n, d) 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""" Args: x(Tensor): Shape [B, L, num_heads, C / num_heads] @@ -477,12 +477,13 @@ class WanSelfAttention(nn.Module): """ attention_mode = self.attention_mode if attention_mode_override is not None: + print("Overriding attention mode to:", attention_mode_override) attention_mode = attention_mode_override 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, 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: 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 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 - attention_mode_override=None, + attention_mode_override=None, frame_tokens=None, ): r""" Args: @@ -1244,7 +1245,7 @@ class WanAttentionBlock(nn.Module): y = torch.cat([x_ref, x_cond, x_noise], dim=1).contiguous() else: 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 @@ -3041,6 +3042,7 @@ class WanModel(torch.nn.Module): camera_embed=camera_embed, audio_proj=audio_proj, num_latent_frames = F, + frame_tokens=x.shape[1] // F, original_seq_len=self.original_seq_len, enhance_enabled=enhance_enabled, audio_scale=audio_scale,