Fix compile and support radial attn in the DF sampler

This commit is contained in:
kijai
2025-07-27 14:19:31 +03:00
parent c0a4e8aa87
commit 3567202e75
3 changed files with 115 additions and 84 deletions
+6 -80
View File
@@ -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:
+23 -3
View File
@@ -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:
+86 -1
View File
@@ -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
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