From 7c8020f8e8051ae63a06457aec8a5f31165e452f Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 17 Jul 2025 21:49:14 +0300 Subject: [PATCH] works but slow --- nodes.py | 2 +- wanvideo/modules/model.py | 21 ++--- wanvideo/radial_attention/attn_mask.py | 102 ++++++++----------------- 3 files changed, 40 insertions(+), 85 deletions(-) diff --git a/nodes.py b/nodes.py index fce0ac9..fa55a0e 100644 --- a/nodes.py +++ b/nodes.py @@ -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: diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 902b753..5c74303 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -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, diff --git a/wanvideo/radial_attention/attn_mask.py b/wanvideo/radial_attention/attn_mask.py index e3c38ff..e47c5ae 100644 --- a/wanvideo/radial_attention/attn_mask.py +++ b/wanvideo/radial_attention/attn_mask.py @@ -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") \ No newline at end of file +# plt.close() +# # save the mask tensor +# torch.save(temporal_mask, "temporal_mask.pt") \ No newline at end of file