Support custom sigmas on more schedulers

This commit is contained in:
kijai
2025-07-30 16:20:48 +03:00
parent ec066008a8
commit 9a5d752ba1
3 changed files with 24 additions and 16 deletions
+8 -4
View File
@@ -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}")
+9 -6
View File
@@ -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
@@ -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)