From d380c7d2194786c50983ce1d829d4a090fa86200 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 5 Oct 2025 01:14:25 +0300 Subject: [PATCH] Add TSR (Temporal Score Rescaling) to experimental args https://github.com/temporalscorerescaling/TSR --- nodes.py | 5 ++++- nodes_sampler.py | 15 +++++++++++++-- utils.py | 13 +++++++++++++ 3 files changed, 30 insertions(+), 3 deletions(-) diff --git a/nodes.py b/nodes.py index 27131bc..cbab84f 100644 --- a/nodes.py +++ b/nodes.py @@ -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"}), }, } diff --git a/nodes_sampler.py b/nodes_sampler.py index ab3e6e0..d1b7412 100644 --- a/nodes_sampler.py +++ b/nodes_sampler.py @@ -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])) diff --git a/utils.py b/utils.py index e7069a2..1e97fd4 100644 --- a/utils.py +++ b/utils.py @@ -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