Fix radial attention cache not resetting on resolution change, add more robust supported dimension check
This commit is contained in:
@@ -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
|
||||
|
||||
@@ -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
|
||||
@@ -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:
|
||||
|
||||
Reference in New Issue
Block a user