From a80ef83a84a3c99f46ee0420163e0cbe71bea8de Mon Sep 17 00:00:00 2001 From: ssit Date: Mon, 3 Jul 2023 20:18:59 -0400 Subject: [PATCH] Added support for the normal scheduler --- nodes.py | 4 +++- restart_sampling.py | 29 ++++++++++++++++++++--------- 2 files changed, 23 insertions(+), 10 deletions(-) diff --git a/nodes.py b/nodes.py index 8645ba4..c162f0a 100644 --- a/nodes.py +++ b/nodes.py @@ -13,7 +13,9 @@ def get_supported_samplers(): def get_supported_restart_schedulers(): - schedulers = ["karras", "exponential"] + schedulers = comfy.samplers.KSampler.SCHEDULERS.copy() + schedulers.remove("simple") + schedulers.remove("ddim_uniform") return schedulers diff --git a/restart_sampling.py b/restart_sampling.py index 5586fda..78787b7 100644 --- a/restart_sampling.py +++ b/restart_sampling.py @@ -3,6 +3,7 @@ import torch from tqdm.auto import trange from nodes import common_ksampler from comfy.k_diffusion import sampling as k_diffusion_sampling +from comfy.samplers import simple_scheduler, ddim_scheduler from comfy.utils import ProgressBar @@ -36,14 +37,24 @@ def round_restart_segments(sigmas, restart_segments): return t_min_mapping -def calc_sigmas(scheduler, n, sigma_min, sigma_max, device): - if scheduler == "karras": - sigmas = k_diffusion_sampling.get_sigmas_karras(n, sigma_min, sigma_max, device=device) - elif scheduler == "exponential": - sigmas = k_diffusion_sampling.get_sigmas_exponential(n, sigma_min, sigma_max, device=device) - else: - raise ValueError("Unsupported scheduler") - return sigmas +def calc_sigmas(scheduler, n, sigma_min, sigma_max, model, device): + match scheduler: + case "karras": + return k_diffusion_sampling.get_sigmas_karras(n, sigma_min, sigma_max, device=device) + case "exponential": + return k_diffusion_sampling.get_sigmas_exponential(n, sigma_min, sigma_max, device=device) + case "normal": + def get_sigmas(model, n, s_min, s_max): + t_min, t_max = model.sigma_to_t(torch.tensor([s_min, s_max], device=device)) + t = torch.linspace(t_max, t_min, n, device=device) + return k_diffusion_sampling.append_zero(model.t_to_sigma(t)) + return get_sigmas(model.inner_model, n, sigma_min, sigma_max) + # case "simple": + # sigmas = simple_scheduler(model.inner_model, steps) + # case "ddim_uniform": + # sigmas = ddim_scheduler(model.inner_model, steps) + case _: + raise ValueError("Unsupported scheduler") def calc_restart_steps(restart_segments): @@ -81,7 +92,7 @@ def restart_sampling(model, seed, steps, cfg, sampler_name, scheduler, positive, if sigmas[i + 1].item() in segments: seg = segments[sigmas[i + 1].item()] s_min, s_max, k, n_restart = sigmas[i + 1], seg['t_max'], seg['k'], seg['n'] - seg_sigmas = calc_sigmas(restart_scheduler, n_restart, s_min, s_max, device=x.device) + seg_sigmas = calc_sigmas(restart_scheduler, n_restart, s_min, s_max, model, device=x.device) for _ in range(k): x += torch.randn_like(x) * (s_max ** 2 - s_min ** 2) ** 0.5 for j in range(n_restart - 1):