Added support for ddim_uniform scheduler

Also renamed schedulers.py to be more specific.
This commit is contained in:
ssit
2023-07-05 12:39:31 -04:00
parent c1df472c17
commit 446cdc3f19
2 changed files with 21 additions and 3 deletions
+7 -3
View File
@@ -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):
+14
View File
@@ -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,
}