allow setting dense attention mode
This commit is contained in:
@@ -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"
|
||||
}
|
||||
|
||||
@@ -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:
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user