diff --git a/nodes.py b/nodes.py index 0ae9e3a..fce0ac9 100644 --- a/nodes.py +++ b/nodes.py @@ -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}") diff --git a/nodes_model_loading.py b/nodes_model_loading.py index fa89277..d5cb4aa 100644 --- a/nodes_model_loading.py +++ b/nodes_model_loading.py @@ -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", ), diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 13322b6..902b753 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -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, diff --git a/wanvideo/radial_attention/attn_mask.py b/wanvideo/radial_attention/attn_mask.py new file mode 100644 index 0000000..e3c38ff --- /dev/null +++ b/wanvideo/radial_attention/attn_mask.py @@ -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") \ No newline at end of file