Experimental RAAG (Ratio Aware Adaptive Guidance) implementation
https://arxiv.org/abs/2508.03442 Available in experimental_args, alpha > 0 = enabled, default value is 1.0
This commit is contained in:
@@ -2,6 +2,7 @@
|
||||
|
||||
import torch
|
||||
import torch.nn as nn
|
||||
from .utils import log
|
||||
|
||||
def fp8_linear_forward(cls, original_dtype, input):
|
||||
weight_dtype = cls.weight.dtype
|
||||
@@ -120,6 +121,10 @@ def convert_linear_with_lora_and_scale(module, scale_weight_keys=None, patches=N
|
||||
|
||||
|
||||
def remove_lora_from_module(module):
|
||||
unloaded = False
|
||||
for name, submodule in module.named_modules():
|
||||
if hasattr(submodule, "lora"):
|
||||
if not unloaded:
|
||||
log.info("Unloading all LoRAs")
|
||||
unloaded = True
|
||||
delattr(submodule, "lora")
|
||||
|
||||
@@ -14,7 +14,7 @@ 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,
|
||||
add_noise_to_reference_video, optimized_scale, setup_radial_attention,
|
||||
compile_model, dict_to_device, tangential_projection, set_module_tensor_to_device)
|
||||
compile_model, dict_to_device, tangential_projection, set_module_tensor_to_device, get_raag_guidance)
|
||||
from .cache_methods.cache_methods import cache_report
|
||||
from .enhance_a_video.globals import set_enhance_weight, set_num_frames
|
||||
from .taehv import TAEHV
|
||||
@@ -1403,6 +1403,7 @@ class WanVideoExperimentalArgs:
|
||||
"fresca_scale_high": ("FLOAT", {"default": 1.25, "min": 0.0, "max": 10.0, "step": 0.01}),
|
||||
"fresca_freq_cutoff": ("INT", {"default": 20, "min": 0, "max": 10000, "step": 1}),
|
||||
"use_tcfg": ("BOOLEAN", {"default": False, "tooltip": "https://arxiv.org/abs/2503.18137 TCFG: Tangential Damping Classifier-free Guidance. CFG artifacts reduction."}),
|
||||
"raag_alpha": ("FLOAT", {"default": 0.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "Alpha value for RAAG, 1.0 is default, 0.0 is disabled."}),
|
||||
},
|
||||
}
|
||||
|
||||
@@ -1508,7 +1509,6 @@ class WanVideoSampler:
|
||||
else:
|
||||
set_lora_params(transformer, patcher.patches)
|
||||
else:
|
||||
log.info("Unloading all LoRAs")
|
||||
remove_lora_from_module(transformer)
|
||||
|
||||
transformer.lora_scheduling_enabled = transformer_options.get("lora_scheduling_enabled", False)
|
||||
@@ -2098,6 +2098,7 @@ class WanVideoSampler:
|
||||
|
||||
# Experimental args
|
||||
use_cfg_zero_star = use_tangential = use_fresca = False
|
||||
raag_alpha = 0.0
|
||||
if experimental_args is not None:
|
||||
video_attention_split_steps = experimental_args.get("video_attention_split_steps", [])
|
||||
if video_attention_split_steps:
|
||||
@@ -2109,6 +2110,7 @@ class WanVideoSampler:
|
||||
use_cfg_zero_star = experimental_args.get("cfg_zero_star", False)
|
||||
use_tangential = experimental_args.get("use_tcfg", False)
|
||||
zero_star_steps = experimental_args.get("zero_star_steps", 0)
|
||||
raag_alpha = experimental_args.get("raag_alpha", 0.0)
|
||||
|
||||
use_fresca = experimental_args.get("use_fresca", False)
|
||||
if use_fresca:
|
||||
@@ -2407,6 +2409,11 @@ class WanVideoSampler:
|
||||
if use_tangential:
|
||||
noise_pred_uncond_scaled = tangential_projection(noise_pred_cond, noise_pred_uncond_scaled)
|
||||
|
||||
# RAAG (RATIO-aware Adaptive Guidance)
|
||||
if raag_alpha > 0.0:
|
||||
cfg_scale = get_raag_guidance(noise_pred_cond, noise_pred_uncond_scaled, cfg_scale, raag_alpha)
|
||||
log.info(f"RAAG modified cfg: {cfg_scale}")
|
||||
|
||||
#https://github.com/WikiChao/FreSca
|
||||
if use_fresca:
|
||||
filtered_cond = fourier_filter(
|
||||
|
||||
@@ -1,6 +1,7 @@
|
||||
import importlib.metadata
|
||||
import torch
|
||||
import logging
|
||||
import math
|
||||
from tqdm import tqdm
|
||||
import types, collections
|
||||
from comfy.utils import ProgressBar, copy_to_param, set_attr_param
|
||||
@@ -501,6 +502,7 @@ def compile_model(transformer, compile_args=None):
|
||||
transformer = torch.compile(transformer, fullgraph=compile_args["fullgraph"], dynamic=compile_args["dynamic"], backend=compile_args["backend"], mode=compile_args["mode"])
|
||||
return transformer
|
||||
|
||||
#https://5410tiffany.github.io/tcfg.github.io/
|
||||
def tangential_projection(pred_cond: torch.Tensor, pred_uncond: torch.Tensor) -> torch.Tensor:
|
||||
cond_dtype = pred_cond.dtype
|
||||
preds = torch.stack([pred_cond, pred_uncond], dim=1).float()
|
||||
@@ -511,3 +513,13 @@ def tangential_projection(pred_cond: torch.Tensor, pred_uncond: torch.Tensor) ->
|
||||
Vh_modified[:, 1] = 0
|
||||
recon = U @ torch.diag_embed(S) @ Vh_modified
|
||||
return recon[:, 1].view(pred_uncond.shape).to(cond_dtype)
|
||||
|
||||
#https://arxiv.org/abs/2508.03442
|
||||
def get_raag_guidance(noise_pred_cond, noise_pred_uncond, w_max, alpha=1.0, eps=1e-8):
|
||||
delta = noise_pred_cond - noise_pred_uncond
|
||||
norm_delta = torch.norm(delta.flatten(1), dim=1, keepdim=True)
|
||||
norm_uncond = torch.norm(noise_pred_uncond.flatten(1), dim=1, keepdim=True)
|
||||
ratio = norm_delta / (norm_uncond + eps)
|
||||
ratio_mean = ratio.mean().item()
|
||||
adaptive_w = 1.0 + (w_max - 1.0) * math.exp(-alpha * ratio_mean)
|
||||
return adaptive_w
|
||||
Reference in New Issue
Block a user