Add TSR (Temporal Score Rescaling) to experimental args

https://github.com/temporalscorerescaling/TSR
This commit is contained in:
kijai
2025-10-05 01:14:25 +03:00
parent 71c8a4961b
commit d380c7d219
3 changed files with 30 additions and 3 deletions
+4 -1
View File
@@ -1738,7 +1738,10 @@ class WanVideoExperimentalArgs:
"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."}),
"bidirectional_sampling": ("BOOLEAN", {"default": False, "tooltip": "Enable bidirectional sampling, based on https://github.com/ff2416/WanFM"})
"bidirectional_sampling": ("BOOLEAN", {"default": False, "tooltip": "Enable bidirectional sampling, based on https://github.com/ff2416/WanFM"}),
"temporal_score_rescaling": ("BOOLEAN", {"default": False, "tooltip": "Enable temporal score rescaling: https://github.com/temporalscorerescaling/TSR/"}),
"tsr_k": ("FLOAT", {"default": 1.0, "min": 0.0, "max": 10.0, "step": 0.01, "tooltip": "The sampling temperature"}),
"tsr_sigma": ("FLOAT", {"default": 0.1, "min": 0.0, "max": 1.0, "step": 0.01, "tooltip": "How early TSR steer the sampling process"}),
},
}
+13 -2
View File
@@ -12,7 +12,7 @@ from .wanvideo.schedulers import get_scheduler, get_sampling_sigmas, retrieve_ti
from .gguf.gguf import set_lora_params_gguf
from .multitalk.multitalk import timestep_transform, add_noise
from .utils import(log, print_memory, apply_lora, clip_encode_image_tiled, fourier_filter, optimized_scale, setup_radial_attention,
compile_model, dict_to_device, tangential_projection, set_module_tensor_to_device, get_raag_guidance)
compile_model, dict_to_device, tangential_projection, set_module_tensor_to_device, get_raag_guidance, temporal_score_rescaling)
from .cache_methods.cache_methods import cache_report
from .nodes_model_loading import load_weights
from .enhance_a_video.globals import set_enhance_weight, set_num_frames
@@ -933,7 +933,7 @@ class WanVideoSampler:
timesteps[-drift_steps:] = drift_timesteps[-drift_steps:]
# Experimental args
use_cfg_zero_star = use_tangential = use_fresca = bidirectional_sampling =False
use_cfg_zero_star = use_tangential = use_fresca = bidirectional_sampling = use_tsr = False
raag_alpha = 0.0
if experimental_args is not None:
video_attention_split_steps = experimental_args.get("video_attention_split_steps", [])
@@ -957,6 +957,9 @@ class WanVideoSampler:
bidirectional_sampling = experimental_args.get("bidirectional_sampling", False)
if bidirectional_sampling:
sample_scheduler_flipped = copy.deepcopy(sample_scheduler)
use_tsr = experimental_args.get("temporal_score_rescaling", False)
tsr_k = experimental_args.get("tsr_k", 1.0)
tsr_sigma = experimental_args.get("tsr_sigma", 1.0)
# Rotary positional embeddings (RoPE)
@@ -2157,6 +2160,8 @@ class WanVideoSampler:
step_iteration_count += 1
# update latent
if use_tsr:
noise_pred = temporal_score_rescaling(noise_pred, latent, timestep, tsr_k, tsr_sigma)
if scheduler == "multitalk":
noise_pred = -noise_pred
dt = (timesteps[i] - timesteps[i + 1]) / 1000
@@ -2627,6 +2632,9 @@ class WanVideoSampler:
sampling_pbar.update(1)
step_iteration_count += 1
if use_tsr:
noise_pred = temporal_score_rescaling(noise_pred, latent, timestep, tsr_k, tsr_sigma)
latent = sample_scheduler.step(noise_pred.unsqueeze(0), timestep, latent.unsqueeze(0).to(noise_pred.device), **scheduler_step_args)[0].squeeze(0)
del noise_pred, latent_model_input, timestep
@@ -2731,6 +2739,9 @@ class WanVideoSampler:
if flowedit_args is None:
latent = latent.to(intermediate_device)
if use_tsr:
noise_pred = temporal_score_rescaling(noise_pred, latent, timestep, tsr_k, tsr_sigma)
if len(timestep.shape) != 1 and not is_pusa: #5b
# all_indices is a list of indices to skip
total_indices = list(range(latent.shape[1]))
+13
View File
@@ -588,3 +588,16 @@ def check_duplicate_nodes():
wanvideo_dirs.append(str(path))
return wanvideo_dirs
#https://github.com/temporalscorerescaling/TSR/
def temporal_score_rescaling(model_output, sample, timestep, k=1.0, tsr_sigma=0.1):
t = (timestep / 1000)
if t == 0.0:
ratio = k
else:
snr_t = (1 - t)**2 / t**2
ratio = (snr_t * tsr_sigma**2 + 1) / (snr_t * tsr_sigma**2 / k + 1)
if not t == 1.0:
model_output = (ratio * ((1-t) * model_output + sample) - sample) / (1 - t)
return model_output