From edd9b206910c8832f27f4850c96234daca909a9a Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Sun, 20 Jul 2025 01:41:17 +0300 Subject: [PATCH] Add res_multistep --- nodes.py | 7 +- wanvideo/schedulers/__init__.py | 37 ++++-- .../schedulers/flowmatch_res_multistep.py | 105 ++++++++++++++++++ 3 files changed, 135 insertions(+), 14 deletions(-) create mode 100644 wanvideo/schedulers/flowmatch_res_multistep.py diff --git a/nodes.py b/nodes.py index 18b59c6..be7b510 100644 --- a/nodes.py +++ b/nodes.py @@ -9,7 +9,7 @@ from diffusers.schedulers import FlowMatchEulerDiscreteScheduler from .wanvideo.modules.model import rope_params -from .wanvideo.schedulers import get_scheduler, get_sampling_sigmas, retrieve_timesteps +from .wanvideo.schedulers import get_scheduler, get_sampling_sigmas, retrieve_timesteps, scheduler_list from .multitalk.multitalk import timestep_transform, add_noise from .utils import log, print_memory, apply_lora, clip_encode_image_tiled, fourier_filter, is_image_black, add_noise_to_reference_video, optimized_scale @@ -1097,10 +1097,7 @@ class WanVideoSampler: "shift": ("FLOAT", {"default": 5.0, "min": 0.0, "max": 1000.0, "step": 0.01}), "seed": ("INT", {"default": 0, "min": 0, "max": 0xffffffffffffffff}), "force_offload": ("BOOLEAN", {"default": True, "tooltip": "Moves the model to the offload device after sampling"}), - "scheduler": (["unipc", "unipc/beta", "dpm++", "dpm++/beta","dpm++_sde", "dpm++_sde/beta", "euler", "euler/beta", "euler/accvideo", "deis", "lcm", "lcm/beta", "flowmatch_causvid", "flowmatch_distill", "flowmatch_pusa", "multitalk"], - { - "default": 'unipc' - }), + "scheduler": (scheduler_list, {"default": "uni_pc",}), "riflex_freq_index": ("INT", {"default": 0, "min": 0, "max": 1000, "step": 1, "tooltip": "Frequency index for RIFLEX, disabled when 0, default 6. Allows for new frames to be generated after without looping"}), }, "optional": { diff --git a/wanvideo/schedulers/__init__.py b/wanvideo/schedulers/__init__.py index 1188b96..48cd37d 100644 --- a/wanvideo/schedulers/__init__.py +++ b/wanvideo/schedulers/__init__.py @@ -3,11 +3,27 @@ from .fm_solvers import (FlowDPMSolverMultistepScheduler, get_sampling_sigmas, r from .fm_solvers_unipc import FlowUniPCMultistepScheduler from .basic_flowmatch import FlowMatchScheduler from .flowmatch_pusa import FlowMatchSchedulerPusa +from .flowmatch_res_multistep import FlowMatchSchedulerResMultistep from .scheduling_flow_match_lcm import FlowMatchLCMScheduler from diffusers.schedulers import FlowMatchEulerDiscreteScheduler, DEISMultistepScheduler from ...utils import log +scheduler_list = [ + "unipc", "unipc/beta", + "dpm++", "dpm++/beta", + "dpm++_sde", "dpm++_sde/beta", + "euler", "euler/beta", + #"euler/accvideo", + "deis", + "lcm", "lcm/beta", + "res_multistep", + "flowmatch_causvid", + "flowmatch_distill", + "flowmatch_pusa", + "multitalk" +] + def get_scheduler(scheduler, steps, shift, device, transformer_dim, flowedit_args, denoise_strength, sigmas=None): timesteps = None if 'unipc' in scheduler: @@ -25,15 +41,15 @@ def get_scheduler(scheduler, steps, shift, device, transformer_dim, flowedit_arg timesteps, _ = retrieve_timesteps(sample_scheduler, device=device, sigmas=get_sampling_sigmas(steps, shift)) else: sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None) - elif scheduler in ['euler/accvideo']: - if steps != 50: - raise Exception("Steps must be set to 50 for accvideo scheduler, 10 actual steps are used") - sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, use_beta_sigmas=(scheduler == 'euler/beta')) - sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None) - start_latent_list = [0, 5, 10, 15, 20, 25, 30, 35, 40, 45, 50] - sample_scheduler.sigmas = sample_scheduler.sigmas[start_latent_list] - steps = len(start_latent_list) - 1 - sample_scheduler.timesteps = timesteps = sample_scheduler.timesteps[start_latent_list[:steps]] + # elif scheduler in ['euler/accvideo']: + # if steps != 50: + # raise Exception("Steps must be set to 50 for accvideo scheduler, 10 actual steps are used") + # sample_scheduler = FlowMatchEulerDiscreteScheduler(shift=shift, use_beta_sigmas=(scheduler == 'euler/beta')) + # sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None) + # start_latent_list = [0, 5, 10, 15, 20, 25, 30, 35, 40, 45, 50] + # sample_scheduler.sigmas = sample_scheduler.sigmas[start_latent_list] + # steps = len(start_latent_list) - 1 + # sample_scheduler.timesteps = timesteps = sample_scheduler.timesteps[start_latent_list[:steps]] elif 'dpm++' in scheduler: if 'sde' in scheduler: algorithm_type = "sde-dpmsolver++" @@ -84,6 +100,9 @@ def get_scheduler(scheduler, steps, shift, device, transformer_dim, flowedit_arg shift=shift, sigma_min=0.0, extra_one_step=True ) sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength, shift=shift) + elif scheduler == 'res_multistep': + sample_scheduler = FlowMatchSchedulerResMultistep(shift=shift) + sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength) if timesteps is None: timesteps = sample_scheduler.timesteps log.info(f"timesteps: {timesteps}") diff --git a/wanvideo/schedulers/flowmatch_res_multistep.py b/wanvideo/schedulers/flowmatch_res_multistep.py new file mode 100644 index 0000000..f109aa2 --- /dev/null +++ b/wanvideo/schedulers/flowmatch_res_multistep.py @@ -0,0 +1,105 @@ +import torch + +sigma_fn = lambda t: t.neg().exp() +t_fn = lambda sigma: sigma.log().neg() +phi1_fn = lambda t: torch.expm1(t) / t +phi2_fn = lambda t: (phi1_fn(t) - 1.0) / t + +class FlowMatchSchedulerResMultistep(): + + def __init__(self, num_inference_steps=100, num_train_timesteps=1000, shift=3.0, sigma_max=1.0, sigma_min=0.003 / 1.002, extra_one_step=False): + self.num_train_timesteps = num_train_timesteps + self.shift = shift + self.sigma_max = sigma_max + self.sigma_min = sigma_min + self.extra_one_step = extra_one_step + self.set_timesteps(num_inference_steps) + self.prev_model_output = None + self.old_sigma_next = None + + def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0): + #Generate the full sigma schedule (from max to min) + if self.extra_one_step: + sigma_start = self.sigma_min + \ + (self.sigma_max - self.sigma_min) * denoising_strength + self.sigmas = torch.linspace( + sigma_start, self.sigma_min, num_inference_steps + 1)[:-1] + full_sigmas = torch.linspace(self.sigma_max, self.sigma_min, self.num_train_timesteps) + ss = len(full_sigmas) / num_inference_steps + sigmas = [] + for x in range(num_inference_steps): + idx = int(round(x * ss)) + sigmas.append(float(full_sigmas[idx])) + sigmas.append(0.0) + self.sigmas = torch.FloatTensor(sigmas) + self.sigmas = self.shift * self.sigmas / \ + (1 + (self.shift - 1) * self.sigmas) + self.timesteps = self.sigmas * self.num_train_timesteps + #print(f"Timesteps: {self.timesteps}, Sigmas: {self.sigmas}") + + + def step(self, model_output, timestep, sample): + if timestep.ndim == 2: + timestep = timestep.flatten(0, 1) + self.sigmas = self.sigmas.to(model_output.device) + self.timesteps = self.timesteps.to(model_output.device) + if timestep.ndim == 0: + timestep_id = torch.argmin((self.timesteps - timestep).abs(), dim=0) + else: + timestep_id = torch.argmin((self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1) + + sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1) + sigma_prev = self.sigmas[timestep_id - 1].reshape(-1, 1, 1, 1) if timestep_id > 0 else sigma + if (timestep_id + 1 >= len(self.timesteps)).any(): + sigma_next = torch.tensor(0) + else: + sigma_next = self.sigmas[timestep_id + 1].reshape(-1, 1, 1, 1) + + x0_pred = (sample - sigma * model_output) + + if sigma_next == 0 or self.prev_model_output is None: + x = sample + model_output * (sigma_next - sigma) + else: + t, t_old, t_next, t_prev = t_fn(sigma), t_fn(self.old_sigma_next), t_fn(sigma_next), t_fn(sigma_prev) + h = t_next - t + c2 = (t_prev - t_old) / h + phi1_val, phi2_val = phi1_fn(-h), phi2_fn(-h) + b1 = torch.nan_to_num(phi1_val - phi2_val / c2, nan=0.0) + b2 = torch.nan_to_num(phi2_val / c2, nan=0.0) + + x = sigma_fn(h) * sample + h * (b1 * x0_pred + b2 * self.prev_model_output) + + self.old_sigma_next = sigma_next + self.prev_model_output = x0_pred + return x + + + def add_noise(self, original_samples, noise, timestep): + """ + Diffusion forward corruption process. + Input: + - clean_latent: the clean latent with shape [B*T, C, H, W] + - noise: the noise with shape [B*T, C, H, W] + - timestep: the timestep with shape [B*T] + Output: the corrupted latent with shape [B*T, C, H, W] + """ + if timestep.ndim == 2: + timestep = timestep.flatten(0, 1) + self.sigmas = self.sigmas.to(noise.device) + self.timesteps = self.timesteps.to(noise.device) + timestep_id = torch.argmin( + (self.timesteps.unsqueeze(0) - timestep.unsqueeze(1)).abs(), dim=1) + sigma = self.sigmas[timestep_id].reshape(-1, 1, 1, 1) + sample = (1 - sigma) * original_samples + sigma * noise + return sample.type_as(noise) + + def training_target(self, sample, noise, timestep): + target = noise - sample + return target + + def training_weight(self, timestep): + timestep_id = torch.argmin( + (self.timesteps - timestep.to(self.timesteps.device)).abs()) + weights = self.linear_timesteps_weights[timestep_id] + return weights +