Fix compile and support radial attn in the DF sampler
This commit is contained in:
@@ -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
@@ -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:
|
||||
|
||||
@@ -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
|
||||
Reference in New Issue
Block a user