works but slow

This commit is contained in:
kijai
2025-07-17 21:49:14 +03:00
parent 6ee7ad508e
commit 7c8020f8e8
3 changed files with 40 additions and 85 deletions
+1 -1
View File
@@ -2500,7 +2500,7 @@ class WanVideoSampler:
if transformer.attention_mode == "radial_sage_attention":
from .wanvideo.radial_attention.attn_mask import MaskMap
for block in transformer.blocks:
block.self_attn.mask_map = MaskMap(video_token_num=seq_len)
block.self_attn.mask_map = MaskMap(video_token_num=seq_len, num_frame=latent_video_length)
self.cache_state = [None, None]
if phantom_latents is not None:
+8 -13
View File
@@ -262,9 +262,9 @@ class WanSelfAttention(nn.Module):
#radial attention
self.layer_idx = layer_idx
self.dense_timestep = 0
self.dense_block = 0
self.decay_factor = 1
self.dense_timestep = 12
self.dense_block = 1
self.decay_factor = 0.2
self.sparse_type = "radial"
self.mask_map = None
@@ -284,7 +284,7 @@ class WanSelfAttention(nn.Module):
v = self.v(x).view(b, s, n, d)
return q, k, v
def forward(self, q, k, v, seq_lens, block_mask=None, timestep=0, numeral_timestep=0):
def forward(self, q, k, v, seq_lens, block_mask=None, current_step=0):
r"""
Args:
x(Tensor): Shape [B, L, num_heads, C / num_heads]
@@ -322,16 +322,15 @@ class WanSelfAttention(nn.Module):
q = rearrange(q, "b s h d -> (b s) h d")
k = rearrange(k, "b s h d -> (b s) h d")
v = rearrange(v, "b s h d -> (b s) h d")
if numeral_timestep < self.dense_timestep or self.layer_idx < self.dense_block or self.sparse_type == "dense":
if current_step < self.dense_timestep or self.layer_idx < self.dense_block or self.sparse_type == "dense":
x = RadialAttention(
query=q, key=k, value=v, mask_map=self.mask_map, sparsity_type="dense", block_size=128, decay_factor=self.decay_factor, model_type="wan", pre_defined_mask=None, use_sage_attention=True
query=q, key=k, value=v, mask_map=self.mask_map, sparsity_type="dense", block_size=128, decay_factor=self.decay_factor, model_type="wan", pre_defined_mask=None
)
else:
# apply radial attention
x = RadialAttention(
query=q, key=k, value=v, mask_map=self.mask_map, sparsity_type="radial", block_size=128, decay_factor=self.decay_factor, model_type="wan", pre_defined_mask=None, use_sage_attention=True
query=q, key=k, value=v, mask_map=self.mask_map, sparsity_type="radial", block_size=128, decay_factor=self.decay_factor, model_type="wan", pre_defined_mask=None
)
# transform back to (batch_size, num_heads, seq_len, head_dim)
x = rearrange(x, "(b s) h d -> b s h d", b=batch_size)
else:
@@ -667,8 +666,6 @@ class WanAttentionBlock(nn.Module):
context,
context_lens,
current_step,
timestep=0,
numeral_stimestep=0,
video_attention_split_steps=[],
clip_embed=None,
camera_embed=None,
@@ -737,7 +734,7 @@ class WanAttentionBlock(nn.Module):
elif ref_target_masks is not None:
y, x_ref_attn_map = self.self_attn.forward_multitalk(q, k, v, seq_lens, grid_sizes, ref_target_masks)
else:
y = self.self_attn.forward(q, k, v, seq_lens, block_mask=block_mask, timestep=timestep, numeral_stimestep=numeral_stimestep)
y = self.self_attn.forward(q, k, v, seq_lens, block_mask=block_mask, current_step=current_step)
# FETA
if enhance_enabled:
@@ -1739,8 +1736,6 @@ class WanModel(ModelMixin, ConfigMixin):
context=context,
context_lens=context_lens,
clip_embed=clip_embed,
timestep=t,
numera_timestep=10,
current_step=current_step,
video_attention_split_steps=self.video_attention_split_steps,
camera_embed=camera_embed,
+31 -71
View File
@@ -1,5 +1,5 @@
import torch
import flashinfer
#import flashinfer
import matplotlib.pyplot as plt
try:
from sparse_sageattn import sparse_sageattn
@@ -252,49 +252,9 @@ def SpargeSageAttnBackend(query, key, value, mask_map=None, video_mask=None, pre
# )
# return torch.cat([output_video, output_text], dim=0)
def FlashInferBackend(query, key, value, mask_map=None, pre_defined_mask=None, bsr_wrapper=None):
if pre_defined_mask is not None:
video_video_o, video_video_o_lse = bsr_wrapper.run(
query[:mask_map.video_token_num, :, :],
key[:mask_map.video_token_num, :, :],
value[:mask_map.video_token_num, :, :],
return_lse=True
)
# perform non-causal flashinfer on the text tokens
video_text_o, video_text_o_lse = flashinfer.single_prefill_with_kv_cache(
q=query[:mask_map.video_token_num, :, :],
k=key[mask_map.video_token_num:, :, :],
v=value[mask_map.video_token_num:, :, :],
causal=False,
return_lse=True,
custom_mask=pre_defined_mask[:mask_map.video_token_num, mask_map.video_token_num:]
)
# merge the two results
o_video, _ = flashinfer.merge_state(v_a=video_video_o, s_a=video_video_o_lse, v_b=video_text_o, s_b=video_text_o_lse)
o_text = flashinfer.single_prefill_with_kv_cache(
q=query[mask_map.video_token_num:, :, :],
k=key,
v=value,
causal=False,
return_lse=False,
custom_mask=pre_defined_mask[mask_map.video_token_num:, :]
)
return torch.cat([o_video, o_text], dim=0)
else:
o = bsr_wrapper.run(
query[:mask_map.video_token_num, :, :],
key[:mask_map.video_token_num, :, :],
value[:mask_map.video_token_num, :, :]
)
return o
def RadialAttention(query, key, value, mask_map=None, sparsity_type="radial", block_size=128, decay_factor=1, model_type=None, pre_defined_mask=None, use_sage_attention=False):
orig_seqlen, num_head, hidden_dim = query.shape
#orig_seqlen, num_head, hidden_dim = query.shape
if sparsity_type == "dense":
video_mask = torch.ones((mask_map.video_token_num // block_size, mask_map.video_token_num // block_size), device=query.device, dtype=torch.bool)
@@ -331,35 +291,35 @@ def RadialAttention(query, key, value, mask_map=None, sparsity_type="radial", bl
# elif backend == "sparse_sageattn":
return SpargeSageAttnBackend(query, key, value, mask_map, video_mask, pre_defined_mask, block_size=block_size)
if __name__ == "__main__":
query = torch.randn(1, 2, 4, 64).cuda()
# mask = torch.tensor([
# [True, False, True, False],
# [False, True, False, True],
# [True, False, False, True],
# [False, True, True, False]
# ], dtype=torch.bool)
# indices = get_indices_from_mask(mask, query)
# indptr = get_indptr_from_mask(mask, query)
# print("Indices: ", indices)
# print("Indptr: ", indptr)
video_token_num = 3840 * 30
num_frame = 30
token_per_frame = video_token_num / num_frame
padded_video_token_num = ((video_token_num + 1) // 128 + 1) * 128
print("padded: ", padded_video_token_num)
temporal_mask = gen_log_mask_shrinked(query, padded_video_token_num, video_token_num, num_frame, sparse_type="radial", decay_factor=1, model_type="hunyuan")
plt.figure(figsize=(10, 8), dpi=500)
# if __name__ == "__main__":
# query = torch.randn(1, 2, 4, 64).cuda()
# # mask = torch.tensor([
# # [True, False, True, False],
# # [False, True, False, True],
# # [True, False, False, True],
# # [False, True, True, False]
# # ], dtype=torch.bool)
# # indices = get_indices_from_mask(mask, query)
# # indptr = get_indptr_from_mask(mask, query)
# # print("Indices: ", indices)
# # print("Indptr: ", indptr)
# video_token_num = 3840 * 30
# num_frame = 30
# token_per_frame = video_token_num / num_frame
# padded_video_token_num = ((video_token_num + 1) // 128 + 1) * 128
# print("padded: ", padded_video_token_num)
# temporal_mask = gen_log_mask_shrinked(query, padded_video_token_num, video_token_num, num_frame, sparse_type="radial", decay_factor=1, model_type="hunyuan")
# plt.figure(figsize=(10, 8), dpi=500)
plt.imshow(temporal_mask.cpu().numpy()[:, :], cmap='hot')
plt.colorbar()
plt.title("Temporal Mask")
# plt.imshow(temporal_mask.cpu().numpy()[:, :], cmap='hot')
# plt.colorbar()
# plt.title("Temporal Mask")
plt.savefig("temporal_mask.png",
dpi=300,
bbox_inches='tight',
pad_inches=0.1)
# plt.savefig("temporal_mask.png",
# dpi=300,
# bbox_inches='tight',
# pad_inches=0.1)
plt.close()
# save the mask tensor
torch.save(temporal_mask, "temporal_mask.pt")
# plt.close()
# # save the mask tensor
# torch.save(temporal_mask, "temporal_mask.pt")