From d3aba9011998a8b2502063531d6f5825af19e9e1 Mon Sep 17 00:00:00 2001 From: komikndr Date: Sat, 19 Jul 2025 20:05:19 +0700 Subject: [PATCH 1/2] adding 64 block_size --- nodes.py | 13 ++-- wanvideo/radial_attention/attn_mask.py | 97 +++++++++++++------------- 2 files changed, 58 insertions(+), 52 deletions(-) diff --git a/nodes.py b/nodes.py index 18b59c6..d79da04 100644 --- a/nodes.py +++ b/nodes.py @@ -92,6 +92,7 @@ class WanVideoSetRadialAttention: "dense_vace_blocks": ("INT", {"default": 1, "min": 0, "max": 15, "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."}), + "block_size":("INT", {"default": 128, "min": 64, "max":128, "step": 64, "tooltip": "Control radial attention block size, either 128 or 64"}), } } @@ -101,7 +102,7 @@ class WanVideoSetRadialAttention: 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): + def loadmodel(self, model, dense_attention_mode, dense_blocks, dense_vace_blocks, dense_timesteps, decay_factor, block_size): if "radial" not in model.model.diffusion_model.attention_mode: raise Exception("Enable radial attention first in the model loader.") @@ -114,6 +115,7 @@ class WanVideoSetRadialAttention: 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 + patcher.model_options["transformer_options"]["block_size"] = block_size return (patcher,) @@ -1639,14 +1641,15 @@ class WanVideoSampler: # 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.") + if latent.shape[2] % 8 != 0 or latent.shape[3] % 8 != 0: + raise Exception(f"Radial attention mode only supports image size divisible by 128 or 64.") 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) + block_size = transformer_options.get("block_size", 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.") @@ -1654,7 +1657,7 @@ class WanVideoSampler: 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.self_attn.mask_map = MaskMap(video_token_num=seq_len, num_frame=latent_video_length, block_size=block_size) block.dense_attention_mode = dense_attention_mode block.dense_timesteps = dense_timesteps block.self_attn.decay_factor = decay_factor @@ -1662,7 +1665,7 @@ class WanVideoSampler: 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.self_attn.mask_map = MaskMap(video_token_num=seq_len, num_frame=latent_video_length, block_size=block_size) block.dense_attention_mode = dense_attention_mode block.dense_timesteps = dense_timesteps block.self_attn.decay_factor = decay_factor diff --git a/wanvideo/radial_attention/attn_mask.py b/wanvideo/radial_attention/attn_mask.py index 1498ef0..830b1a4 100644 --- a/wanvideo/radial_attention/attn_mask.py +++ b/wanvideo/radial_attention/attn_mask.py @@ -11,14 +11,13 @@ except: from comfy import model_management as mm device = mm.get_torch_device() - from tqdm import tqdm -def shrinkMaskStrict(mask, block_size=128): +def shrinkMaskStrict(mask, block_size): 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 + 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 @@ -28,24 +27,22 @@ def shrinkMaskStrict(mask, block_size=128): 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"]) +def get_diagonal_split_mask(i, j, token_per_frame, sparse_type, block_size): + 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 + threshold = block_size # CHANGE, can 64 or 128 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) + return torch.ones((token_per_frame, token_per_frame), device=device, dtype=torch.bool) if modular == 0 \ + else 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"]) +def get_window_width(i, j, token_per_frame, sparse_type, decay_factor, block_size): + assert sparse_type in ["radial"] dist = abs(i - j) if dist < 1: return token_per_frame @@ -53,18 +50,13 @@ def get_window_width(i, j, token_per_frame, sparse_type, decay_factor=1, block_s 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 + return max(decay_length, block_size) -def gen_log_mask_shrinked(device, s, video_token_num, num_frame, block_size=128, sparse_type="log", decay_factor=0.5): +def gen_log_mask_shrinked(device, s, video_token_num, num_frame, block_size, sparse_type, decay_factor): """ 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 @@ -73,87 +65,98 @@ def gen_log_mask_shrinked(device, s, video_token_num, num_frame, block_size=128, 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) + window_width = get_window_width(i, j, token_per_frame, sparse_type, decay_factor, 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) + split_mask = get_diagonal_split_mask(i, j, token_per_frame, sparse_type, block_size) 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) + block_mask = shrinkMaskStrict(padded_local_mask, 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): + def __init__(self, video_token_num=25440, num_frame=16, block_size=128): self.video_token_num = video_token_num self.num_frame = num_frame self.log_mask = None + self.block_size = block_size - def queryLogMask(self, seq_len, sparse_type, block_size=128, decay_factor=0.5): + def queryLogMask(self, seq_len, sparse_type, block_size=None, decay_factor=0.5): + block_size = block_size or self.block_size 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) + self.log_mask = gen_log_mask_shrinked( + device, seq_len, self.video_token_num, self.num_frame, + block_size=block_size, sparse_type=sparse_type, decay_factor=decay_factor + ) 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): +def RadialSpargeSageAttnDense(query, key, value, mask_map): # dense case - output_video = sparse_sageattn( - query[:,:mask_map.video_token_num], - key[:,:key.shape[1], :, :], - value[:,:key.shape[1], :, :], + return 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() + tensor_layout="NHD" + ).contiguous() @torch.compiler.disable() -def RadialSpargeSageAttn(query, key, value, mask_map, block_size=128, decay_factor=1): +def RadialSpargeSageAttn(query, key, value, mask_map, decay_factor): # Simple cache based on function arguments if not hasattr(RadialSpargeSageAttn, "_cache"): RadialSpargeSageAttn._cache = {} - cache_key = (query.shape[-2], block_size, decay_factor) + # print(mask_map.block_size) + block_size = mask_map.block_size + cache_key = (query.shape[-2], mask_map.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 + video_mask = mask_map.queryLogMask(query.shape[0] * query.shape[1], "radial", block_size=block_size, decay_factor=decay_factor) + + # based on https://github.com/mit-han-lab/radial-attention/blob/3ec33ce9633adadadcbb7692c8a1983d5e82d15a/radial_attn/attn_mask.py#L7 + if block_size == 128: + mask = torch.repeat_interleave(video_mask, 2, dim=1) + elif block_size == 64: + reshaped_mask = video_mask.view(video_mask.shape[0] // 2, 2, video_mask.shape[1]) + mask = torch.max(reshaped_mask, dim=1).values + input_mask = mask.unsqueeze(0).unsqueeze(1).expand(1, query.shape[-2], mask.shape[0], mask.shape[1]) RadialSpargeSageAttn._cache[cache_key] = input_mask - output = sparse_sageattn( + return 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 + ).contiguous() From 47978275bb63a7678c3aef50e1738b145f8a6090 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 20 Jul 2025 02:13:07 +0300 Subject: [PATCH 2/2] Update nodes.py --- nodes.py | 20 ++++++++++++++------ 1 file changed, 14 insertions(+), 6 deletions(-) diff --git a/nodes.py b/nodes.py index d79da04..7e4df58 100644 --- a/nodes.py +++ b/nodes.py @@ -92,7 +92,7 @@ class WanVideoSetRadialAttention: "dense_vace_blocks": ("INT", {"default": 1, "min": 0, "max": 15, "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."}), - "block_size":("INT", {"default": 128, "min": 64, "max":128, "step": 64, "tooltip": "Control radial attention block size, either 128 or 64"}), + "block_size":([128, 64], {"default": 128, "tooltip": "Radial attention block size"}), } } @@ -1641,17 +1641,25 @@ class WanVideoSampler: # Radial attention setup if transformer.attention_mode == "radial_sage_attention": - if latent.shape[2] % 8 != 0 or latent.shape[3] % 8 != 0: - raise Exception(f"Radial attention mode only supports image size divisible by 128 or 64.") - 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) block_size = transformer_options.get("block_size", 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.") + + if latent.shape[2] % (block_size/8) != 0 or latent.shape[3] % (block_size/8) != 0: + # Calculate closest valid latent sizes + block_div = int(block_size // 8) + closest_h = round(latent.shape[2] / block_div) * block_div + closest_w = round(latent.shape[3] / block_div) * block_div + closest_h_px = closest_h * 8 + closest_w_px = closest_w * 8 + raise Exception( + f"Radial attention mode only supports image size divisible by block size. " + f"Got {latent.shape[3] * 8}x{latent.shape[2] * 8} with block size {block_size}.\n" + f"Closest valid sizes: {closest_w_px}x{closest_h_px} (width x height in pixels)." + ) from .wanvideo.radial_attention.attn_mask import MaskMap for i, block in enumerate(transformer.blocks):