From 9a5d752ba1f70cd9d42410780a9437a0b5b0cb33 Mon Sep 17 00:00:00 2001 From: kijai <40791699+kijai@users.noreply.github.com> Date: Wed, 30 Jul 2025 16:20:48 +0300 Subject: [PATCH] Support custom sigmas on more schedulers --- wanvideo/schedulers/__init__.py | 12 ++++++++---- wanvideo/schedulers/flowmatch_pusa.py | 15 +++++++++------ wanvideo/schedulers/flowmatch_res_multistep.py | 13 +++++++------ 3 files changed, 24 insertions(+), 16 deletions(-) diff --git a/wanvideo/schedulers/__init__.py b/wanvideo/schedulers/__init__.py index 48cd37d..4431889 100644 --- a/wanvideo/schedulers/__init__.py +++ b/wanvideo/schedulers/__init__.py @@ -40,7 +40,7 @@ def get_scheduler(scheduler, steps, shift, device, transformer_dim, flowedit_arg if flowedit_args: #seems to work better 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) + sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas[:-1].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") @@ -68,8 +68,10 @@ def get_scheduler(scheduler, steps, shift, device, transformer_dim, flowedit_arg sample_scheduler.sigmas[-1] = 1e-6 elif 'lcm' in scheduler: sample_scheduler = FlowMatchLCMScheduler(shift=shift, use_beta_sigmas=(scheduler == 'lcm/beta')) - sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas.tolist() if sigmas is not None else None) + sample_scheduler.set_timesteps(steps, device=device, sigmas=sigmas[:-1].tolist() if sigmas is not None else None) elif 'flowmatch_causvid' in scheduler: + if sigmas is not None: + raise NotImplementedError("This scheduler does not support custom sigmas") if transformer_dim == 5120: denoising_list = [999, 934, 862, 756, 603, 410, 250, 140, 74] else: @@ -80,6 +82,8 @@ def get_scheduler(scheduler, steps, shift, device, transformer_dim, flowedit_arg sample_scheduler.timesteps = torch.tensor(denoising_list)[:steps].to(device) sample_scheduler.sigmas = torch.cat([sample_scheduler.timesteps / 1000, torch.tensor([0.0], device=device)]) elif 'flowmatch_distill' in scheduler: + if sigmas is not None: + raise NotImplementedError("This scheduler does not support custom sigmas") sample_scheduler = FlowMatchScheduler( shift=shift, sigma_min=0.0, extra_one_step=True ) @@ -99,10 +103,10 @@ def get_scheduler(scheduler, steps, shift, device, transformer_dim, flowedit_arg sample_scheduler = FlowMatchSchedulerPusa( shift=shift, sigma_min=0.0, extra_one_step=True ) - sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength, shift=shift) + sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength, shift=shift, sigmas=sigmas[:-1].tolist() if sigmas is not None else None) elif scheduler == 'res_multistep': sample_scheduler = FlowMatchSchedulerResMultistep(shift=shift) - sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength) + sample_scheduler.set_timesteps(steps, denoising_strength=denoise_strength, sigmas=sigmas[:-1].tolist() if sigmas is not None else None) if timesteps is None: timesteps = sample_scheduler.timesteps log.info(f"timesteps: {timesteps}") diff --git a/wanvideo/schedulers/flowmatch_pusa.py b/wanvideo/schedulers/flowmatch_pusa.py index 12c3b32..94dd9d9 100644 --- a/wanvideo/schedulers/flowmatch_pusa.py +++ b/wanvideo/schedulers/flowmatch_pusa.py @@ -15,16 +15,19 @@ class FlowMatchSchedulerPusa(): self.set_timesteps(num_inference_steps) - def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0, training=False, shift=None): + def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0, training=False, shift=None, sigmas=None): if shift is not None: self.shift = shift sigma_start = self.sigma_min + (self.sigma_max - self.sigma_min) * denoising_strength - if self.extra_one_step: - self.sigmas = torch.linspace(sigma_start, self.sigma_min, num_inference_steps + 1)[:-1] + if sigmas is None: + if self.extra_one_step: + self.sigmas = torch.linspace(sigma_start, self.sigma_min, num_inference_steps + 1)[:-1] + else: + self.sigmas = torch.linspace(sigma_start, self.sigma_min, num_inference_steps) + if self.inverse_timesteps: + self.sigmas = torch.flip(self.sigmas, dims=[0]) else: - self.sigmas = torch.linspace(sigma_start, self.sigma_min, num_inference_steps) - if self.inverse_timesteps: - self.sigmas = torch.flip(self.sigmas, dims=[0]) + self.sigmas = torch.tensor(sigmas, dtype=torch.float32) self.sigmas = self.shift * self.sigmas / (1 + (self.shift - 1) * self.sigmas) if self.reverse_sigmas: self.sigmas = 1 - self.sigmas diff --git a/wanvideo/schedulers/flowmatch_res_multistep.py b/wanvideo/schedulers/flowmatch_res_multistep.py index f109aa2..8a644d1 100644 --- a/wanvideo/schedulers/flowmatch_res_multistep.py +++ b/wanvideo/schedulers/flowmatch_res_multistep.py @@ -17,7 +17,7 @@ class FlowMatchSchedulerResMultistep(): self.prev_model_output = None self.old_sigma_next = None - def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0): + def set_timesteps(self, num_inference_steps=100, denoising_strength=1.0, sigmas=None): #Generate the full sigma schedule (from max to min) if self.extra_one_step: sigma_start = self.sigma_min + \ @@ -26,11 +26,12 @@ class FlowMatchSchedulerResMultistep(): 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) + if sigmas is None: + 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)