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:
kijai
2025-08-08 00:17:32 +03:00
parent ab06ac2f64
commit d446e97309
3 changed files with 26 additions and 2 deletions
+5
View File
@@ -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")
+9 -2
View File
@@ -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(
+12
View File
@@ -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