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