Support custom sigmas on more schedulers
This commit is contained in:
@@ -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}")
|
||||
|
||||
@@ -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)
|
||||
|
||||
Reference in New Issue
Block a user