Merge branch 'radial_attn'
This commit is contained in:
@@ -267,6 +267,47 @@ 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_blocks": ("INT", {"default": 1, "min": 0, "max": 40, "step": 1, "tooltip": "Number of blocks to apply normal attention to"}),
|
||||
"dense_vace_blocks": ("INT", {"default": 15, "min": 0, "max": 40, "step": 1, "tooltip": "Number of vace blocks to apply normal attention to"}),
|
||||
"dense_timesteps": ("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": "Controls how quickly the attention window shrinks as the distance between frames increases in the sparse attention mask."}),
|
||||
}
|
||||
}
|
||||
|
||||
RETURN_TYPES = ("WANVIDEOMODEL",)
|
||||
RETURN_NAMES = ("model", )
|
||||
FUNCTION = "loadmodel"
|
||||
CATEGORY = "WanVideoWrapper"
|
||||
DESCRIPTION = "Sets radial attention parameters, dense attention refers to normal attention"
|
||||
|
||||
def loadmodel(self, model, dense_attention_mode, dense_blocks, dense_vace_blocks, dense_timesteps, 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'] = {}
|
||||
|
||||
patcher.model_options["transformer_options"]["dense_attention_mode"] = dense_attention_mode
|
||||
patcher.model_options["transformer_options"]["dense_blocks"] = dense_blocks
|
||||
patcher.model_options["transformer_options"]["dense_vace_blocks"] = dense_vace_blocks
|
||||
patcher.model_options["transformer_options"]["dense_timesteps"] = dense_timesteps
|
||||
patcher.model_options["transformer_options"]["decay_factor"] = decay_factor
|
||||
|
||||
return (patcher,)
|
||||
|
||||
class WanVideoTorchCompileSettings:
|
||||
@classmethod
|
||||
@@ -2340,8 +2381,11 @@ class WanVideoSampler:
|
||||
if mask.shape[2] != noise.shape[1]:
|
||||
mask = torch.cat([torch.zeros(1, noise.shape[0], noise.shape[1] - mask.shape[2], noise.shape[2], noise.shape[3]), mask], dim=2)
|
||||
|
||||
if (extra_latents := image_embeds.get("extra_latents", None)) is not None:
|
||||
|
||||
if (extra_latents := image_embeds.get("extra_latents", None)) is not None:
|
||||
encoded_image_latents = extra_latents["samples"].squeeze(0).to(noise)
|
||||
if (empty_latent_indices := extra_latents.get("empty_latent_indices", None)) is not None and len(empty_latent_indices) > 0:
|
||||
if (empty_latent_indices := extra_latents.get("empty_latent_indices", None)) is not None and len(empty_latent_indices) > 0:
|
||||
noise_out = encoded_image_latents.clone()
|
||||
for idx in empty_latent_indices:
|
||||
@@ -2497,6 +2541,38 @@ class WanVideoSampler:
|
||||
else:
|
||||
transformer.slg_blocks = None
|
||||
|
||||
# Radial attention setup
|
||||
if transformer.attention_mode == "radial_sage_attention":
|
||||
if latent.shape[2] % 16 != 0 or latent.shape[3] % 16 != 0:
|
||||
raise Exception(f"Radial attention mode only supports image size divisible by 128.")
|
||||
|
||||
dense_timesteps = transformer_options.get("dense_timesteps", None)
|
||||
dense_blocks = transformer_options.get("dense_blocks", None)
|
||||
dense_vace_blocks = transformer_options.get("dense_vace_blocks", None)
|
||||
decay_factor = transformer_options.get("decay_factor", None)
|
||||
dense_attention_mode = transformer_options.get("dense_attention_mode", None)
|
||||
if dense_timesteps is None:
|
||||
raise Exception("Radial attention mode is enabled, but no parameters are provided. Add the `WanVideoSetRadialAttention` node to the model to set the parameters.")
|
||||
|
||||
from .wanvideo.radial_attention.attn_mask import MaskMap
|
||||
for i, block in enumerate(transformer.blocks):
|
||||
block.self_attn.mask_map = block.dense_attention_mode = block.dense_timesteps = block.self_attn.decay_factor = None
|
||||
block.dense_block = True if i < dense_blocks else False
|
||||
block.self_attn.mask_map = MaskMap(video_token_num=seq_len, num_frame=latent_video_length)
|
||||
block.dense_attention_mode = dense_attention_mode
|
||||
block.dense_timesteps = dense_timesteps
|
||||
block.self_attn.decay_factor = decay_factor
|
||||
if transformer.vace_layers is not None:
|
||||
for i, block in enumerate(transformer.vace_blocks):
|
||||
block.self_attn.mask_map = block.dense_attention_mode = block.dense_timesteps = block.self_attn.decay_factor = None
|
||||
block.dense_block = True if i < dense_vace_blocks else False
|
||||
block.self_attn.mask_map = MaskMap(video_token_num=seq_len, num_frame=latent_video_length)
|
||||
block.dense_attention_mode = dense_attention_mode
|
||||
block.dense_timesteps = dense_timesteps
|
||||
block.self_attn.decay_factor = decay_factor
|
||||
|
||||
log.info(f"Radial attention mode enabled. dense_attention_mode: {dense_attention_mode}, dense_timesteps: {dense_timesteps}, dense_blocks: {dense_blocks}, decay_factor: {decay_factor}")
|
||||
|
||||
self.cache_state = [None, None]
|
||||
if phantom_latents is not None:
|
||||
log.info(f"Phantom latents shape: {phantom_latents.shape}")
|
||||
@@ -3814,6 +3890,7 @@ NODE_CLASS_MAPPINGS = {
|
||||
"WanVideoApplyNAG": WanVideoApplyNAG,
|
||||
"WanVideoMiniMaxRemoverEmbeds": WanVideoMiniMaxRemoverEmbeds,
|
||||
"WanVideoFreeInitArgs": WanVideoFreeInitArgs,
|
||||
"WanVideoSetRadialAttention": WanVideoSetRadialAttention
|
||||
}
|
||||
NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoSampler": "WanVideo Sampler",
|
||||
@@ -3853,4 +3930,5 @@ NODE_DISPLAY_NAME_MAPPINGS = {
|
||||
"WanVideoApplyNAG": "WanVideo Apply NAG",
|
||||
"WanVideoMiniMaxRemoverEmbeds": "WanVideo MiniMax Remover Embeds",
|
||||
"WanVideoFreeInitArgs": "WanVideo Free Init Args",
|
||||
"WanVideoSetRadialAttention": "WanVideo Set Radial Attention"
|
||||
}
|
||||
|
||||
@@ -472,8 +472,7 @@ class WanVideoModelLoader:
|
||||
"flash_attn_3",
|
||||
"sageattn",
|
||||
"flex_attention",
|
||||
#"spargeattn", needs tuning
|
||||
#"spargeattn_tune",
|
||||
"radial_sage_attention",
|
||||
], {"default": "sdpa"}),
|
||||
"compile_args": ("WANCOMPILEARGS", ),
|
||||
"block_swap_args": ("BLOCKSWAPARGS", ),
|
||||
|
||||
@@ -18,9 +18,9 @@ try:
|
||||
@torch.compiler.disable()
|
||||
def sageattn_func(q, k, v, attn_mask=None, dropout_p=0, is_causal=False):
|
||||
if q.dtype == torch.float32:
|
||||
return sageattn(q.to(torch.float16), k.to(torch.float16), v.to(torch.float16), attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal).to(torch.float32)
|
||||
return sageattn(q.to(torch.float16), k.to(torch.float16), v.to(torch.float16), attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout="NHD").to(torch.float32)
|
||||
else:
|
||||
return sageattn(q, k, v, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal)
|
||||
return sageattn(q, k, v, attn_mask=attn_mask, dropout_p=dropout_p, is_causal=is_causal, tensor_layout="NHD")
|
||||
except Exception as e:
|
||||
print(f"Warning: Could not load sageattention: {str(e)}")
|
||||
if isinstance(e, ModuleNotFoundError):
|
||||
@@ -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':
|
||||
return sageattn_func(q.transpose(1, 2), k.transpose(1, 2), v.transpose(1, 2)).transpose(1, 2).contiguous()
|
||||
else:
|
||||
return sageattn_func(q, k, v).contiguous()
|
||||
|
||||
@@ -15,6 +15,10 @@ try:
|
||||
except:
|
||||
BlockMask = create_block_mask = flex_attention = None
|
||||
pass
|
||||
try:
|
||||
from ..radial_attention.attn_mask import RadialSpargeSageAttn, RadialSpargeSageAttnDense
|
||||
except:
|
||||
pass
|
||||
|
||||
from .attention import attention
|
||||
import numpy as np
|
||||
@@ -259,6 +263,10 @@ class WanSelfAttention(nn.Module):
|
||||
self.eps = eps
|
||||
self.attention_mode = attention_mode
|
||||
|
||||
#radial attention
|
||||
self.mask_map = None
|
||||
self.decay_factor = 0.2
|
||||
|
||||
# layers
|
||||
self.q = nn.Linear(in_features, out_features)
|
||||
self.k = nn.Linear(in_features, out_features)
|
||||
@@ -319,6 +327,16 @@ class WanSelfAttention(nn.Module):
|
||||
|
||||
return x
|
||||
|
||||
def forward_radial(self, q, k, v, dense_step=False):
|
||||
if dense_step:
|
||||
x = RadialSpargeSageAttnDense(q, k, v, self.mask_map)
|
||||
else:
|
||||
x = RadialSpargeSageAttn(q, k, v, self.mask_map, decay_factor=self.decay_factor)
|
||||
|
||||
x = self.o(x.flatten(2))
|
||||
|
||||
return x
|
||||
|
||||
def forward_multitalk(self, q, k, v, seq_lens, grid_sizes, ref_target_masks):
|
||||
x = attention(
|
||||
q, k, v,
|
||||
@@ -560,7 +578,7 @@ class WanAttentionBlock(nn.Module):
|
||||
cross_attn_norm=False,
|
||||
eps=1e-6,
|
||||
attention_mode='sdpa',
|
||||
rope_func="comfy"
|
||||
rope_func="comfy",
|
||||
):
|
||||
super().__init__()
|
||||
self.dim = out_features
|
||||
@@ -571,6 +589,10 @@ class WanAttentionBlock(nn.Module):
|
||||
self.eps = eps
|
||||
self.attention_mode = attention_mode
|
||||
self.rope_func = rope_func
|
||||
#radial attn
|
||||
self.dense_timesteps = 10
|
||||
self.dense_block = False
|
||||
self.dense_attention_mode = "sageattn"
|
||||
|
||||
# layers
|
||||
self.norm1 = WanLayerNorm(out_features, eps)
|
||||
@@ -704,6 +726,14 @@ 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 or self.dense_timesteps is not None and current_step < self.dense_timesteps:
|
||||
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)
|
||||
|
||||
|
||||
@@ -0,0 +1,156 @@
|
||||
# based on https://github.com/mit-han-lab/radial-attention/blob/main/radial_attn/attn_mask.py
|
||||
import torch
|
||||
try:
|
||||
from sparse_sageattn import sparse_sageattn
|
||||
except:
|
||||
sparse_sageattn = None
|
||||
raise ImportError("Package is not installed: https://github.com/jt-zhang/Sparse_SageAttention_API")
|
||||
|
||||
from comfy import model_management as mm
|
||||
device = mm.get_torch_device()
|
||||
|
||||
from tqdm import tqdm
|
||||
|
||||
def shrinkMaskStrict(mask, block_size=128):
|
||||
seqlen = mask.shape[0]
|
||||
block_num = seqlen // block_size
|
||||
mask = mask[:block_num * block_size, :block_num * block_size].view(block_num, block_size, block_num, block_size)
|
||||
col_densities = mask.sum(dim = 1) / block_size
|
||||
# we want the minimum non-zero column density in the block
|
||||
non_zero_densities = col_densities > 0
|
||||
high_density_cols = col_densities > 1/3
|
||||
frac_high_density_cols = high_density_cols.sum(dim=-1) / (non_zero_densities.sum(dim=-1) + 1e-9)
|
||||
block_mask = frac_high_density_cols > 0.6
|
||||
block_mask[0:0] = True
|
||||
block_mask[-1:-1] = True
|
||||
return block_mask
|
||||
|
||||
def get_diagonal_split_mask(i, j, token_per_frame, sparse_type):
|
||||
assert(sparse_type in ["radial"])
|
||||
dist = abs(i - j)
|
||||
group = dist.bit_length()
|
||||
threshold = 128 # hardcoded threshold for now, which is equal to block-size
|
||||
decay_length = 2 ** token_per_frame.bit_length() / 2 ** group
|
||||
if decay_length >= threshold:
|
||||
return torch.ones((token_per_frame, token_per_frame), device=device, dtype=torch.bool)
|
||||
|
||||
split_factor = int(threshold / decay_length)
|
||||
modular = dist % split_factor
|
||||
if modular == 0:
|
||||
return torch.ones((token_per_frame, token_per_frame), device=device, dtype=torch.bool)
|
||||
else:
|
||||
return torch.zeros((token_per_frame, token_per_frame), device=device, dtype=torch.bool)
|
||||
|
||||
def get_window_width(i, j, token_per_frame, sparse_type, decay_factor=1, block_size=128):
|
||||
assert(sparse_type in ["radial"])
|
||||
dist = abs(i - j)
|
||||
if dist < 1:
|
||||
return token_per_frame
|
||||
if dist == 1:
|
||||
return token_per_frame // 2
|
||||
group = dist.bit_length()
|
||||
decay_length = 2 ** token_per_frame.bit_length() / 2 ** group * decay_factor
|
||||
threshold = block_size
|
||||
if decay_length >= threshold:
|
||||
return decay_length
|
||||
else:
|
||||
return threshold
|
||||
|
||||
def gen_log_mask_shrinked(device, s, video_token_num, num_frame, block_size=128, sparse_type="log", decay_factor=0.5):
|
||||
"""
|
||||
A more memory friendly version, we generate the attention mask of each frame pair at a time,
|
||||
shrinks it, and stores it into the final result
|
||||
"""
|
||||
|
||||
final_log_mask = torch.zeros((s // block_size, s // block_size), device=device, dtype=torch.bool)
|
||||
token_per_frame = video_token_num // num_frame
|
||||
video_text_border = video_token_num // block_size
|
||||
|
||||
col_indices = torch.arange(0, token_per_frame, device=device).view(1, -1)
|
||||
row_indices = torch.arange(0, token_per_frame, device=device).view(-1, 1)
|
||||
final_log_mask[video_text_border:] = True
|
||||
final_log_mask[:, video_text_border:] = True
|
||||
for i in tqdm(range(num_frame), desc="Frames (i)"):
|
||||
for j in range(num_frame):
|
||||
local_mask = torch.zeros((token_per_frame, token_per_frame), device=device, dtype=torch.bool)
|
||||
if j == 0: # this is attention sink
|
||||
local_mask = torch.ones((token_per_frame, token_per_frame), device=device, dtype=torch.bool)
|
||||
else:
|
||||
window_width = get_window_width(i, j, token_per_frame, sparse_type, decay_factor=decay_factor, block_size=block_size)
|
||||
local_mask = torch.abs(col_indices - row_indices) <= window_width
|
||||
split_mask = get_diagonal_split_mask(i, j, token_per_frame, sparse_type)
|
||||
local_mask = torch.logical_and(local_mask, split_mask)
|
||||
|
||||
remainder_row = (i * token_per_frame) % block_size
|
||||
remainder_col = (j * token_per_frame) % block_size
|
||||
# get the padded size
|
||||
all_length_row = remainder_row + ((token_per_frame - 1) // block_size + 1) * block_size
|
||||
all_length_col = remainder_col + ((token_per_frame - 1) // block_size + 1) * block_size
|
||||
padded_local_mask = torch.zeros((all_length_row, all_length_col), device=device, dtype=torch.bool)
|
||||
padded_local_mask[remainder_row:remainder_row + token_per_frame, remainder_col:remainder_col + token_per_frame] = local_mask
|
||||
# shrink the mask
|
||||
block_mask = shrinkMaskStrict(padded_local_mask, block_size=block_size)
|
||||
# set the block mask to the final log mask
|
||||
block_row_start = (i * token_per_frame) // block_size
|
||||
block_col_start = (j * token_per_frame) // block_size
|
||||
block_row_end = block_row_start + block_mask.shape[0]
|
||||
block_col_end = block_col_start + block_mask.shape[1]
|
||||
final_log_mask[block_row_start:block_row_end, block_col_start:block_col_end] = torch.logical_or(
|
||||
final_log_mask[block_row_start:block_row_end, block_col_start:block_col_end], block_mask)
|
||||
#print(f"mask sparsity: {1 - final_log_mask.sum() / final_log_mask.numel()}")
|
||||
return final_log_mask
|
||||
|
||||
class MaskMap:
|
||||
|
||||
def __init__(self, video_token_num=25440, num_frame=16):
|
||||
self.video_token_num = video_token_num
|
||||
self.num_frame = num_frame
|
||||
self.log_mask = None
|
||||
|
||||
def queryLogMask(self, seq_len, sparse_type, block_size=128, decay_factor=0.5):
|
||||
log_mask = torch.ones((seq_len // block_size, seq_len // block_size), device=device, dtype=torch.bool)
|
||||
if self.log_mask is None:
|
||||
self.log_mask = gen_log_mask_shrinked(device, seq_len, self.video_token_num, self.num_frame, sparse_type=sparse_type, decay_factor=decay_factor, block_size=block_size)
|
||||
block_bound = self.video_token_num // block_size
|
||||
log_mask[:block_bound, :block_bound] = self.log_mask[:block_bound, :block_bound]
|
||||
return log_mask
|
||||
|
||||
@torch.compiler.disable()
|
||||
def RadialSpargeSageAttnDense(query, key, value, mask_map=None):
|
||||
# dense case
|
||||
output_video = sparse_sageattn(
|
||||
query[:,:mask_map.video_token_num],
|
||||
key[:,:key.shape[1], :, :],
|
||||
value[:,:key.shape[1], :, :],
|
||||
mask_id=None,
|
||||
is_causal=False,
|
||||
tensor_layout="NHD",
|
||||
)
|
||||
|
||||
return output_video.contiguous()
|
||||
|
||||
@torch.compiler.disable()
|
||||
def RadialSpargeSageAttn(query, key, value, mask_map, block_size=128, decay_factor=1):
|
||||
# Simple cache based on function arguments
|
||||
if not hasattr(RadialSpargeSageAttn, "_cache"):
|
||||
RadialSpargeSageAttn._cache = {}
|
||||
cache_key = (query.shape[-2], block_size, decay_factor)
|
||||
if cache_key in RadialSpargeSageAttn._cache:
|
||||
input_mask = RadialSpargeSageAttn._cache[cache_key]
|
||||
else:
|
||||
print("Radial Attention: Generating block mask")
|
||||
video_mask = mask_map.queryLogMask(query.shape[0] * query.shape[1], "radial", block_size=block_size, decay_factor=decay_factor)
|
||||
mask = torch.repeat_interleave(video_mask, 2, dim=1) #s, t
|
||||
input_mask = mask.unsqueeze(0).unsqueeze(1).expand(1, query.shape[-2], mask.shape[0], mask.shape[1]) # b, h, s, t
|
||||
RadialSpargeSageAttn._cache[cache_key] = input_mask
|
||||
|
||||
output = sparse_sageattn(
|
||||
query[:, :, :mask_map.video_token_num, :],
|
||||
key[:, :, :mask_map.video_token_num, :],
|
||||
value[:, :, :mask_map.video_token_num, :],
|
||||
mask_id=input_mask.to(torch.int8),
|
||||
is_causal=False,
|
||||
tensor_layout="NHD"
|
||||
)
|
||||
|
||||
return output.contiguous()
|
||||
Reference in New Issue
Block a user