From 3567202e759b7a606b953c3f9767de59c695750e Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 27 Jul 2025 14:19:31 +0300 Subject: [PATCH] Fix compile and support radial attn in the DF sampler --- nodes.py | 86 ++++------------------------------------------ skyreels/nodes.py | 26 ++++++++++++-- utils.py | 87 ++++++++++++++++++++++++++++++++++++++++++++++- 3 files changed, 115 insertions(+), 84 deletions(-) diff --git a/nodes.py b/nodes.py index 647eee7..aa52baf 100644 --- a/nodes.py +++ b/nodes.py @@ -12,7 +12,7 @@ from .fp8_optimization import convert_linear_with_lora_and_scale, remove_lora_fr from .wanvideo.schedulers import get_scheduler, get_sampling_sigmas, retrieve_timesteps, scheduler_list from .gguf.gguf import set_lora_params 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, find_closest_valid_dim +from .utils import log, print_memory, apply_lora, clip_encode_image_tiled, fourier_filter, is_image_black, add_noise_to_reference_video, optimized_scale, setup_radial_attention, compile_model from .cache_methods.cache_methods import cache_report from .enhance_a_video.globals import set_enhance_weight, set_num_frames from .taehv import TAEHV @@ -1310,27 +1310,9 @@ class WanVideoSampler: log.info("Unloading all LoRAs") remove_lora_from_module(transformer) - #compile - compile_args = model["compile_args"] - if compile_args is not None and model["auto_cpu_offload"] is False: - torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"] - try: - if hasattr(torch, '_dynamo') and hasattr(torch._dynamo, 'config'): - torch._dynamo.config.recompile_limit = compile_args["dynamo_recompile_limit"] - except Exception as e: - log.warning(f"Could not set recompile_limit: {e}") - if compile_args["compile_transformer_blocks_only"]: - for i, block in enumerate(transformer.blocks): - if hasattr(block, "_orig_mod"): - block = block._orig_mod - transformer.blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) - if transformer.vace_layers is not None: - for i, block in enumerate(transformer.vace_blocks): - if hasattr(block, "_orig_mod"): - block = block._orig_mod - transformer.vace_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) - else: - transformer = torch.compile(transformer, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) + #torch.compile + if model["auto_cpu_offload"] is False: + transformer = compile_model(transformer, model["compile_args"]) multitalk_sampling = image_embeds.get("multitalk_sampling", False) if not multitalk_sampling and scheduler == "multitalk": @@ -1843,65 +1825,9 @@ class WanVideoSampler: else: transformer.slg_blocks = None - # Radial attention setup + # Setup radial attention if transformer.attention_mode == "radial_sage_attention": - dense_timesteps = transformer_options.get("dense_timesteps", None) - dense_blocks = transformer_options.get("dense_blocks", None) - dense_vace_blocks = transformer_options.get("dense_vace_blocks", None) - decay_factor = transformer_options.get("decay_factor", None) - 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: - 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 - 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 * 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 - for i, block in enumerate(transformer.blocks): - block.self_attn.mask_map = block.dense_attention_mode = block.dense_timesteps = block.self_attn.decay_factor = None - if isinstance(dense_blocks, list): - block.dense_block = i in dense_blocks - else: - block.dense_block = i < dense_blocks - block.self_attn.mask_map = MaskMap(video_token_num=seq_len, num_frame=latent_video_length if context_options is None else context_frames, block_size=block_size) - 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 - if isinstance(dense_vace_blocks, list): - block.dense_block = i in dense_vace_blocks - else: - block.dense_block = i < dense_vace_blocks - block.self_attn.mask_map = MaskMap(video_token_num=seq_len, num_frame=latent_video_length if context_options is None else context_frames, block_size=block_size) - 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.") - log.info(f"dense_attention_mode: {dense_attention_mode}, dense_timesteps: {dense_timesteps}, decay_factor: {decay_factor}") - log.info(f"dense_blocks: {[i for i, block in enumerate(transformer.blocks) if getattr(block, 'dense_block', False)]})") - + setup_radial_attention(transformer, transformer_options, latent, seq_len, latent_video_length, context_options=context_options) # FlowEdit setup if flowedit_args is not None: diff --git a/skyreels/nodes.py b/skyreels/nodes.py index dceb035..aee53ab 100644 --- a/skyreels/nodes.py +++ b/skyreels/nodes.py @@ -1,15 +1,16 @@ import os import torch import gc -from ..utils import log, print_memory, fourier_filter, optimized_scale +from ..utils import log, print_memory, fourier_filter, optimized_scale, setup_radial_attention, compile_model import math from tqdm import tqdm from ..wanvideo.modules.model import rope_params from ..wanvideo.schedulers.fm_solvers_unipc import FlowUniPCMultistepScheduler from diffusers.schedulers import FlowMatchEulerDiscreteScheduler +from ..fp8_optimization import convert_linear_with_lora_and_scale, remove_lora_from_module from ..wanvideo.schedulers.scheduling_flow_match_lcm import FlowMatchLCMScheduler - +from ..gguf.gguf import set_lora_params from einops import rearrange from ..enhance_a_video.globals import disable_enhance @@ -144,6 +145,23 @@ class WanVideoDiffusionForcingSampler: dtype = model["dtype"] device = mm.get_torch_device() offload_device = mm.unet_offload_device() + + gguf = model["gguf"] + transformer_options = patcher.model_options.get("transformer_options", None) + + if len(patcher.patches) != 0 and transformer_options.get("linear_with_lora", False) is True: + log.info(f"Using {len(patcher.patches)} LoRA weight patches for WanVideo model") + if not gguf: + convert_linear_with_lora_and_scale(transformer, patches=patcher.patches) + else: + set_lora_params(transformer, patcher.patches) + else: + log.info("Unloading all LoRAs") + remove_lora_from_module(transformer) + + #torch.compile + if model["auto_cpu_offload"] is False: + transformer = compile_model(transformer, model["compile_args"]) steps = int(steps/denoise_strength) @@ -343,7 +361,6 @@ class WanVideoDiffusionForcingSampler: callback = prepare_callback(patcher, steps) #blockswap init - transformer_options = patcher.model_options.get("transformer_options", None) if transformer_options is not None: block_swap_args = transformer_options.get("block_swap_args", None) @@ -408,6 +425,9 @@ class WanVideoDiffusionForcingSampler: self.teacache_state_source = [None, None] self.teacache_states_context = [] + if transformer.attention_mode == "radial_sage_attention": + setup_radial_attention(transformer, transformer_options, latents, seq_len, latent_video_length) + use_cfg_zero_star, use_fresca = False, False if experimental_args is not None: diff --git a/utils.py b/utils.py index 6fbc5c0..71721fb 100644 --- a/utils.py +++ b/utils.py @@ -308,4 +308,89 @@ def find_closest_valid_dim(fixed_dim, var_dim, block_size): 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 + return var_dim + + # Radial attention setup +def setup_radial_attention(transformer, transformer_options, latent, seq_len, latent_video_length, context_options=None): + if context_options is not None: + context_frames = (context_options["context_frames"] - 1) // 4 + 1 + + dense_timesteps = transformer_options.get("dense_timesteps", 1) + dense_blocks = transformer_options.get("dense_blocks", 1) + dense_vace_blocks = transformer_options.get("dense_vace_blocks", 1) + decay_factor = transformer_options.get("decay_factor", 0.2) + dense_attention_mode = transformer_options.get("dense_attention_mode", "sageattn") + block_size = transformer_options.get("block_size", 128) + + # Calculate closest valid latent sizes + if latent.shape[2] % (block_size/8) != 0 or latent.shape[3] % (block_size/8) != 0: + 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 + 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 * 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 + for i, block in enumerate(transformer.blocks): + block.self_attn.mask_map = block.dense_attention_mode = block.dense_timesteps = block.self_attn.decay_factor = None + if isinstance(dense_blocks, list): + block.dense_block = i in dense_blocks + else: + block.dense_block = i < dense_blocks + block.self_attn.mask_map = MaskMap(video_token_num=seq_len, num_frame=latent_video_length if context_options is None else context_frames, block_size=block_size) + 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 + if isinstance(dense_vace_blocks, list): + block.dense_block = i in dense_vace_blocks + else: + block.dense_block = i < dense_vace_blocks + block.self_attn.mask_map = MaskMap(video_token_num=seq_len, num_frame=latent_video_length if context_options is None else context_frames, block_size=block_size) + 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.") + log.info(f"dense_attention_mode: {dense_attention_mode}, dense_timesteps: {dense_timesteps}, decay_factor: {decay_factor}") + log.info(f"dense_blocks: {[i for i, block in enumerate(transformer.blocks) if getattr(block, 'dense_block', False)]})") + +def compile_model(transformer, compile_args=None): + if compile_args is None: + return transformer + torch._dynamo.config.cache_size_limit = compile_args["dynamo_cache_size_limit"] + try: + if hasattr(torch, '_dynamo') and hasattr(torch._dynamo, 'config'): + torch._dynamo.config.recompile_limit = compile_args["dynamo_recompile_limit"] + except Exception as e: + log.warning(f"Could not set recompile_limit: {e}") + if compile_args["compile_transformer_blocks_only"]: + for i, block in enumerate(transformer.blocks): + if hasattr(block, "_orig_mod"): + block = block._orig_mod + transformer.blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) + if transformer.vace_layers is not None: + for i, block in enumerate(transformer.vace_blocks): + if hasattr(block, "_orig_mod"): + block = block._orig_mod + transformer.vace_blocks[i] = torch.compile(block, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) + else: + transformer = torch.compile(transformer, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"]) + return transformer \ No newline at end of file