force using the set node

This commit is contained in:
kijai
2025-07-18 14:26:45 +03:00
parent 0e920e5b18
commit a1cc320c36
2 changed files with 16 additions and 13 deletions
+15 -12
View File
@@ -2540,11 +2540,13 @@ class WanVideoSampler:
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.")
if transformer_options is not None:
dense_timesteps = transformer_options.get("dense_timesteps", 10)
dense_blocks = transformer_options.get("dense_blocks", 1)
decay_factor = transformer_options.get("decay_factor", 0.2)
dense_attention_mode = transformer_options.get("dense_attention_mode", "sageattn")
dense_timesteps = transformer_options.get("dense_timesteps", None)
dense_blocks = transformer_options.get("dense_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):
@@ -2554,13 +2556,14 @@ class WanVideoSampler:
block.dense_attention_mode = dense_attention_mode
block.dense_timesteps = dense_timesteps
block.self_attn.decay_factor = decay_factor
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_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_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}")
+1 -1
View File
@@ -138,7 +138,7 @@ def RadialSpargeSageAttn(query, key, value, mask_map, block_size=128, decay_fact
if cache_key in RadialSpargeSageAttn._cache:
input_mask = RadialSpargeSageAttn._cache[cache_key]
else:
print("generating input mask")
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