Add TSR (Temporal Score Rescaling) to experimental args
https://github.com/temporalscorerescaling/TSR
This commit is contained in:
@@ -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
@@ -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]))
|
||||
|
||||
@@ -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
|
||||
|
||||
Reference in New Issue
Block a user