From 605011a23742501a02c43c389bb3109663710e99 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Thu, 17 Jul 2025 23:47:08 +0300 Subject: [PATCH] allow setting dense attention mode --- nodes.py | 49 ++++++++++++++++++++++++++ wanvideo/modules/model.py | 10 +++++- wanvideo/radial_attention/attn_mask.py | 1 + 3 files changed, 59 insertions(+), 1 deletion(-) diff --git a/nodes.py b/nodes.py index fa55a0e..30ab570 100644 --- a/nodes.py +++ b/nodes.py @@ -267,6 +267,42 @@ 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_block": ("INT", {"default": 1, "min": 0, "max": 8, "step": 1, "tooltip": "The dense block to apply normal attention"}), + "dense_timestep": ("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": "The dense block to apply normal attention"}), + } + } + + RETURN_TYPES = ("WANVIDEOMODEL",) + RETURN_NAMES = ("model", ) + FUNCTION = "loadmodel" + CATEGORY = "WanVideoWrapper" + + def loadmodel(self, model, dense_attention_mode, dense_block, dense_timestep, decay_factor): + + 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_block"] = dense_block + patcher.model_options["transformer_options"]["dense_timestep"] = dense_timestep + patcher.model_options["transformer_options"]["decay_factor"] = decay_factor + + return (patcher,) class WanVideoTorchCompileSettings: @classmethod @@ -2498,9 +2534,20 @@ class WanVideoSampler: transformer.slg_blocks = None if transformer.attention_mode == "radial_sage_attention": + if transformer_options is not None: + dense_timestep = transformer_options.get("dense_timestep", 10) + dense_block = transformer_options.get("dense_block", 1) + decay_factor = transformer_options.get("decay_factor", 0.2) + dense_attention_mode = transformer_options.get("dense_attention_mode", "sageattn") + from .wanvideo.radial_attention.attn_mask import MaskMap for block in transformer.blocks: block.self_attn.mask_map = MaskMap(video_token_num=seq_len, num_frame=latent_video_length) + block.self_attn.dense_attention_mode = dense_attention_mode + block.self_attn.dense_timestep = dense_timestep + block.self_attn.dense_block = dense_block + block.self_attn.decay_factor = decay_factor + log.info(f"Radial attention mode enabled. dense_attention_mode: {dense_attention_mode}, dense_timestep: {dense_timestep}, dense_block: {dense_block}, decay_factor: {decay_factor}") self.cache_state = [None, None] if phantom_latents is not None: @@ -3819,6 +3866,7 @@ NODE_CLASS_MAPPINGS = { "WanVideoApplyNAG": WanVideoApplyNAG, "WanVideoMiniMaxRemoverEmbeds": WanVideoMiniMaxRemoverEmbeds, "WanVideoFreeInitArgs": WanVideoFreeInitArgs, + "WanVideoSetRadialAttention": WanVideoSetRadialAttention } NODE_DISPLAY_NAME_MAPPINGS = { "WanVideoSampler": "WanVideo Sampler", @@ -3858,4 +3906,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/wanvideo/modules/model.py b/wanvideo/modules/model.py index e6d867b..3173409 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -270,6 +270,7 @@ class WanSelfAttention(nn.Module): self.dense_block = 1 self.decay_factor = 0.2 self.sparse_type = "radial" + self.dense_attention_mode = "sageattn" self.mask_map = None @@ -323,7 +324,14 @@ class WanSelfAttention(nn.Module): elif self.attention_mode == 'radial_sage_attention': dense_step = current_step < self.dense_timestep or self.layer_idx < self.dense_block or self.sparse_type == "dense" if dense_step: - x = RadialAttention(query=q, key=k, value=v, mask_map=self.mask_map, sparsity_type="dense", block_size=128, decay_factor=self.decay_factor) + if self.dense_attention_mode == "sparse_sage_attn": + x = RadialAttention(query=q, key=k, value=v, mask_map=self.mask_map, sparsity_type="dense", block_size=128, decay_factor=self.decay_factor) + else: + x = attention( + q, k, v, + k_lens=seq_lens, + attention_mode=self.dense_attention_mode + ) else: x = RadialAttention(query=q, key=k, value=v, mask_map=self.mask_map, sparsity_type="radial", block_size=128, decay_factor=self.decay_factor) else: diff --git a/wanvideo/radial_attention/attn_mask.py b/wanvideo/radial_attention/attn_mask.py index 40384b4..fcffc32 100644 --- a/wanvideo/radial_attention/attn_mask.py +++ b/wanvideo/radial_attention/attn_mask.py @@ -191,6 +191,7 @@ def SpargeSageAttnBackend(query, key, value, mask_map=None, video_mask=None, blo return output.transpose(1, 2).contiguous() +@torch.compiler.disable() def RadialAttention(query, key, value, mask_map=None, sparsity_type="radial", block_size=128, decay_factor=1, use_sage_attention=False): if sparsity_type == "dense": video_mask = torch.ones((mask_map.video_token_num // block_size, mask_map.video_token_num // block_size), device=device, dtype=torch.bool)