From 712eecf64b29154212d78b6138a96437cf10966a Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 18 Jul 2025 00:53:12 +0300 Subject: [PATCH] more compile friendly --- nodes.py | 21 +++++++---- wanvideo/modules/attention.py | 2 +- wanvideo/modules/model.py | 71 ++++++++++++++++++++--------------- 3 files changed, 55 insertions(+), 39 deletions(-) diff --git a/nodes.py b/nodes.py index 30ab570..975c046 100644 --- a/nodes.py +++ b/nodes.py @@ -273,16 +273,16 @@ class WanVideoSetRadialAttention: return { "required": { "model": ("WANVIDEOMODEL", ), - "dense_attention_mode": ([ + "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_block": ("INT", {"default": 1, "min": 0, "max": 8, "step": 1, "tooltip": "Number of blocks to apply normal attention to"}), "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"}), + "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."}), } } @@ -292,7 +292,9 @@ class WanVideoSetRadialAttention: CATEGORY = "WanVideoWrapper" def loadmodel(self, model, dense_attention_mode, dense_block, dense_timestep, 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'] = {} @@ -2533,6 +2535,7 @@ class WanVideoSampler: else: transformer.slg_blocks = None + # Radial attention setup if transformer.attention_mode == "radial_sage_attention": if transformer_options is not None: dense_timestep = transformer_options.get("dense_timestep", 10) @@ -2541,12 +2544,14 @@ class WanVideoSampler: dense_attention_mode = transformer_options.get("dense_attention_mode", "sageattn") from .wanvideo.radial_attention.attn_mask import MaskMap - for block in transformer.blocks: + for i, block in enumerate(transformer.blocks): + block.self_attn.mask_map = block.dense_attention_mode = block.dense_timestep = block.self_attn.decay_factor = None + block.dense_block = True if i < dense_block else False 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.dense_attention_mode = dense_attention_mode + block.dense_timestep = dense_timestep 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] diff --git a/wanvideo/modules/attention.py b/wanvideo/modules/attention.py index 722a610..a715457 100644 --- a/wanvideo/modules/attention.py +++ b/wanvideo/modules/attention.py @@ -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': + else: return sageattn_func(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)).transpose(1, 2).contiguous() diff --git a/wanvideo/modules/model.py b/wanvideo/modules/model.py index 3173409..9d098c9 100644 --- a/wanvideo/modules/model.py +++ b/wanvideo/modules/model.py @@ -253,8 +253,7 @@ class WanSelfAttention(nn.Module): num_heads, qk_norm=True, eps=1e-6, - attention_mode='sdpa', - layer_idx=0): + attention_mode='sdpa'): assert out_features % num_heads == 0 super().__init__() self.dim = out_features @@ -265,14 +264,8 @@ class WanSelfAttention(nn.Module): self.attention_mode = attention_mode #radial attention - self.layer_idx = layer_idx - self.dense_timestep = 10 - self.dense_block = 1 - self.decay_factor = 0.2 - self.sparse_type = "radial" - self.dense_attention_mode = "sageattn" self.mask_map = None - + self.decay_factor = 0.2 # layers self.q = nn.Linear(in_features, out_features) @@ -289,7 +282,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, current_step=0): + def forward(self, q, k, v, seq_lens, block_mask=None): r""" Args: x(Tensor): Shape [B, L, num_heads, C / num_heads] @@ -321,19 +314,6 @@ 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': - dense_step = current_step < self.dense_timestep or self.layer_idx < self.dense_block or self.sparse_type == "dense" - if dense_step: - 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: x = attention( q, k, v, @@ -347,6 +327,26 @@ class WanSelfAttention(nn.Module): return x + def forward_radial(self, q, k, v, dense_step=False): + r""" + Args: + x(Tensor): Shape [B, L, num_heads, C / num_heads] + seq_lens(Tensor): Shape [B] + grid_sizes(Tensor): Shape [B, 3], the second dimension contains (F, H, W) + freqs(Tensor): Rope freqs, shape [1024, C / num_heads / 2] + """ + + 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) + 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) + + # output + x = x.flatten(2) + x = self.o(x) + + return x + def forward_multitalk(self, q, k, v, seq_lens, grid_sizes, ref_target_masks): x = attention( q, k, v, @@ -589,7 +589,6 @@ class WanAttentionBlock(nn.Module): eps=1e-6, attention_mode='sdpa', rope_func="comfy", - block_idx=0 ): super().__init__() self.dim = out_features @@ -600,12 +599,15 @@ class WanAttentionBlock(nn.Module): self.eps = eps self.attention_mode = attention_mode self.rope_func = rope_func - self.block_idx = block_idx + #radial attn + self.dense_timestep = 10 + self.dense_block = False + self.dense_attention_mode = "sageattn" # layers self.norm1 = WanLayerNorm(out_features, eps) self.self_attn = WanSelfAttention(in_features, out_features, num_heads, qk_norm, - eps, self.attention_mode, self.block_idx) + eps, self.attention_mode) if cross_attn_type != "no_cross_attn": self.norm3 = WanLayerNorm( out_features, eps, @@ -660,6 +662,7 @@ class WanAttentionBlock(nn.Module): def forward( self, x, + block_idx, e, seq_lens, grid_sizes, @@ -734,8 +737,16 @@ 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 and self.dense_timestep is not None and current_step < self.dense_timestep: + 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, current_step=current_step) + y = self.self_attn.forward(q, k, v, seq_lens, block_mask=block_mask) # FETA if enhance_enabled: @@ -1151,8 +1162,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, block_idx=i) - for i in range(num_layers) + attention_mode=self.attention_mode, rope_func=self.rope_func) + for _ in range(num_layers) ]) # head @@ -1796,7 +1807,7 @@ class WanModel(ModelMixin, ConfigMixin): continue if b <= self.blocks_to_swap and self.blocks_to_swap >= 0: block.to(self.main_device) - x = block(x, **kwargs) + x = block(x, b, **kwargs) #uni3c controlnet if pdc_controlnet_states is not None and b < len(pdc_controlnet_states):