more compile friendly

This commit is contained in:
kijai
2025-07-18 00:53:12 +03:00
parent 605011a237
commit 712eecf64b
3 changed files with 55 additions and 39 deletions
+13 -8
View File
@@ -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]
+1 -1
View File
@@ -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()
+41 -30
View File
@@ -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):