init (not working)
This commit is contained in:
@@ -2339,10 +2339,10 @@ class WanVideoSampler:
|
||||
if mask is not None:
|
||||
if mask.shape[2] != noise.shape[1]:
|
||||
mask = torch.cat([torch.zeros(1, noise.shape[0], noise.shape[1] - mask.shape[2], noise.shape[2], noise.shape[3]), mask], dim=2)
|
||||
|
||||
if (extra_latents := image_embeds.get("extra_latents"), None) is not None:
|
||||
|
||||
if (extra_latents := image_embeds.get("extra_latents", None)) is not None:
|
||||
encoded_image_latents = extra_latents["samples"].squeeze(0).to(noise)
|
||||
if (empty_latent_indices := extra_latents.get("empty_latent_indices"), None) is not None and len(empty_latent_indices) > 0:
|
||||
if (empty_latent_indices := extra_latents.get("empty_latent_indices", None)) is not None and len(empty_latent_indices) > 0:
|
||||
noise_out = encoded_image_latents.clone()
|
||||
for idx in empty_latent_indices:
|
||||
#print(f"Adding noise to Empty latent index: {idx}")
|
||||
@@ -2497,6 +2497,11 @@ class WanVideoSampler:
|
||||
else:
|
||||
transformer.slg_blocks = None
|
||||
|
||||
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)
|
||||
|
||||
self.cache_state = [None, None]
|
||||
if phantom_latents is not None:
|
||||
log.info(f"Phantom latents shape: {phantom_latents.shape}")
|
||||
|
||||
@@ -472,8 +472,7 @@ class WanVideoModelLoader:
|
||||
"flash_attn_3",
|
||||
"sageattn",
|
||||
"flex_attention",
|
||||
#"spargeattn", needs tuning
|
||||
#"spargeattn_tune",
|
||||
"radial_sage_attention",
|
||||
], {"default": "sdpa"}),
|
||||
"compile_args": ("WANCOMPILEARGS", ),
|
||||
"block_swap_args": ("BLOCKSWAPARGS", ),
|
||||
|
||||
@@ -249,7 +249,8 @@ class WanSelfAttention(nn.Module):
|
||||
num_heads,
|
||||
qk_norm=True,
|
||||
eps=1e-6,
|
||||
attention_mode='sdpa'):
|
||||
attention_mode='sdpa',
|
||||
layer_idx=0):
|
||||
assert out_features % num_heads == 0
|
||||
super().__init__()
|
||||
self.dim = out_features
|
||||
@@ -259,6 +260,15 @@ class WanSelfAttention(nn.Module):
|
||||
self.eps = eps
|
||||
self.attention_mode = attention_mode
|
||||
|
||||
#radial attention
|
||||
self.layer_idx = layer_idx
|
||||
self.dense_timestep = 0
|
||||
self.dense_block = 0
|
||||
self.decay_factor = 1
|
||||
self.sparse_type = "radial"
|
||||
self.mask_map = None
|
||||
|
||||
|
||||
# layers
|
||||
self.q = nn.Linear(in_features, out_features)
|
||||
self.k = nn.Linear(in_features, out_features)
|
||||
@@ -274,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):
|
||||
def forward(self, q, k, v, seq_lens, block_mask=None, timestep=0, numeral_timestep=0):
|
||||
r"""
|
||||
Args:
|
||||
x(Tensor): Shape [B, L, num_heads, C / num_heads]
|
||||
@@ -306,6 +316,24 @@ class WanSelfAttention(nn.Module):
|
||||
value=padded_v.transpose(2, 1),
|
||||
block_mask=block_mask
|
||||
)[:, :, :-padded_length].transpose(2, 1)
|
||||
elif self.attention_mode == 'radial_sage_attention':
|
||||
from ..radial_attention.attn_mask import RadialAttention
|
||||
batch_size = q.shape[0]
|
||||
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":
|
||||
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
|
||||
)
|
||||
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
|
||||
)
|
||||
# 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:
|
||||
x = attention(
|
||||
q, k, v,
|
||||
@@ -560,7 +588,8 @@ class WanAttentionBlock(nn.Module):
|
||||
cross_attn_norm=False,
|
||||
eps=1e-6,
|
||||
attention_mode='sdpa',
|
||||
rope_func="comfy"
|
||||
rope_func="comfy",
|
||||
block_idx=0
|
||||
):
|
||||
super().__init__()
|
||||
self.dim = out_features
|
||||
@@ -571,11 +600,12 @@ class WanAttentionBlock(nn.Module):
|
||||
self.eps = eps
|
||||
self.attention_mode = attention_mode
|
||||
self.rope_func = rope_func
|
||||
self.block_idx = block_idx
|
||||
|
||||
# layers
|
||||
self.norm1 = WanLayerNorm(out_features, eps)
|
||||
self.self_attn = WanSelfAttention(in_features, out_features, num_heads, qk_norm,
|
||||
eps, self.attention_mode)
|
||||
eps, self.attention_mode, self.block_idx)
|
||||
if cross_attn_type != "no_cross_attn":
|
||||
self.norm3 = WanLayerNorm(
|
||||
out_features, eps,
|
||||
@@ -637,6 +667,8 @@ class WanAttentionBlock(nn.Module):
|
||||
context,
|
||||
context_lens,
|
||||
current_step,
|
||||
timestep=0,
|
||||
numeral_stimestep=0,
|
||||
video_attention_split_steps=[],
|
||||
clip_embed=None,
|
||||
camera_embed=None,
|
||||
@@ -705,7 +737,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)
|
||||
y = self.self_attn.forward(q, k, v, seq_lens, block_mask=block_mask, timestep=timestep, numeral_stimestep=numeral_stimestep)
|
||||
|
||||
# FETA
|
||||
if enhance_enabled:
|
||||
@@ -1121,8 +1153,8 @@ class WanModel(ModelMixin, ConfigMixin):
|
||||
self.blocks = nn.ModuleList([
|
||||
WanAttentionBlock(cross_attn_type, self.in_features, self.out_features, ffn_dim, ffn2_dim, num_heads,
|
||||
qk_norm, cross_attn_norm, eps,
|
||||
attention_mode=self.attention_mode, rope_func=self.rope_func)
|
||||
for _ in range(num_layers)
|
||||
attention_mode=self.attention_mode, rope_func=self.rope_func, block_idx=i)
|
||||
for i in range(num_layers)
|
||||
])
|
||||
|
||||
# head
|
||||
@@ -1707,6 +1739,8 @@ 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,
|
||||
|
||||
@@ -0,0 +1,365 @@
|
||||
import torch
|
||||
import flashinfer
|
||||
import matplotlib.pyplot as plt
|
||||
try:
|
||||
from sparse_sageattn import sparse_sageattn
|
||||
except:
|
||||
sparse_sageattn = None
|
||||
from einops import rearrange, repeat
|
||||
|
||||
def sparge_mask_convert(mask: torch.Tensor, block_size: int = 128) -> torch.Tensor:
|
||||
assert block_size in [128, 64], "Radial Attention only supports block size of 128 or 64"
|
||||
assert mask.shape[0] == mask.shape[1], "Input mask must be square."
|
||||
|
||||
if block_size == 128:
|
||||
new_mask = torch.repeat_interleave(mask, 2, dim=1)
|
||||
|
||||
elif block_size == 64:
|
||||
num_row, num_col = mask.shape
|
||||
reshaped_mask = mask.view(num_row // 2, 2, num_col)
|
||||
new_mask = torch.max(reshaped_mask, dim=1).values
|
||||
|
||||
return new_mask
|
||||
|
||||
def get_indptr_from_mask(mask, query):
|
||||
# query shows the device of the indptr
|
||||
# indptr (torch.Tensor) - the block index pointer of the block-sparse matrix on row dimension,
|
||||
# shape `(MB + 1,)`, where `MB` is the number of blocks in the row dimension.
|
||||
# The first element is always 0, and the last element is the number of blocks in the row dimension.
|
||||
# The rest of the elements are the number of blocks in each row.
|
||||
# the mask is already a block sparse mask
|
||||
indptr = torch.zeros(mask.shape[0] + 1, device=query.device, dtype=torch.int32)
|
||||
indptr[0] = 0
|
||||
row_counts = mask.sum(dim=1).flatten() # Ensure 1D output [num_blocks_row]
|
||||
indptr[1:] = torch.cumsum(row_counts, dim=0)
|
||||
return indptr
|
||||
|
||||
def get_indices_from_mask(mask, query):
|
||||
# indices (torch.Tensor) - the block indices of the block-sparse matrix on column dimension,
|
||||
# shape `(nnz,),` where `nnz` is the number of non-zero blocks.
|
||||
# The elements in `indices` array should be less than `NB`: the number of blocks in the column dimension.
|
||||
nonzero_indices = torch.nonzero(mask)
|
||||
indices = nonzero_indices[:, 1].to(dtype=torch.int32, device=query.device)
|
||||
return indices
|
||||
|
||||
def shrinkMaskStrict(mask, block_size=128):
|
||||
seqlen = mask.shape[0]
|
||||
block_num = seqlen // block_size
|
||||
mask = mask[:block_num * block_size, :block_num * block_size].view(block_num, block_size, block_num, block_size)
|
||||
col_densities = mask.sum(dim = 1) / block_size
|
||||
# we want the minimum non-zero column density in the block
|
||||
non_zero_densities = col_densities > 0
|
||||
high_density_cols = col_densities > 1/3
|
||||
frac_high_density_cols = high_density_cols.sum(dim=-1) / (non_zero_densities.sum(dim=-1) + 1e-9)
|
||||
block_mask = frac_high_density_cols > 0.6
|
||||
block_mask[0:0] = True
|
||||
block_mask[-1:-1] = True
|
||||
return block_mask
|
||||
|
||||
def pad_qkv(input_tensor, block_size=128):
|
||||
"""
|
||||
Pad the input tensor to be a multiple of the block size.
|
||||
input shape: (seqlen, num_heads, hidden_dim)
|
||||
"""
|
||||
seqlen, num_heads, hidden_dim = input_tensor.shape
|
||||
# Calculate the necessary padding
|
||||
padding_length = (block_size - (seqlen % block_size)) % block_size
|
||||
# Create a padded tensor with zeros
|
||||
padded_tensor = torch.zeros((seqlen + padding_length, num_heads, hidden_dim), device=input_tensor.device, dtype=input_tensor.dtype)
|
||||
# Copy the original tensor into the padded tensor
|
||||
padded_tensor[:seqlen, :, :] = input_tensor
|
||||
|
||||
return padded_tensor
|
||||
|
||||
def get_diagonal_split_mask(i, j, token_per_frame, sparse_type, query):
|
||||
assert(sparse_type in ["radial"])
|
||||
dist = abs(i - j)
|
||||
group = dist.bit_length()
|
||||
threshold = 128 # hardcoded threshold for now, which is equal to block-size
|
||||
decay_length = 2 ** token_per_frame.bit_length() / 2 ** group
|
||||
if decay_length >= threshold:
|
||||
return torch.ones((token_per_frame, token_per_frame), device=query.device, dtype=torch.bool)
|
||||
|
||||
split_factor = int(threshold / decay_length)
|
||||
modular = dist % split_factor
|
||||
if modular == 0:
|
||||
return torch.ones((token_per_frame, token_per_frame), device=query.device, dtype=torch.bool)
|
||||
else:
|
||||
return torch.zeros((token_per_frame, token_per_frame), device=query.device, dtype=torch.bool)
|
||||
|
||||
def get_window_width(i, j, token_per_frame, sparse_type, num_frame, decay_factor=1, block_size=128, model_type=None):
|
||||
assert(sparse_type in ["radial"])
|
||||
dist = abs(i - j)
|
||||
if model_type == "wan":
|
||||
if dist < 1:
|
||||
return token_per_frame
|
||||
if dist == 1:
|
||||
return token_per_frame // 2
|
||||
elif model_type == "hunyuan":
|
||||
if dist <= 1:
|
||||
return token_per_frame
|
||||
else:
|
||||
raise ValueError(f"Unknown model type: {model_type}")
|
||||
group = dist.bit_length()
|
||||
decay_length = 2 ** token_per_frame.bit_length() / 2 ** group * decay_factor
|
||||
threshold = block_size
|
||||
if decay_length >= threshold:
|
||||
return decay_length
|
||||
else:
|
||||
return threshold
|
||||
|
||||
def gen_log_mask_shrinked(query, s, video_token_num, num_frame, block_size=128, sparse_type="log", decay_factor=0.5, model_type=None):
|
||||
"""
|
||||
A more memory friendly version, we generate the attention mask of each frame pair at a time,
|
||||
shrinks it, and stores it into the final result
|
||||
"""
|
||||
final_log_mask = torch.zeros((s // block_size, s // block_size), device=query.device, dtype=torch.bool)
|
||||
token_per_frame = video_token_num // num_frame
|
||||
video_text_border = video_token_num // block_size
|
||||
|
||||
col_indices = torch.arange(0, token_per_frame, device=query.device).view(1, -1)
|
||||
row_indices = torch.arange(0, token_per_frame, device=query.device).view(-1, 1)
|
||||
final_log_mask[video_text_border:] = True
|
||||
final_log_mask[:, video_text_border:] = True
|
||||
for i in range(num_frame):
|
||||
for j in range(num_frame):
|
||||
local_mask = torch.zeros((token_per_frame, token_per_frame), device=query.device, dtype=torch.bool)
|
||||
if j == 0 and model_type == "wan": # this is attention sink
|
||||
local_mask = torch.ones((token_per_frame, token_per_frame), device=query.device, dtype=torch.bool)
|
||||
else:
|
||||
window_width = get_window_width(i, j, token_per_frame, sparse_type, num_frame, decay_factor=decay_factor, block_size=block_size, model_type=model_type)
|
||||
local_mask = torch.abs(col_indices - row_indices) <= window_width
|
||||
split_mask = get_diagonal_split_mask(i, j, token_per_frame, sparse_type, query)
|
||||
local_mask = torch.logical_and(local_mask, split_mask)
|
||||
|
||||
remainder_row = (i * token_per_frame) % block_size
|
||||
remainder_col = (j * token_per_frame) % block_size
|
||||
# get the padded size
|
||||
all_length_row = remainder_row + ((token_per_frame - 1) // block_size + 1) * block_size
|
||||
all_length_col = remainder_col + ((token_per_frame - 1) // block_size + 1) * block_size
|
||||
padded_local_mask = torch.zeros((all_length_row, all_length_col), device=query.device, dtype=torch.bool)
|
||||
padded_local_mask[remainder_row:remainder_row + token_per_frame, remainder_col:remainder_col + token_per_frame] = local_mask
|
||||
# shrink the mask
|
||||
block_mask = shrinkMaskStrict(padded_local_mask, block_size=block_size)
|
||||
# set the block mask to the final log mask
|
||||
block_row_start = (i * token_per_frame) // block_size
|
||||
block_col_start = (j * token_per_frame) // block_size
|
||||
block_row_end = block_row_start + block_mask.shape[0]
|
||||
block_col_end = block_col_start + block_mask.shape[1]
|
||||
final_log_mask[block_row_start:block_row_end, block_col_start:block_col_end] = torch.logical_or(
|
||||
final_log_mask[block_row_start:block_row_end, block_col_start:block_col_end], block_mask)
|
||||
print(f"mask sparsity: {1 - final_log_mask.sum() / final_log_mask.numel()}")
|
||||
return final_log_mask
|
||||
|
||||
class MaskMap:
|
||||
|
||||
def __init__(self, video_token_num=25440, num_frame=16):
|
||||
self.video_token_num = video_token_num
|
||||
self.num_frame = num_frame
|
||||
self.log_mask = None
|
||||
|
||||
def queryLogMask(self, query, sparse_type, block_size=128, decay_factor=0.5, model_type=None):
|
||||
log_mask = torch.ones((query.shape[0] // block_size, query.shape[0] // block_size), device=query.device, dtype=torch.bool)
|
||||
if self.log_mask is None:
|
||||
self.log_mask = gen_log_mask_shrinked(query, query.shape[0], self.video_token_num, self.num_frame, sparse_type=sparse_type, decay_factor=decay_factor, model_type=model_type, block_size=block_size)
|
||||
block_bound = self.video_token_num // block_size
|
||||
log_mask[:block_bound, :block_bound] = self.log_mask[:block_bound, :block_bound]
|
||||
return log_mask
|
||||
|
||||
def SpargeSageAttnBackend(query, key, value, mask_map=None, video_mask=None, pre_defined_mask=None, block_size=128):
|
||||
if video_mask.all():
|
||||
# dense case
|
||||
kv_border = pre_defined_mask[0].sum() if pre_defined_mask is not None else key.shape[0]
|
||||
output_video = sparse_sageattn(
|
||||
query[:mask_map.video_token_num].unsqueeze(0),
|
||||
key[:kv_border, :, :].unsqueeze(0),
|
||||
value[:kv_border, :, :].unsqueeze(0),
|
||||
mask_id=None,
|
||||
is_causal=False,
|
||||
tensor_layout="NHD",
|
||||
)[0]
|
||||
|
||||
# if pre_defined_mask is not None:
|
||||
# output_text = flashinfer.single_prefill_with_kv_cache(
|
||||
# q=query[mask_map.video_token_num:, :, :],
|
||||
# k=key[:pre_defined_mask[0].sum(), :, :],
|
||||
# v=value[:pre_defined_mask[0].sum(), :, :],
|
||||
# causal=False,
|
||||
# return_lse=False,
|
||||
# )
|
||||
# return torch.cat([output_video, output_text], dim=0)
|
||||
# else:
|
||||
return output_video
|
||||
|
||||
# sparse-sageattention only supports (b, h, s, d) layout, need rearrange first
|
||||
query_hnd = rearrange(query.unsqueeze(0), "b s h d -> b h s d")
|
||||
key_hnd = rearrange(key.unsqueeze(0), "b s h d -> b h s d")
|
||||
value_hnd = rearrange(value.unsqueeze(0), "b s h d -> b h s d")
|
||||
converted_mask = repeat(sparge_mask_convert(mask=video_mask, block_size=block_size), "s t -> b h s t", b=query_hnd.shape[0], h=query_hnd.shape[1])
|
||||
|
||||
converted_mask = converted_mask.to(torch.int8)
|
||||
if pre_defined_mask is None:
|
||||
# wan case
|
||||
output = sparse_sageattn(
|
||||
query_hnd[:, :, :mask_map.video_token_num, :],
|
||||
key_hnd[:, :, :mask_map.video_token_num, :],
|
||||
value_hnd[:, :, :mask_map.video_token_num, :],
|
||||
mask_id=converted_mask,
|
||||
is_causal=False,
|
||||
tensor_layout="HND",
|
||||
)
|
||||
|
||||
# rearrange back to (s, h, d), we know that b = 1
|
||||
output = rearrange(output, "b h s d -> s (b h) d", b=1)
|
||||
return output
|
||||
|
||||
# query_video = query_hnd[:, :, :mask_map.video_token_num, :]
|
||||
# key_video = key_hnd
|
||||
# value_video = value_hnd
|
||||
# kv_border = (pre_defined_mask[0].sum() + 63) // 64
|
||||
# converted_mask[:, :, :, kv_border:] = False
|
||||
# output_video = sparse_sageattn(
|
||||
# query_video,
|
||||
# key_video,
|
||||
# value_video,
|
||||
# mask_id=converted_mask[:, :, :mask_map.video_token_num // block_size, :],
|
||||
# is_causal=False,
|
||||
# tensor_layout="HND",
|
||||
# )
|
||||
|
||||
# # rearrange back to (s, h, d), we know that b = 1
|
||||
# output_video = rearrange(output_video, "b h s d -> s (b h) d", b=1)
|
||||
|
||||
# # gt = sparse_sageattn(
|
||||
# # query_video,
|
||||
# # key_video,
|
||||
# # value_video,
|
||||
# # mask_id=None,
|
||||
# # is_causal=False,
|
||||
# # tensor_layout="HND",
|
||||
# # )[0]
|
||||
|
||||
|
||||
|
||||
# # import pdb; pdb.set_trace()
|
||||
|
||||
# output_text = flashinfer.single_prefill_with_kv_cache(
|
||||
# q=query[mask_map.video_token_num:, :, :],
|
||||
# k=key[:pre_defined_mask[0].sum(), :, :],
|
||||
# v=value[:pre_defined_mask[0].sum(), :, :],
|
||||
# causal=False,
|
||||
# return_lse=False,
|
||||
# )
|
||||
|
||||
# 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
|
||||
|
||||
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)
|
||||
else:
|
||||
video_mask = mask_map.queryLogMask(query, sparsity_type, block_size=block_size, decay_factor=decay_factor, model_type=model_type) if mask_map else None
|
||||
|
||||
# backend = "sparse_sageattn" if use_sage_attention else "flashinfer"
|
||||
|
||||
# if backend == "flashinfer":
|
||||
# video_mask = video_mask[:mask_map.video_token_num // block_size, :mask_map.video_token_num // block_size]
|
||||
# # perform block-sparse attention on the video tokens
|
||||
# workspace_buffer = torch.empty(128 * 1024 * 1024, device=query.device, dtype=torch.uint8)
|
||||
# bsr_wrapper = flashinfer.BlockSparseAttentionWrapper(
|
||||
# workspace_buffer,
|
||||
# backend="fa2",
|
||||
# )
|
||||
|
||||
# indptr = get_indptr_from_mask(video_mask, query)
|
||||
# indices = get_indices_from_mask(video_mask, query)
|
||||
|
||||
# bsr_wrapper.plan(
|
||||
# indptr=indptr,
|
||||
# indices=indices,
|
||||
# M=mask_map.video_token_num,
|
||||
# N=mask_map.video_token_num,
|
||||
# R=block_size,
|
||||
# C=block_size,
|
||||
# num_qo_heads=num_head,
|
||||
# num_kv_heads=num_head,
|
||||
# head_dim=hidden_dim,
|
||||
# )
|
||||
|
||||
# return FlashInferBackend(query, key, value, mask_map, pre_defined_mask, bsr_wrapper)
|
||||
# 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)
|
||||
|
||||
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.close()
|
||||
# save the mask tensor
|
||||
torch.save(temporal_mask, "temporal_mask.pt")
|
||||
Reference in New Issue
Block a user