Added support for ddim_uniform scheduler
Also renamed schedulers.py to be more specific.
This commit is contained in:
+7
-3
@@ -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):
|
||||
|
||||
@@ -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,
|
||||
}
|
||||
Reference in New Issue
Block a user