more compile friendly
This commit is contained in:
@@ -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]
|
||||
|
||||
@@ -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
@@ -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):
|
||||
|
||||
Reference in New Issue
Block a user