From 446cdc3f1984d07ec87e5215d921713105c1fec6 Mon Sep 17 00:00:00 2001 From: ssit Date: Wed, 5 Jul 2023 12:39:31 -0400 Subject: [PATCH] Added support for ddim_uniform scheduler Also renamed schedulers.py to be more specific. --- restart_sampling.py | 10 +++++++--- schedulers.py => restart_schedulers.py | 14 ++++++++++++++ 2 files changed, 21 insertions(+), 3 deletions(-) rename schedulers.py => restart_schedulers.py (77%) diff --git a/restart_sampling.py b/restart_sampling.py index c30ffc5..7b0762c 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -4,7 +4,7 @@ from tqdm.auto import trange from nodes import common_ksampler from comfy.k_diffusion import sampling as k_diffusion_sampling from comfy.utils import ProgressBar -from .schedulers import SCHEDULER_MAPPING +from .restart_schedulers import SCHEDULER_MAPPING def add_restart_segment(restart_segments, n_restart, k, t_min, t_max): @@ -48,18 +48,22 @@ def calc_restart_steps(restart_segments): return restart_steps +total_steps = 0 + + def restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, negative, latent_image, denoise, restart_info, restart_scheduler): sample_func_name = "sample_{}".format(sampler_name) sampler = getattr(k_diffusion_sampling, sample_func_name) restart_segments = prepare_restart_segments(restart_info) + global total_steps total_steps = steps @torch.no_grad() def restart_wrapper(model, x, sigmas, extra_args=None, callback=None, disable=None): extra_args = {} if extra_args is None else extra_args segments = round_restart_segments(sigmas, restart_segments) - nonlocal total_steps - total_steps = steps + calc_restart_steps(segments) + global total_steps + total_steps = len(sigmas) - 1 + calc_restart_steps(segments) step = 0 def callback_wrapper(x): diff --git a/schedulers.py b/restart_schedulers.py similarity index 77% rename from schedulers.py rename to restart_schedulers.py index 55b4169..58eb628 100644 --- a/schedulers.py +++ b/restart_schedulers.py @@ -1,5 +1,6 @@ import torch from comfy.k_diffusion import sampling as k_diffusion_sampling +import warnings def get_sigmas_karras(model, n, s_min, s_max, device): @@ -28,6 +29,18 @@ def get_sigmas_simple(model, n, s_min, s_max, device): return torch.tensor(sigs, device=device) +def get_sigmas_ddim_uniform(model, n, s_min, s_max, device): + t_min, t_max = model.sigma_to_t(torch.tensor([s_min, s_max], device=device)) + ddim_timesteps = torch.linspace(t_max, t_min, n, dtype=torch.int16, device=device) + sigs = [] + for ts in ddim_timesteps: + if ts > 999: + ts = 999 + sigs.append(model.t_to_sigma(ts)) + sigs += [0.0] + return torch.tensor(sigs, device=device) + + def get_sigmas_simple_test(model, n, s_min, s_max, device): min_idx = torch.argmin(torch.abs(model.sigmas - s_min)) max_idx = torch.argmin(torch.abs(model.sigmas - s_max)) @@ -45,5 +58,6 @@ SCHEDULER_MAPPING = { "karras": get_sigmas_karras, "exponential": get_sigmas_exponential, "simple": get_sigmas_simple, + "ddim_uniform": get_sigmas_ddim_uniform, "simple_test": get_sigmas_simple_test, }