Fix radial attention cache not resetting on resolution change, add more robust supported dimension check

This commit is contained in:
kijai
2025-07-20 12:15:01 +03:00
parent 9f75b7e051
commit 768a4b2d90
3 changed files with 30 additions and 6 deletions
+15 -5
View File
@@ -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
+8
View File
@@ -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
+7 -1
View File
@@ -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: