From 768a4b2d90b8caba84cab73ed768dfcd12e70bc6 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 20 Jul 2025 12:15:01 +0300 Subject: [PATCH] Fix radial attention cache not resetting on resolution change, add more robust supported dimension check --- nodes.py | 20 +++++++++++++++----- utils.py | 8 ++++++++ wanvideo/radial_attention/attn_mask.py | 8 +++++++- 3 files changed, 30 insertions(+), 6 deletions(-) diff --git a/nodes.py b/nodes.py index f6fa21b..2a1e324 100644 --- a/nodes.py +++ b/nodes.py @@ -12,7 +12,7 @@ from .wanvideo.modules.model import rope_params from .wanvideo.schedulers import get_scheduler, get_sampling_sigmas, retrieve_timesteps, scheduler_list from .multitalk.multitalk import timestep_transform, add_noise -from .utils import log, print_memory, apply_lora, clip_encode_image_tiled, fourier_filter, is_image_black, add_noise_to_reference_video, optimized_scale +from .utils import log, print_memory, apply_lora, clip_encode_image_tiled, fourier_filter, is_image_black, add_noise_to_reference_video, optimized_scale, find_closest_valid_dim from .cache_methods.cache_methods import cache_report from .enhance_a_video.globals import set_enhance_weight, set_num_frames from .taehv import TAEHV @@ -1645,17 +1645,27 @@ class WanVideoSampler: dense_attention_mode = transformer_options.get("dense_attention_mode", None) block_size = transformer_options.get("block_size", None) + # Calculate closest valid latent sizes 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)." + f"Closest valid sizes: {closest_w * 8}x{closest_h * 8} (width x height in pixels)." + ) + tokens_per_frame = (latent.shape[2] * latent.shape[3]) // 4 + if tokens_per_frame % block_size != 0: + closest_latent_h = find_closest_valid_dim(latent.shape[3], latent.shape[2], block_size) + closest_latent_w = find_closest_valid_dim(latent.shape[2], latent.shape[3], block_size) + raise Exception( + f"Radial attention mode requires tokens per frame ((latent_h * latent_w) // 4) to be divisible by block size ({block_size}).\n" + f"Current size in latent space:{latent.shape[3]}x{latent.shape[2]}, pixel space: {latent.shape[3]*8}x{latent.shape[2]*8} tokens_per_frame={tokens_per_frame}.\n" + f"Try adjusting to one of these latent sizes (in pixels):\n" + f" Height: {latent.shape[2]*8} -> {closest_latent_h * 8}\n" + f" Width: {latent.shape[3]*8} -> {closest_latent_w * 8}\n" + f"Or choose another resolution so that (latent_h * latent_w) // 4 is divisible by {block_size}." ) from .wanvideo.radial_attention.attn_mask import MaskMap diff --git a/utils.py b/utils.py index daca9d7..6f4e273 100644 --- a/utils.py +++ b/utils.py @@ -261,3 +261,11 @@ def optimized_scale(positive_flat, negative_flat): st_star = dot_product / squared_norm return st_star + +def find_closest_valid_dim(fixed_dim, var_dim, block_size): + for delta in range(1, 17): + for sign in [-1, 1]: + candidate = var_dim + sign * delta + if candidate > 0 and ((fixed_dim * candidate) // 4) % block_size == 0: + return candidate + return var_dim \ No newline at end of file diff --git a/wanvideo/radial_attention/attn_mask.py b/wanvideo/radial_attention/attn_mask.py index 830b1a4..747a611 100644 --- a/wanvideo/radial_attention/attn_mask.py +++ b/wanvideo/radial_attention/attn_mask.py @@ -136,7 +136,13 @@ def RadialSpargeSageAttn(query, key, value, mask_map, decay_factor): RadialSpargeSageAttn._cache = {} # print(mask_map.block_size) block_size = mask_map.block_size - cache_key = (query.shape[-2], mask_map.block_size, decay_factor) + cache_key = ( + query.shape[-2], + mask_map.block_size, + decay_factor, + mask_map.video_token_num, + mask_map.num_frame + ) if cache_key in RadialSpargeSageAttn._cache: input_mask = RadialSpargeSageAttn._cache[cache_key] else: