Merge branch 'pr/828'

This commit is contained in:
kijai
2025-07-20 02:13:26 +03:00
2 changed files with 69 additions and 55 deletions
+19 -8
View File
@@ -92,6 +92,7 @@ class WanVideoSetRadialAttention:
"dense_vace_blocks": ("INT", {"default": 1, "min": 0, "max": 15, "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."}),
"block_size":([128, 64], {"default": 128, "tooltip": "Radial attention block size"}),
}
}
@@ -101,7 +102,7 @@ class WanVideoSetRadialAttention:
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):
def loadmodel(self, model, dense_attention_mode, dense_blocks, dense_vace_blocks, dense_timesteps, decay_factor, block_size):
if "radial" not in model.model.diffusion_model.attention_mode:
raise Exception("Enable radial attention first in the model loader.")
@@ -114,6 +115,7 @@ class WanVideoSetRadialAttention:
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
patcher.model_options["transformer_options"]["block_size"] = block_size
return (patcher,)
@@ -1636,22 +1638,31 @@ class WanVideoSampler:
# 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.")
block_size = transformer_options.get("block_size", None)
if latent.shape[2] % (block_size/8) != 0 or latent.shape[3] % (block_size/8) != 0:
# Calculate closest valid latent sizes
block_div = int(block_size // 8)
closest_h = round(latent.shape[2] / block_div) * block_div
closest_w = round(latent.shape[3] / block_div) * block_div
closest_h_px = closest_h * 8
closest_w_px = closest_w * 8
raise Exception(
f"Radial attention mode only supports image size divisible by block size. "
f"Got {latent.shape[3] * 8}x{latent.shape[2] * 8} with block size {block_size}.\n"
f"Closest valid sizes: {closest_w_px}x{closest_h_px} (width x height in pixels)."
)
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.self_attn.mask_map = MaskMap(video_token_num=seq_len, num_frame=latent_video_length, block_size=block_size)
block.dense_attention_mode = dense_attention_mode
block.dense_timesteps = dense_timesteps
block.self_attn.decay_factor = decay_factor
@@ -1659,7 +1670,7 @@ class WanVideoSampler:
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.self_attn.mask_map = MaskMap(video_token_num=seq_len, num_frame=latent_video_length, block_size=block_size)
block.dense_attention_mode = dense_attention_mode
block.dense_timesteps = dense_timesteps
block.self_attn.decay_factor = decay_factor
+50 -47
View File
@@ -11,14 +11,13 @@ except:
from comfy import model_management as mm
device = mm.get_torch_device()
from tqdm import tqdm
def shrinkMaskStrict(mask, block_size=128):
def shrinkMaskStrict(mask, block_size):
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
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
@@ -28,24 +27,22 @@ def shrinkMaskStrict(mask, block_size=128):
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"])
def get_diagonal_split_mask(i, j, token_per_frame, sparse_type, block_size):
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
threshold = block_size # CHANGE, can 64 or 128
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)
return torch.ones((token_per_frame, token_per_frame), device=device, dtype=torch.bool) if modular == 0 \
else 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"])
def get_window_width(i, j, token_per_frame, sparse_type, decay_factor, block_size):
assert sparse_type in ["radial"]
dist = abs(i - j)
if dist < 1:
return token_per_frame
@@ -53,18 +50,13 @@ def get_window_width(i, j, token_per_frame, sparse_type, decay_factor=1, block_s
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
return max(decay_length, block_size)
def gen_log_mask_shrinked(device, s, video_token_num, num_frame, block_size=128, sparse_type="log", decay_factor=0.5):
def gen_log_mask_shrinked(device, s, video_token_num, num_frame, block_size, sparse_type, decay_factor):
"""
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
@@ -73,87 +65,98 @@ def gen_log_mask_shrinked(device, s, video_token_num, num_frame, block_size=128,
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)
window_width = get_window_width(i, j, token_per_frame, sparse_type, decay_factor, 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)
split_mask = get_diagonal_split_mask(i, j, token_per_frame, sparse_type, block_size)
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)
block_mask = shrinkMaskStrict(padded_local_mask, 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):
def __init__(self, video_token_num=25440, num_frame=16, block_size=128):
self.video_token_num = video_token_num
self.num_frame = num_frame
self.log_mask = None
self.block_size = block_size
def queryLogMask(self, seq_len, sparse_type, block_size=128, decay_factor=0.5):
def queryLogMask(self, seq_len, sparse_type, block_size=None, decay_factor=0.5):
block_size = block_size or self.block_size
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)
self.log_mask = gen_log_mask_shrinked(
device, seq_len, self.video_token_num, self.num_frame,
block_size=block_size, sparse_type=sparse_type, decay_factor=decay_factor
)
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):
def RadialSpargeSageAttnDense(query, key, value, mask_map):
# dense case
output_video = sparse_sageattn(
query[:,:mask_map.video_token_num],
key[:,:key.shape[1], :, :],
value[:,:key.shape[1], :, :],
return 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()
tensor_layout="NHD"
).contiguous()
@torch.compiler.disable()
def RadialSpargeSageAttn(query, key, value, mask_map, block_size=128, decay_factor=1):
def RadialSpargeSageAttn(query, key, value, mask_map, decay_factor):
# Simple cache based on function arguments
if not hasattr(RadialSpargeSageAttn, "_cache"):
RadialSpargeSageAttn._cache = {}
cache_key = (query.shape[-2], block_size, decay_factor)
# print(mask_map.block_size)
block_size = mask_map.block_size
cache_key = (query.shape[-2], mask_map.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
video_mask = mask_map.queryLogMask(query.shape[0] * query.shape[1], "radial", block_size=block_size, decay_factor=decay_factor)
# based on https://github.com/mit-han-lab/radial-attention/blob/3ec33ce9633adadadcbb7692c8a1983d5e82d15a/radial_attn/attn_mask.py#L7
if block_size == 128:
mask = torch.repeat_interleave(video_mask, 2, dim=1)
elif block_size == 64:
reshaped_mask = video_mask.view(video_mask.shape[0] // 2, 2, video_mask.shape[1])
mask = torch.max(reshaped_mask, dim=1).values
input_mask = mask.unsqueeze(0).unsqueeze(1).expand(1, query.shape[-2], mask.shape[0], mask.shape[1])
RadialSpargeSageAttn._cache[cache_key] = input_mask
output = sparse_sageattn(
return 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()
).contiguous()