From d446e97309824fa936fec2b52efeb19adb9ba38f Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Fri, 8 Aug 2025 00:17:32 +0300 Subject: [PATCH] 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 --- fp8_optimization.py | 5 +++++ nodes.py | 11 +++++++++-- utils.py | 12 ++++++++++++ 3 files changed, 26 insertions(+), 2 deletions(-) diff --git a/fp8_optimization.py b/fp8_optimization.py index 0f823be..b63e91d 100644 --- a/fp8_optimization.py +++ b/fp8_optimization.py @@ -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") diff --git a/nodes.py b/nodes.py index 233338f..e333c88 100644 --- a/nodes.py +++ b/nodes.py @@ -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( diff --git a/utils.py b/utils.py index d2cc243..540cf27 100644 --- a/utils.py +++ b/utils.py @@ -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 \ No newline at end of file