Adjust ultravico frame_tokens
This commit is contained in:
@@ -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"}),
|
||||
},
|
||||
|
||||
@@ -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)
|
||||
|
||||
|
||||
|
||||
@@ -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
|
||||
|
||||
@@ -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,
|
||||
|
||||
Reference in New Issue
Block a user