works but slow
This commit is contained in:
@@ -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:
|
||||
|
||||
@@ -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,
|
||||
|
||||
@@ -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")
|
||||
Reference in New Issue
Block a user