diff --git a/nodes.py b/nodes.py index f4ae73d..426f29c 100644 --- a/nodes.py +++ b/nodes.py @@ -267,6 +267,47 @@ class WanVideoSetBlockSwap: return (patcher,) +class WanVideoSetRadialAttention: + @classmethod + def INPUT_TYPES(s): + return { + "required": { + "model": ("WANVIDEOMODEL", ), + "dense_attention_mode": ([ + "sdpa", + "flash_attn_2", + "flash_attn_3", + "sageattn", + "sparse_sage_attention", + ], {"default": "sageattn", "tooltip": "The attention mode for dense attention"}), + "dense_blocks": ("INT", {"default": 1, "min": 0, "max": 40, "step": 1, "tooltip": "Number of blocks to apply normal attention to"}), + "dense_vace_blocks": ("INT", {"default": 15, "min": 0, "max": 40, "step": 1, "tooltip": "Number of vace blocks to apply normal attention to"}), + "dense_timesteps": ("INT", {"default": 10, "min": 0, "max": 100, "step": 1, "tooltip": "The step to start applying sparse attention"}), + "decay_factor": ("FLOAT", {"default": 0.2, "min": 0, "max": 1, "step": 0.01, "tooltip": "Controls how quickly the attention window shrinks as the distance between frames increases in the sparse attention mask."}), + } + } + + RETURN_TYPES = ("WANVIDEOMODEL",) + RETURN_NAMES = ("model", ) + FUNCTION = "loadmodel" + CATEGORY = "WanVideoWrapper" + DESCRIPTION = "Sets radial attention parameters, dense attention refers to normal attention" + + def loadmodel(self, model, dense_attention_mode, dense_blocks, dense_vace_blocks, dense_timesteps, decay_factor): + if "radial" not in model.model.diffusion_model.attention_mode: + raise Exception("Enable radial attention first in the model loader.") + + patcher = model.clone() + if 'transformer_options' not in patcher.model_options: + patcher.model_options['transformer_options'] = {} + + patcher.model_options["transformer_options"]["dense_attention_mode"] = dense_attention_mode + patcher.model_options["transformer_options"]["dense_blocks"] = dense_blocks + patcher.model_options["transformer_options"]["dense_vace_blocks"] = dense_vace_blocks + patcher.model_options["transformer_options"]["dense_timesteps"] = dense_timesteps + patcher.model_options["transformer_options"]["decay_factor"] = decay_factor + + return (patcher,) class WanVideoTorchCompileSettings: @classmethod @@ -2340,8 +2381,11 @@ class WanVideoSampler: 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: @@ -2497,6 +2541,38 @@ class WanVideoSampler: else: transformer.slg_blocks = None + # Radial attention setup + if transformer.attention_mode == "radial_sage_attention": + if latent.shape[2] % 16 != 0 or latent.shape[3] % 16 != 0: + raise Exception(f"Radial attention mode only supports image size divisible by 128.") + + dense_timesteps = transformer_options.get("dense_timesteps", None) + dense_blocks = transformer_options.get("dense_blocks", None) + dense_vace_blocks = transformer_options.get("dense_vace_blocks", None) + decay_factor = transformer_options.get("decay_factor", None) + dense_attention_mode = transformer_options.get("dense_attention_mode", None) + if dense_timesteps is None: + raise Exception("Radial attention mode is enabled, but no parameters are provided. Add the `WanVideoSetRadialAttention` node to the model to set the parameters.") + + from .wanvideo.radial_attention.attn_mask import MaskMap + for i, block in enumerate(transformer.blocks): + block.self_attn.mask_map = block.dense_attention_mode = block.dense_timesteps = block.self_attn.decay_factor = None + block.dense_block = True if i < dense_blocks else False + block.self_attn.mask_map = MaskMap(video_token_num=seq_len, num_frame=latent_video_length) + block.dense_attention_mode = dense_attention_mode + block.dense_timesteps = dense_timesteps + block.self_attn.decay_factor = decay_factor + if transformer.vace_layers is not None: + for i, block in enumerate(transformer.vace_blocks): + block.self_attn.mask_map = block.dense_attention_mode = block.dense_timesteps = block.self_attn.decay_factor = None + block.dense_block = True if i < dense_vace_blocks else False + block.self_attn.mask_map = MaskMap(video_token_num=seq_len, num_frame=latent_video_length) + block.dense_attention_mode = dense_attention_mode + block.dense_timesteps = dense_timesteps + block.self_attn.decay_factor = decay_factor + + log.info(f"Radial attention mode enabled. dense_attention_mode: {dense_attention_mode}, dense_timesteps: {dense_timesteps}, dense_blocks: {dense_blocks}, decay_factor: {decay_factor}") + self.cache_state = [None, None] if phantom_latents is not None: log.info(f"Phantom latents shape: {phantom_latents.shape}") @@ -3814,6 +3890,7 @@ NODE_CLASS_MAPPINGS = { "WanVideoApplyNAG": WanVideoApplyNAG, "WanVideoMiniMaxRemoverEmbeds": WanVideoMiniMaxRemoverEmbeds, "WanVideoFreeInitArgs": WanVideoFreeInitArgs, + "WanVideoSetRadialAttention": WanVideoSetRadialAttention } NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoSampler": "WanVideo Sampler", @@ -3853,4 +3930,5 @@ NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoApplyNAG": "WanVideo Apply NAG", "WanVideoMiniMaxRemoverEmbeds": "WanVideo MiniMax Remover Embeds", "WanVideoFreeInitArgs": "WanVideo Free Init Args", + "WanVideoSetRadialAttention": "WanVideo Set Radial Attention" } 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/attention.py b/wanvideo/modules/attention.py index 722a610..39a3a80 100644 --- a/wanvideo/modules/attention.py +++ b/wanvideo/modules/attention.py @@ -18,9 +18,9 @@ try: @torch.compiler.disable() def sageattn_func(q, k, v, attn_mask=None, dropout_p=0, is_causal=False): if q.dtype == torch.float32: - return sageattn(q.to(torch.float16), k.to(torch.float16), v.to(torch.float16), attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal).to(torch.float32) + return sageattn(q.to(torch.float16), k.to(torch.float16), v.to(torch.float16), attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout="NHD").to(torch.float32) else: - return sageattn(q, k, v, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal) + return sageattn(q, k, v, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout="NHD") except Exception as e: print(f"Warning: Could not load sageattention: {str(e)}") if isinstance(e, ModuleNotFoundError): @@ -182,5 +182,5 @@ def attention( ) elif attention_mode == 'sdpa': return torch.nn.functional.scaled_dot_product_attention(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)).transpose(1, 2).contiguous() - elif attention_mode == 'sageattn': - return sageattn_func(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)).transpose(1, 2).contiguous() + else: + return sageattn_func(q, k, v).contiguous() diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 13322b6..31664ea 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -15,6 +15,10 @@ try: except: BlockMask = create_block_mask = flex_attention = None pass +try: + from ..radial_attention.attn_mask import RadialSpargeSageAttn, RadialSpargeSageAttnDense +except: + pass from .attention import attention import numpy as np @@ -259,6 +263,10 @@ class WanSelfAttention(nn.Module): self.eps = eps self.attention_mode = attention_mode + #radial attention + self.mask_map = None + self.decay_factor = 0.2 + # layers self.q = nn.Linear(in_features, out_features) self.k = nn.Linear(in_features, out_features) @@ -319,6 +327,16 @@ class WanSelfAttention(nn.Module): return x + def forward_radial(self, q, k, v, dense_step=False): + if dense_step: + x = RadialSpargeSageAttnDense(q, k, v, self.mask_map) + else: + x = RadialSpargeSageAttn(q, k, v, self.mask_map, decay_factor=self.decay_factor) + + x = self.o(x.flatten(2)) + + return x + def forward_multitalk(self, q, k, v, seq_lens, grid_sizes, ref_target_masks): x = attention( q, k, v, @@ -560,7 +578,7 @@ class WanAttentionBlock(nn.Module): cross_attn_norm=False, eps=1e-6, attention_mode='sdpa', - rope_func="comfy" + rope_func="comfy", ): super().__init__() self.dim = out_features @@ -571,6 +589,10 @@ class WanAttentionBlock(nn.Module): self.eps = eps self.attention_mode = attention_mode self.rope_func = rope_func + #radial attn + self.dense_timesteps = 10 + self.dense_block = False + self.dense_attention_mode = "sageattn" # layers self.norm1 = WanLayerNorm(out_features, eps) @@ -704,6 +726,14 @@ 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) + elif self.attention_mode == "radial_sage_attention": + if self.dense_block or self.dense_timesteps is not None and current_step < self.dense_timesteps: + if self.dense_attention_mode == "sparse_sage_attn": + y = self.self_attn.forward_radial(q, k, v, dense_step=True) + else: + y = self.self_attn.forward(q, k, v, seq_lens, block_mask=block_mask) + else: + y = self.self_attn.forward_radial(q, k, v, dense_step=False) else: y = self.self_attn.forward(q, k, v, seq_lens, block_mask=block_mask) diff --git a/wanvideo/radial_attention/attn_mask.py b/wanvideo/radial_attention/attn_mask.py new file mode 100644 index 0000000..4e13674 --- /dev/null +++ b/wanvideo/radial_attention/attn_mask.py @@ -0,0 +1,156 @@ +# based on https://github.com/mit-han-lab/radial-attention/blob/main/radial_attn/attn_mask.py +import torch +try: + from sparse_sageattn import sparse_sageattn +except: + sparse_sageattn = None + raise ImportError("Package is not installed: https://github.com/jt-zhang/Sparse_SageAttention_API") + +from comfy import model_management as mm +device = mm.get_torch_device() + +from tqdm import tqdm + +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 get_diagonal_split_mask(i, j, token_per_frame, sparse_type): + 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=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=device, dtype=torch.bool) + else: + return torch.zeros((token_per_frame, token_per_frame), device=device, dtype=torch.bool) + +def get_window_width(i, j, token_per_frame, sparse_type, decay_factor=1, block_size=128): + assert(sparse_type in ["radial"]) + dist = abs(i - j) + if dist < 1: + return token_per_frame + if dist == 1: + return token_per_frame // 2 + 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(device, s, video_token_num, num_frame, block_size=128, sparse_type="log", decay_factor=0.5): + """ + 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=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=device).view(1, -1) + row_indices = torch.arange(0, token_per_frame, device=device).view(-1, 1) + final_log_mask[video_text_border:] = True + final_log_mask[:, video_text_border:] = True + for i in tqdm(range(num_frame), desc="Frames (i)"): + for j in range(num_frame): + local_mask = torch.zeros((token_per_frame, token_per_frame), device=device, dtype=torch.bool) + if j == 0: # this is attention sink + local_mask = torch.ones((token_per_frame, token_per_frame), device=device, dtype=torch.bool) + else: + window_width = get_window_width(i, j, token_per_frame, sparse_type, decay_factor=decay_factor, block_size=block_size) + local_mask = torch.abs(col_indices - row_indices) <= window_width + split_mask = get_diagonal_split_mask(i, j, token_per_frame, sparse_type) + 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=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, seq_len, sparse_type, block_size=128, decay_factor=0.5): + log_mask = torch.ones((seq_len // block_size, seq_len // block_size), device=device, dtype=torch.bool) + if self.log_mask is None: + self.log_mask = gen_log_mask_shrinked(device, seq_len, self.video_token_num, self.num_frame, sparse_type=sparse_type, decay_factor=decay_factor, 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 + +@torch.compiler.disable() +def RadialSpargeSageAttnDense(query, key, value, mask_map=None): + # dense case + output_video = sparse_sageattn( + query[:,:mask_map.video_token_num], + key[:,:key.shape[1], :, :], + value[:,:key.shape[1], :, :], + mask_id=None, + is_causal=False, + tensor_layout="NHD", + ) + + return output_video.contiguous() + +@torch.compiler.disable() +def RadialSpargeSageAttn(query, key, value, mask_map, block_size=128, decay_factor=1): + # Simple cache based on function arguments + if not hasattr(RadialSpargeSageAttn, "_cache"): + RadialSpargeSageAttn._cache = {} + cache_key = (query.shape[-2], block_size, decay_factor) + if cache_key in RadialSpargeSageAttn._cache: + input_mask = RadialSpargeSageAttn._cache[cache_key] + else: + print("Radial Attention: Generating block mask") + video_mask = mask_map.queryLogMask(query.shape[0] * query.shape[1], "radial", block_size=block_size, decay_factor=decay_factor) + mask = torch.repeat_interleave(video_mask, 2, dim=1) #s, t + input_mask = mask.unsqueeze(0).unsqueeze(1).expand(1, query.shape[-2], mask.shape[0], mask.shape[1]) # b, h, s, t + RadialSpargeSageAttn._cache[cache_key] = input_mask + + output = sparse_sageattn( + query[:, :, :mask_map.video_token_num, :], + key[:, :, :mask_map.video_token_num, :], + value[:, :, :mask_map.video_token_num, :], + mask_id=input_mask.to(torch.int8), + is_causal=False, + tensor_layout="NHD" + ) + + return output.contiguous() \ No newline at end of file