diff --git a/nodes.py b/nodes.py index 3f79b19..df96b06 100644 --- a/nodes.py +++ b/nodes.py @@ -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}") diff --git a/wanvideo/radial_attention/attn_mask.py b/wanvideo/radial_attention/attn_mask.py index 6b95cec..4e13674 100644 --- a/wanvideo/radial_attention/attn_mask.py +++ b/wanvideo/radial_attention/attn_mask.py @@ -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