fred-dev/comfy_ui_ali
0
1import math2 3from scipy import integrate4import torch5from torch import nn6import torchsde7from tqdm.auto import trange, tqdm8 9from . import utils10from . import deis11import comfy.model_patcher12import comfy.model_sampling13 14def append_zero(x):15 return torch.cat([x, x.new_zeros([1])])16 17 18def get_sigmas_karras(n, sigma_min, sigma_max, rho=7., device='cpu'):19 """Constructs the noise schedule of Karras et al. (2022)."""20 ramp = torch.linspace(0, 1, n, device=device)21 min_inv_rho = sigma_min ** (1 / rho)22 max_inv_rho = sigma_max ** (1 / rho)23 sigmas = (max_inv_rho + ramp * (min_inv_rho - max_inv_rho)) ** rho24 return append_zero(sigmas).to(device)25 26 27def get_sigmas_exponential(n, sigma_min, sigma_max, device='cpu'):28 """Constructs an exponential noise schedule."""29 sigmas = torch.linspace(math.log(sigma_max), math.log(sigma_min), n, device=device).exp()30 return append_zero(sigmas)31 32 33def get_sigmas_polyexponential(n, sigma_min, sigma_max, rho=1., device='cpu'):34 """Constructs an polynomial in log sigma noise schedule."""35 ramp = torch.linspace(1, 0, n, device=device) ** rho36 sigmas = torch.exp(ramp * (math.log(sigma_max) - math.log(sigma_min)) + math.log(sigma_min))37 return append_zero(sigmas)38 39 40def get_sigmas_vp(n, beta_d=19.9, beta_min=0.1, eps_s=1e-3, device='cpu'):41 """Constructs a continuous VP noise schedule."""42 t = torch.linspace(1, eps_s, n, device=device)43 sigmas = torch.sqrt(torch.special.expm1(beta_d * t ** 2 / 2 + beta_min * t))44 return append_zero(sigmas)45 46 47def get_sigmas_laplace(n, sigma_min, sigma_max, mu=0., beta=0.5, device='cpu'):48 """Constructs the noise schedule proposed by Tiankai et al. (2024). """49 epsilon = 1e-5 # avoid log(0)50 x = torch.linspace(0, 1, n, device=device)51 clamp = lambda x: torch.clamp(x, min=sigma_min, max=sigma_max)52 lmb = mu - beta * torch.sign(0.5-x) * torch.log(1 - 2 * torch.abs(0.5-x) + epsilon)53 sigmas = clamp(torch.exp(lmb))54 return sigmas55 56 57 58def to_d(x, sigma, denoised):59 """Converts a denoiser output to a Karras ODE derivative."""60 return (x - denoised) / utils.append_dims(sigma, x.ndim)61 62 63def get_ancestral_step(sigma_from, sigma_to, eta=1.):64 """Calculates the noise level (sigma_down) to step down to and the amount65 of noise to add (sigma_up) when doing an ancestral sampling step."""66 if not eta:67 return sigma_to, 0.68 sigma_up = min(sigma_to, eta * (sigma_to ** 2 * (sigma_from ** 2 - sigma_to ** 2) / sigma_from ** 2) ** 0.5)69 sigma_down = (sigma_to ** 2 - sigma_up ** 2) ** 0.570 return sigma_down, sigma_up71 72 73def default_noise_sampler(x, seed=None):74 if seed is not None:75 generator = torch.Generator(device=x.device)76 generator.manual_seed(seed)77 else:78 generator = None79 80 return lambda sigma, sigma_next: torch.randn(x.size(), dtype=x.dtype, layout=x.layout, device=x.device, generator=generator)81 82 83class BatchedBrownianTree:84 """A wrapper around torchsde.BrownianTree that enables batches of entropy."""85 86 def __init__(self, x, t0, t1, seed=None, **kwargs):87 self.cpu_tree = True88 if "cpu" in kwargs:89 self.cpu_tree = kwargs.pop("cpu")90 t0, t1, self.sign = self.sort(t0, t1)91 w0 = kwargs.get('w0', torch.zeros_like(x))92 if seed is None:93 seed = torch.randint(0, 2 ** 63 - 1, []).item()94 self.batched = True95 try:96 assert len(seed) == x.shape[0]97 w0 = w0[0]98 except TypeError:99 seed = [seed]100 self.batched = False101 if self.cpu_tree:102 self.trees = [torchsde.BrownianTree(t0.cpu(), w0.cpu(), t1.cpu(), entropy=s, **kwargs) for s in seed]103 else:104 self.trees = [torchsde.BrownianTree(t0, w0, t1, entropy=s, **kwargs) for s in seed]105 106 @staticmethod107 def sort(a, b):108 return (a, b, 1) if a < b else (b, a, -1)109 110 def __call__(self, t0, t1):111 t0, t1, sign = self.sort(t0, t1)112 if self.cpu_tree:113 w = torch.stack([tree(t0.cpu().float(), t1.cpu().float()).to(t0.dtype).to(t0.device) for tree in self.trees]) * (self.sign * sign)114 else:115 w = torch.stack([tree(t0, t1) for tree in self.trees]) * (self.sign * sign)116 117 return w if self.batched else w[0]118 119 120class BrownianTreeNoiseSampler:121 """A noise sampler backed by a torchsde.BrownianTree.122 123 Args:124 x (Tensor): The tensor whose shape, device and dtype to use to generate125 random samples.126 sigma_min (float): The low end of the valid interval.127 sigma_max (float): The high end of the valid interval.128 seed (int or List[int]): The random seed. If a list of seeds is129 supplied instead of a single integer, then the noise sampler will130 use one BrownianTree per batch item, each with its own seed.131 transform (callable): A function that maps sigma to the sampler's132 internal timestep.133 """134 135 def __init__(self, x, sigma_min, sigma_max, seed=None, transform=lambda x: x, cpu=False):136 self.transform = transform137 t0, t1 = self.transform(torch.as_tensor(sigma_min)), self.transform(torch.as_tensor(sigma_max))138 self.tree = BatchedBrownianTree(x, t0, t1, seed, cpu=cpu)139 140 def __call__(self, sigma, sigma_next):141 t0, t1 = self.transform(torch.as_tensor(sigma)), self.transform(torch.as_tensor(sigma_next))142 return self.tree(t0, t1) / (t1 - t0).abs().sqrt()143 144 145@torch.no_grad()146def sample_euler(model, x, sigmas, extra_args=None, callback=None, disable=None, s_churn=0., s_tmin=0., s_tmax=float('inf'), s_noise=1.):147 """Implements Algorithm 2 (Euler steps) from Karras et al. (2022)."""148 extra_args = {} if extra_args is None else extra_args149 s_in = x.new_ones([x.shape[0]])150 for i in trange(len(sigmas) - 1, disable=disable):151 if s_churn > 0:152 gamma = min(s_churn / (len(sigmas) - 1), 2 ** 0.5 - 1) if s_tmin <= sigmas[i] <= s_tmax else 0.153 sigma_hat = sigmas[i] * (gamma + 1)154 else:155 gamma = 0156 sigma_hat = sigmas[i]157 158 if gamma > 0:159 eps = torch.randn_like(x) * s_noise160 x = x + eps * (sigma_hat ** 2 - sigmas[i] ** 2) ** 0.5161 denoised = model(x, sigma_hat * s_in, **extra_args)162 d = to_d(x, sigma_hat, denoised)163 if callback is not None:164 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigma_hat, 'denoised': denoised})165 dt = sigmas[i + 1] - sigma_hat166 # Euler method167 x = x + d * dt168 return x169 170 171@torch.no_grad()172def sample_euler_ancestral(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):173 if isinstance(model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST):174 return sample_euler_ancestral_RF(model, x, sigmas, extra_args, callback, disable, eta, s_noise, noise_sampler)175 """Ancestral sampling with Euler method steps."""176 extra_args = {} if extra_args is None else extra_args177 seed = extra_args.get("seed", None)178 noise_sampler = default_noise_sampler(x, seed=seed) if noise_sampler is None else noise_sampler179 s_in = x.new_ones([x.shape[0]])180 for i in trange(len(sigmas) - 1, disable=disable):181 denoised = model(x, sigmas[i] * s_in, **extra_args)182 sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)183 if callback is not None:184 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})185 186 if sigma_down == 0:187 x = denoised188 else:189 d = to_d(x, sigmas[i], denoised)190 # Euler method191 dt = sigma_down - sigmas[i]192 x = x + d * dt + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up193 return x194 195@torch.no_grad()196def sample_euler_ancestral_RF(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1.0, s_noise=1., noise_sampler=None):197 """Ancestral sampling with Euler method steps."""198 extra_args = {} if extra_args is None else extra_args199 seed = extra_args.get("seed", None)200 noise_sampler = default_noise_sampler(x, seed=seed) if noise_sampler is None else noise_sampler201 s_in = x.new_ones([x.shape[0]])202 for i in trange(len(sigmas) - 1, disable=disable):203 denoised = model(x, sigmas[i] * s_in, **extra_args)204 # sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)205 if callback is not None:206 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})207 208 if sigmas[i + 1] == 0:209 x = denoised210 else:211 downstep_ratio = 1 + (sigmas[i + 1] / sigmas[i] - 1) * eta212 sigma_down = sigmas[i + 1] * downstep_ratio213 alpha_ip1 = 1 - sigmas[i + 1]214 alpha_down = 1 - sigma_down215 renoise_coeff = (sigmas[i + 1]**2 - sigma_down**2 * alpha_ip1**2 / alpha_down**2)**0.5216 # Euler method217 sigma_down_i_ratio = sigma_down / sigmas[i]218 x = sigma_down_i_ratio * x + (1 - sigma_down_i_ratio) * denoised219 if eta > 0:220 x = (alpha_ip1 / alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff221 return x222 223@torch.no_grad()224def sample_heun(model, x, sigmas, extra_args=None, callback=None, disable=None, s_churn=0., s_tmin=0., s_tmax=float('inf'), s_noise=1.):225 """Implements Algorithm 2 (Heun steps) from Karras et al. (2022)."""226 extra_args = {} if extra_args is None else extra_args227 s_in = x.new_ones([x.shape[0]])228 for i in trange(len(sigmas) - 1, disable=disable):229 if s_churn > 0:230 gamma = min(s_churn / (len(sigmas) - 1), 2 ** 0.5 - 1) if s_tmin <= sigmas[i] <= s_tmax else 0.231 sigma_hat = sigmas[i] * (gamma + 1)232 else:233 gamma = 0234 sigma_hat = sigmas[i]235 236 sigma_hat = sigmas[i] * (gamma + 1)237 if gamma > 0:238 eps = torch.randn_like(x) * s_noise239 x = x + eps * (sigma_hat ** 2 - sigmas[i] ** 2) ** 0.5240 denoised = model(x, sigma_hat * s_in, **extra_args)241 d = to_d(x, sigma_hat, denoised)242 if callback is not None:243 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigma_hat, 'denoised': denoised})244 dt = sigmas[i + 1] - sigma_hat245 if sigmas[i + 1] == 0:246 # Euler method247 x = x + d * dt248 else:249 # Heun's method250 x_2 = x + d * dt251 denoised_2 = model(x_2, sigmas[i + 1] * s_in, **extra_args)252 d_2 = to_d(x_2, sigmas[i + 1], denoised_2)253 d_prime = (d + d_2) / 2254 x = x + d_prime * dt255 return x256 257 258@torch.no_grad()259def sample_dpm_2(model, x, sigmas, extra_args=None, callback=None, disable=None, s_churn=0., s_tmin=0., s_tmax=float('inf'), s_noise=1.):260 """A sampler inspired by DPM-Solver-2 and Algorithm 2 from Karras et al. (2022)."""261 extra_args = {} if extra_args is None else extra_args262 s_in = x.new_ones([x.shape[0]])263 for i in trange(len(sigmas) - 1, disable=disable):264 if s_churn > 0:265 gamma = min(s_churn / (len(sigmas) - 1), 2 ** 0.5 - 1) if s_tmin <= sigmas[i] <= s_tmax else 0.266 sigma_hat = sigmas[i] * (gamma + 1)267 else:268 gamma = 0269 sigma_hat = sigmas[i]270 271 if gamma > 0:272 eps = torch.randn_like(x) * s_noise273 x = x + eps * (sigma_hat ** 2 - sigmas[i] ** 2) ** 0.5274 denoised = model(x, sigma_hat * s_in, **extra_args)275 d = to_d(x, sigma_hat, denoised)276 if callback is not None:277 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigma_hat, 'denoised': denoised})278 if sigmas[i + 1] == 0:279 # Euler method280 dt = sigmas[i + 1] - sigma_hat281 x = x + d * dt282 else:283 # DPM-Solver-2284 sigma_mid = sigma_hat.log().lerp(sigmas[i + 1].log(), 0.5).exp()285 dt_1 = sigma_mid - sigma_hat286 dt_2 = sigmas[i + 1] - sigma_hat287 x_2 = x + d * dt_1288 denoised_2 = model(x_2, sigma_mid * s_in, **extra_args)289 d_2 = to_d(x_2, sigma_mid, denoised_2)290 x = x + d_2 * dt_2291 return x292 293 294@torch.no_grad()295def sample_dpm_2_ancestral(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):296 if isinstance(model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST):297 return sample_dpm_2_ancestral_RF(model, x, sigmas, extra_args, callback, disable, eta, s_noise, noise_sampler)298 299 """Ancestral sampling with DPM-Solver second-order steps."""300 extra_args = {} if extra_args is None else extra_args301 seed = extra_args.get("seed", None)302 noise_sampler = default_noise_sampler(x, seed=seed) if noise_sampler is None else noise_sampler303 s_in = x.new_ones([x.shape[0]])304 for i in trange(len(sigmas) - 1, disable=disable):305 denoised = model(x, sigmas[i] * s_in, **extra_args)306 sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)307 if callback is not None:308 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})309 d = to_d(x, sigmas[i], denoised)310 if sigma_down == 0:311 # Euler method312 dt = sigma_down - sigmas[i]313 x = x + d * dt314 else:315 # DPM-Solver-2316 sigma_mid = sigmas[i].log().lerp(sigma_down.log(), 0.5).exp()317 dt_1 = sigma_mid - sigmas[i]318 dt_2 = sigma_down - sigmas[i]319 x_2 = x + d * dt_1320 denoised_2 = model(x_2, sigma_mid * s_in, **extra_args)321 d_2 = to_d(x_2, sigma_mid, denoised_2)322 x = x + d_2 * dt_2323 x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up324 return x325 326@torch.no_grad()327def sample_dpm_2_ancestral_RF(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):328 """Ancestral sampling with DPM-Solver second-order steps."""329 extra_args = {} if extra_args is None else extra_args330 seed = extra_args.get("seed", None)331 noise_sampler = default_noise_sampler(x, seed=seed) if noise_sampler is None else noise_sampler332 s_in = x.new_ones([x.shape[0]])333 for i in trange(len(sigmas) - 1, disable=disable):334 denoised = model(x, sigmas[i] * s_in, **extra_args)335 downstep_ratio = 1 + (sigmas[i+1]/sigmas[i] - 1) * eta336 sigma_down = sigmas[i+1] * downstep_ratio337 alpha_ip1 = 1 - sigmas[i+1]338 alpha_down = 1 - sigma_down339 renoise_coeff = (sigmas[i+1]**2 - sigma_down**2*alpha_ip1**2/alpha_down**2)**0.5340 341 if callback is not None:342 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})343 d = to_d(x, sigmas[i], denoised)344 if sigma_down == 0:345 # Euler method346 dt = sigma_down - sigmas[i]347 x = x + d * dt348 else:349 # DPM-Solver-2350 sigma_mid = sigmas[i].log().lerp(sigma_down.log(), 0.5).exp()351 dt_1 = sigma_mid - sigmas[i]352 dt_2 = sigma_down - sigmas[i]353 x_2 = x + d * dt_1354 denoised_2 = model(x_2, sigma_mid * s_in, **extra_args)355 d_2 = to_d(x_2, sigma_mid, denoised_2)356 x = x + d_2 * dt_2357 x = (alpha_ip1/alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff358 return x359 360def linear_multistep_coeff(order, t, i, j):361 if order - 1 > i:362 raise ValueError(f'Order {order} too high for step {i}')363 def fn(tau):364 prod = 1.365 for k in range(order):366 if j == k:367 continue368 prod *= (tau - t[i - k]) / (t[i - j] - t[i - k])369 return prod370 return integrate.quad(fn, t[i], t[i + 1], epsrel=1e-4)[0]371 372 373@torch.no_grad()374def sample_lms(model, x, sigmas, extra_args=None, callback=None, disable=None, order=4):375 extra_args = {} if extra_args is None else extra_args376 s_in = x.new_ones([x.shape[0]])377 sigmas_cpu = sigmas.detach().cpu().numpy()378 ds = []379 for i in trange(len(sigmas) - 1, disable=disable):380 denoised = model(x, sigmas[i] * s_in, **extra_args)381 d = to_d(x, sigmas[i], denoised)382 ds.append(d)383 if len(ds) > order:384 ds.pop(0)385 if callback is not None:386 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})387 cur_order = min(i + 1, order)388 coeffs = [linear_multistep_coeff(cur_order, sigmas_cpu, i, j) for j in range(cur_order)]389 x = x + sum(coeff * d for coeff, d in zip(coeffs, reversed(ds)))390 return x391 392 393class PIDStepSizeController:394 """A PID controller for ODE adaptive step size control."""395 def __init__(self, h, pcoeff, icoeff, dcoeff, order=1, accept_safety=0.81, eps=1e-8):396 self.h = h397 self.b1 = (pcoeff + icoeff + dcoeff) / order398 self.b2 = -(pcoeff + 2 * dcoeff) / order399 self.b3 = dcoeff / order400 self.accept_safety = accept_safety401 self.eps = eps402 self.errs = []403 404 def limiter(self, x):405 return 1 + math.atan(x - 1)406 407 def propose_step(self, error):408 inv_error = 1 / (float(error) + self.eps)409 if not self.errs:410 self.errs = [inv_error, inv_error, inv_error]411 self.errs[0] = inv_error412 factor = self.errs[0] ** self.b1 * self.errs[1] ** self.b2 * self.errs[2] ** self.b3413 factor = self.limiter(factor)414 accept = factor >= self.accept_safety415 if accept:416 self.errs[2] = self.errs[1]417 self.errs[1] = self.errs[0]418 self.h *= factor419 return accept420 421 422class DPMSolver(nn.Module):423 """DPM-Solver. See https://arxiv.org/abs/2206.00927."""424 425 def __init__(self, model, extra_args=None, eps_callback=None, info_callback=None):426 super().__init__()427 self.model = model428 self.extra_args = {} if extra_args is None else extra_args429 self.eps_callback = eps_callback430 self.info_callback = info_callback431 432 def t(self, sigma):433 return -sigma.log()434 435 def sigma(self, t):436 return t.neg().exp()437 438 def eps(self, eps_cache, key, x, t, *args, **kwargs):439 if key in eps_cache:440 return eps_cache[key], eps_cache441 sigma = self.sigma(t) * x.new_ones([x.shape[0]])442 eps = (x - self.model(x, sigma, *args, **self.extra_args, **kwargs)) / self.sigma(t)443 if self.eps_callback is not None:444 self.eps_callback()445 return eps, {key: eps, **eps_cache}446 447 def dpm_solver_1_step(self, x, t, t_next, eps_cache=None):448 eps_cache = {} if eps_cache is None else eps_cache449 h = t_next - t450 eps, eps_cache = self.eps(eps_cache, 'eps', x, t)451 x_1 = x - self.sigma(t_next) * h.expm1() * eps452 return x_1, eps_cache453 454 def dpm_solver_2_step(self, x, t, t_next, r1=1 / 2, eps_cache=None):455 eps_cache = {} if eps_cache is None else eps_cache456 h = t_next - t457 eps, eps_cache = self.eps(eps_cache, 'eps', x, t)458 s1 = t + r1 * h459 u1 = x - self.sigma(s1) * (r1 * h).expm1() * eps460 eps_r1, eps_cache = self.eps(eps_cache, 'eps_r1', u1, s1)461 x_2 = x - self.sigma(t_next) * h.expm1() * eps - self.sigma(t_next) / (2 * r1) * h.expm1() * (eps_r1 - eps)462 return x_2, eps_cache463 464 def dpm_solver_3_step(self, x, t, t_next, r1=1 / 3, r2=2 / 3, eps_cache=None):465 eps_cache = {} if eps_cache is None else eps_cache466 h = t_next - t467 eps, eps_cache = self.eps(eps_cache, 'eps', x, t)468 s1 = t + r1 * h469 s2 = t + r2 * h470 u1 = x - self.sigma(s1) * (r1 * h).expm1() * eps471 eps_r1, eps_cache = self.eps(eps_cache, 'eps_r1', u1, s1)472 u2 = x - self.sigma(s2) * (r2 * h).expm1() * eps - self.sigma(s2) * (r2 / r1) * ((r2 * h).expm1() / (r2 * h) - 1) * (eps_r1 - eps)473 eps_r2, eps_cache = self.eps(eps_cache, 'eps_r2', u2, s2)474 x_3 = x - self.sigma(t_next) * h.expm1() * eps - self.sigma(t_next) / r2 * (h.expm1() / h - 1) * (eps_r2 - eps)475 return x_3, eps_cache476 477 def dpm_solver_fast(self, x, t_start, t_end, nfe, eta=0., s_noise=1., noise_sampler=None):478 noise_sampler = default_noise_sampler(x, seed=self.extra_args.get("seed", None)) if noise_sampler is None else noise_sampler479 if not t_end > t_start and eta:480 raise ValueError('eta must be 0 for reverse sampling')481 482 m = math.floor(nfe / 3) + 1483 ts = torch.linspace(t_start, t_end, m + 1, device=x.device)484 485 if nfe % 3 == 0:486 orders = [3] * (m - 2) + [2, 1]487 else:488 orders = [3] * (m - 1) + [nfe % 3]489 490 for i in range(len(orders)):491 eps_cache = {}492 t, t_next = ts[i], ts[i + 1]493 if eta:494 sd, su = get_ancestral_step(self.sigma(t), self.sigma(t_next), eta)495 t_next_ = torch.minimum(t_end, self.t(sd))496 su = (self.sigma(t_next) ** 2 - self.sigma(t_next_) ** 2) ** 0.5497 else:498 t_next_, su = t_next, 0.499 500 eps, eps_cache = self.eps(eps_cache, 'eps', x, t)501 denoised = x - self.sigma(t) * eps502 if self.info_callback is not None:503 self.info_callback({'x': x, 'i': i, 't': ts[i], 't_up': t, 'denoised': denoised})504 505 if orders[i] == 1:506 x, eps_cache = self.dpm_solver_1_step(x, t, t_next_, eps_cache=eps_cache)507 elif orders[i] == 2:508 x, eps_cache = self.dpm_solver_2_step(x, t, t_next_, eps_cache=eps_cache)509 else:510 x, eps_cache = self.dpm_solver_3_step(x, t, t_next_, eps_cache=eps_cache)511 512 x = x + su * s_noise * noise_sampler(self.sigma(t), self.sigma(t_next))513 514 return x515 516 def dpm_solver_adaptive(self, x, t_start, t_end, order=3, rtol=0.05, atol=0.0078, h_init=0.05, pcoeff=0., icoeff=1., dcoeff=0., accept_safety=0.81, eta=0., s_noise=1., noise_sampler=None):517 noise_sampler = default_noise_sampler(x, seed=self.extra_args.get("seed", None)) if noise_sampler is None else noise_sampler518 if order not in {2, 3}:519 raise ValueError('order should be 2 or 3')520 forward = t_end > t_start521 if not forward and eta:522 raise ValueError('eta must be 0 for reverse sampling')523 h_init = abs(h_init) * (1 if forward else -1)524 atol = torch.tensor(atol)525 rtol = torch.tensor(rtol)526 s = t_start527 x_prev = x528 accept = True529 pid = PIDStepSizeController(h_init, pcoeff, icoeff, dcoeff, 1.5 if eta else order, accept_safety)530 info = {'steps': 0, 'nfe': 0, 'n_accept': 0, 'n_reject': 0}531 532 while s < t_end - 1e-5 if forward else s > t_end + 1e-5:533 eps_cache = {}534 t = torch.minimum(t_end, s + pid.h) if forward else torch.maximum(t_end, s + pid.h)535 if eta:536 sd, su = get_ancestral_step(self.sigma(s), self.sigma(t), eta)537 t_ = torch.minimum(t_end, self.t(sd))538 su = (self.sigma(t) ** 2 - self.sigma(t_) ** 2) ** 0.5539 else:540 t_, su = t, 0.541 542 eps, eps_cache = self.eps(eps_cache, 'eps', x, s)543 denoised = x - self.sigma(s) * eps544 545 if order == 2:546 x_low, eps_cache = self.dpm_solver_1_step(x, s, t_, eps_cache=eps_cache)547 x_high, eps_cache = self.dpm_solver_2_step(x, s, t_, eps_cache=eps_cache)548 else:549 x_low, eps_cache = self.dpm_solver_2_step(x, s, t_, r1=1 / 3, eps_cache=eps_cache)550 x_high, eps_cache = self.dpm_solver_3_step(x, s, t_, eps_cache=eps_cache)551 delta = torch.maximum(atol, rtol * torch.maximum(x_low.abs(), x_prev.abs()))552 error = torch.linalg.norm((x_low - x_high) / delta) / x.numel() ** 0.5553 accept = pid.propose_step(error)554 if accept:555 x_prev = x_low556 x = x_high + su * s_noise * noise_sampler(self.sigma(s), self.sigma(t))557 s = t558 info['n_accept'] += 1559 else:560 info['n_reject'] += 1561 info['nfe'] += order562 info['steps'] += 1563 564 if self.info_callback is not None:565 self.info_callback({'x': x, 'i': info['steps'] - 1, 't': s, 't_up': s, 'denoised': denoised, 'error': error, 'h': pid.h, **info})566 567 return x, info568 569 570@torch.no_grad()571def sample_dpm_fast(model, x, sigma_min, sigma_max, n, extra_args=None, callback=None, disable=None, eta=0., s_noise=1., noise_sampler=None):572 """DPM-Solver-Fast (fixed step size). See https://arxiv.org/abs/2206.00927."""573 if sigma_min <= 0 or sigma_max <= 0:574 raise ValueError('sigma_min and sigma_max must not be 0')575 with tqdm(total=n, disable=disable) as pbar:576 dpm_solver = DPMSolver(model, extra_args, eps_callback=pbar.update)577 if callback is not None:578 dpm_solver.info_callback = lambda info: callback({'sigma': dpm_solver.sigma(info['t']), 'sigma_hat': dpm_solver.sigma(info['t_up']), **info})579 return dpm_solver.dpm_solver_fast(x, dpm_solver.t(torch.tensor(sigma_max)), dpm_solver.t(torch.tensor(sigma_min)), n, eta, s_noise, noise_sampler)580 581 582@torch.no_grad()583def sample_dpm_adaptive(model, x, sigma_min, sigma_max, extra_args=None, callback=None, disable=None, order=3, rtol=0.05, atol=0.0078, h_init=0.05, pcoeff=0., icoeff=1., dcoeff=0., accept_safety=0.81, eta=0., s_noise=1., noise_sampler=None, return_info=False):584 """DPM-Solver-12 and 23 (adaptive step size). See https://arxiv.org/abs/2206.00927."""585 if sigma_min <= 0 or sigma_max <= 0:586 raise ValueError('sigma_min and sigma_max must not be 0')587 with tqdm(disable=disable) as pbar:588 dpm_solver = DPMSolver(model, extra_args, eps_callback=pbar.update)589 if callback is not None:590 dpm_solver.info_callback = lambda info: callback({'sigma': dpm_solver.sigma(info['t']), 'sigma_hat': dpm_solver.sigma(info['t_up']), **info})591 x, info = dpm_solver.dpm_solver_adaptive(x, dpm_solver.t(torch.tensor(sigma_max)), dpm_solver.t(torch.tensor(sigma_min)), order, rtol, atol, h_init, pcoeff, icoeff, dcoeff, accept_safety, eta, s_noise, noise_sampler)592 if return_info:593 return x, info594 return x595 596 597@torch.no_grad()598def sample_dpmpp_2s_ancestral(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):599 if isinstance(model.inner_model.inner_model.model_sampling, comfy.model_sampling.CONST):600 return sample_dpmpp_2s_ancestral_RF(model, x, sigmas, extra_args, callback, disable, eta, s_noise, noise_sampler)601 602 """Ancestral sampling with DPM-Solver++(2S) second-order steps."""603 extra_args = {} if extra_args is None else extra_args604 seed = extra_args.get("seed", None)605 noise_sampler = default_noise_sampler(x, seed=seed) if noise_sampler is None else noise_sampler606 s_in = x.new_ones([x.shape[0]])607 sigma_fn = lambda t: t.neg().exp()608 t_fn = lambda sigma: sigma.log().neg()609 610 for i in trange(len(sigmas) - 1, disable=disable):611 denoised = model(x, sigmas[i] * s_in, **extra_args)612 sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)613 if callback is not None:614 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})615 if sigma_down == 0:616 # Euler method617 d = to_d(x, sigmas[i], denoised)618 dt = sigma_down - sigmas[i]619 x = x + d * dt620 else:621 # DPM-Solver++(2S)622 t, t_next = t_fn(sigmas[i]), t_fn(sigma_down)623 r = 1 / 2624 h = t_next - t625 s = t + r * h626 x_2 = (sigma_fn(s) / sigma_fn(t)) * x - (-h * r).expm1() * denoised627 denoised_2 = model(x_2, sigma_fn(s) * s_in, **extra_args)628 x = (sigma_fn(t_next) / sigma_fn(t)) * x - (-h).expm1() * denoised_2629 # Noise addition630 if sigmas[i + 1] > 0:631 x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up632 return x633 634 635@torch.no_grad()636def sample_dpmpp_2s_ancestral_RF(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):637 """Ancestral sampling with DPM-Solver++(2S) second-order steps."""638 extra_args = {} if extra_args is None else extra_args639 seed = extra_args.get("seed", None)640 noise_sampler = default_noise_sampler(x, seed=seed) if noise_sampler is None else noise_sampler641 s_in = x.new_ones([x.shape[0]])642 sigma_fn = lambda lbda: (lbda.exp() + 1) ** -1643 lambda_fn = lambda sigma: ((1-sigma)/sigma).log()644 645 # logged_x = x.unsqueeze(0)646 647 for i in trange(len(sigmas) - 1, disable=disable):648 denoised = model(x, sigmas[i] * s_in, **extra_args)649 downstep_ratio = 1 + (sigmas[i+1]/sigmas[i] - 1) * eta650 sigma_down = sigmas[i+1] * downstep_ratio651 alpha_ip1 = 1 - sigmas[i+1]652 alpha_down = 1 - sigma_down653 renoise_coeff = (sigmas[i+1]**2 - sigma_down**2*alpha_ip1**2/alpha_down**2)**0.5654 # sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)655 if callback is not None:656 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})657 if sigmas[i + 1] == 0:658 # Euler method659 d = to_d(x, sigmas[i], denoised)660 dt = sigma_down - sigmas[i]661 x = x + d * dt662 else:663 # DPM-Solver++(2S)664 if sigmas[i] == 1.0:665 sigma_s = 0.9999666 else:667 t_i, t_down = lambda_fn(sigmas[i]), lambda_fn(sigma_down)668 r = 1 / 2669 h = t_down - t_i670 s = t_i + r * h671 sigma_s = sigma_fn(s)672 # sigma_s = sigmas[i+1]673 sigma_s_i_ratio = sigma_s / sigmas[i]674 u = sigma_s_i_ratio * x + (1 - sigma_s_i_ratio) * denoised675 D_i = model(u, sigma_s * s_in, **extra_args)676 sigma_down_i_ratio = sigma_down / sigmas[i]677 x = sigma_down_i_ratio * x + (1 - sigma_down_i_ratio) * D_i678 # print("sigma_i", sigmas[i], "sigma_ip1", sigmas[i+1],"sigma_down", sigma_down, "sigma_down_i_ratio", sigma_down_i_ratio, "sigma_s_i_ratio", sigma_s_i_ratio, "renoise_coeff", renoise_coeff)679 # Noise addition680 if sigmas[i + 1] > 0 and eta > 0:681 x = (alpha_ip1/alpha_down) * x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * renoise_coeff682 # logged_x = torch.cat((logged_x, x.unsqueeze(0)), dim=0)683 return x684 685@torch.no_grad()686def sample_dpmpp_sde(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, r=1 / 2):687 """DPM-Solver++ (stochastic)."""688 if len(sigmas) <= 1:689 return x690 691 extra_args = {} if extra_args is None else extra_args692 sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()693 seed = extra_args.get("seed", None)694 noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=seed, cpu=True) if noise_sampler is None else noise_sampler695 s_in = x.new_ones([x.shape[0]])696 sigma_fn = lambda t: t.neg().exp()697 t_fn = lambda sigma: sigma.log().neg()698 699 for i in trange(len(sigmas) - 1, disable=disable):700 denoised = model(x, sigmas[i] * s_in, **extra_args)701 if callback is not None:702 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})703 if sigmas[i + 1] == 0:704 # Euler method705 d = to_d(x, sigmas[i], denoised)706 dt = sigmas[i + 1] - sigmas[i]707 x = x + d * dt708 else:709 # DPM-Solver++710 t, t_next = t_fn(sigmas[i]), t_fn(sigmas[i + 1])711 h = t_next - t712 s = t + h * r713 fac = 1 / (2 * r)714 715 # Step 1716 sd, su = get_ancestral_step(sigma_fn(t), sigma_fn(s), eta)717 s_ = t_fn(sd)718 x_2 = (sigma_fn(s_) / sigma_fn(t)) * x - (t - s_).expm1() * denoised719 x_2 = x_2 + noise_sampler(sigma_fn(t), sigma_fn(s)) * s_noise * su720 denoised_2 = model(x_2, sigma_fn(s) * s_in, **extra_args)721 722 # Step 2723 sd, su = get_ancestral_step(sigma_fn(t), sigma_fn(t_next), eta)724 t_next_ = t_fn(sd)725 denoised_d = (1 - fac) * denoised + fac * denoised_2726 x = (sigma_fn(t_next_) / sigma_fn(t)) * x - (t - t_next_).expm1() * denoised_d727 x = x + noise_sampler(sigma_fn(t), sigma_fn(t_next)) * s_noise * su728 return x729 730 731@torch.no_grad()732def sample_dpmpp_2m(model, x, sigmas, extra_args=None, callback=None, disable=None):733 """DPM-Solver++(2M)."""734 extra_args = {} if extra_args is None else extra_args735 s_in = x.new_ones([x.shape[0]])736 sigma_fn = lambda t: t.neg().exp()737 t_fn = lambda sigma: sigma.log().neg()738 old_denoised = None739 740 for i in trange(len(sigmas) - 1, disable=disable):741 denoised = model(x, sigmas[i] * s_in, **extra_args)742 if callback is not None:743 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})744 t, t_next = t_fn(sigmas[i]), t_fn(sigmas[i + 1])745 h = t_next - t746 if old_denoised is None or sigmas[i + 1] == 0:747 x = (sigma_fn(t_next) / sigma_fn(t)) * x - (-h).expm1() * denoised748 else:749 h_last = t - t_fn(sigmas[i - 1])750 r = h_last / h751 denoised_d = (1 + 1 / (2 * r)) * denoised - (1 / (2 * r)) * old_denoised752 x = (sigma_fn(t_next) / sigma_fn(t)) * x - (-h).expm1() * denoised_d753 old_denoised = denoised754 return x755 756@torch.no_grad()757def sample_dpmpp_2m_sde(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, solver_type='midpoint'):758 """DPM-Solver++(2M) SDE."""759 if len(sigmas) <= 1:760 return x761 762 if solver_type not in {'heun', 'midpoint'}:763 raise ValueError('solver_type must be \'heun\' or \'midpoint\'')764 765 extra_args = {} if extra_args is None else extra_args766 seed = extra_args.get("seed", None)767 sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()768 noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=seed, cpu=True) if noise_sampler is None else noise_sampler769 s_in = x.new_ones([x.shape[0]])770 771 old_denoised = None772 h_last = None773 h = None774 775 for i in trange(len(sigmas) - 1, disable=disable):776 denoised = model(x, sigmas[i] * s_in, **extra_args)777 if callback is not None:778 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})779 if sigmas[i + 1] == 0:780 # Denoising step781 x = denoised782 else:783 # DPM-Solver++(2M) SDE784 t, s = -sigmas[i].log(), -sigmas[i + 1].log()785 h = s - t786 eta_h = eta * h787 788 x = sigmas[i + 1] / sigmas[i] * (-eta_h).exp() * x + (-h - eta_h).expm1().neg() * denoised789 790 if old_denoised is not None:791 r = h_last / h792 if solver_type == 'heun':793 x = x + ((-h - eta_h).expm1().neg() / (-h - eta_h) + 1) * (1 / r) * (denoised - old_denoised)794 elif solver_type == 'midpoint':795 x = x + 0.5 * (-h - eta_h).expm1().neg() * (1 / r) * (denoised - old_denoised)796 797 if eta:798 x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[i + 1] * (-2 * eta_h).expm1().neg().sqrt() * s_noise799 800 old_denoised = denoised801 h_last = h802 return x803 804@torch.no_grad()805def sample_dpmpp_3m_sde(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):806 """DPM-Solver++(3M) SDE."""807 808 if len(sigmas) <= 1:809 return x810 811 extra_args = {} if extra_args is None else extra_args812 seed = extra_args.get("seed", None)813 sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()814 noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=seed, cpu=True) if noise_sampler is None else noise_sampler815 s_in = x.new_ones([x.shape[0]])816 817 denoised_1, denoised_2 = None, None818 h, h_1, h_2 = None, None, None819 820 for i in trange(len(sigmas) - 1, disable=disable):821 denoised = model(x, sigmas[i] * s_in, **extra_args)822 if callback is not None:823 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})824 if sigmas[i + 1] == 0:825 # Denoising step826 x = denoised827 else:828 t, s = -sigmas[i].log(), -sigmas[i + 1].log()829 h = s - t830 h_eta = h * (eta + 1)831 832 x = torch.exp(-h_eta) * x + (-h_eta).expm1().neg() * denoised833 834 if h_2 is not None:835 r0 = h_1 / h836 r1 = h_2 / h837 d1_0 = (denoised - denoised_1) / r0838 d1_1 = (denoised_1 - denoised_2) / r1839 d1 = d1_0 + (d1_0 - d1_1) * r0 / (r0 + r1)840 d2 = (d1_0 - d1_1) / (r0 + r1)841 phi_2 = h_eta.neg().expm1() / h_eta + 1842 phi_3 = phi_2 / h_eta - 0.5843 x = x + phi_2 * d1 - phi_3 * d2844 elif h_1 is not None:845 r = h_1 / h846 d = (denoised - denoised_1) / r847 phi_2 = h_eta.neg().expm1() / h_eta + 1848 x = x + phi_2 * d849 850 if eta:851 x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * sigmas[i + 1] * (-2 * h * eta).expm1().neg().sqrt() * s_noise852 853 denoised_1, denoised_2 = denoised, denoised_1854 h_1, h_2 = h, h_1855 return x856 857@torch.no_grad()858def sample_dpmpp_3m_sde_gpu(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):859 if len(sigmas) <= 1:860 return x861 extra_args = {} if extra_args is None else extra_args862 sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()863 noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=extra_args.get("seed", None), cpu=False) if noise_sampler is None else noise_sampler864 return sample_dpmpp_3m_sde(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler)865 866@torch.no_grad()867def sample_dpmpp_2m_sde_gpu(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, solver_type='midpoint'):868 if len(sigmas) <= 1:869 return x870 extra_args = {} if extra_args is None else extra_args871 sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()872 noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=extra_args.get("seed", None), cpu=False) if noise_sampler is None else noise_sampler873 return sample_dpmpp_2m_sde(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler, solver_type=solver_type)874 875@torch.no_grad()876def sample_dpmpp_sde_gpu(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None, r=1 / 2):877 if len(sigmas) <= 1:878 return x879 extra_args = {} if extra_args is None else extra_args880 sigma_min, sigma_max = sigmas[sigmas > 0].min(), sigmas.max()881 noise_sampler = BrownianTreeNoiseSampler(x, sigma_min, sigma_max, seed=extra_args.get("seed", None), cpu=False) if noise_sampler is None else noise_sampler882 return sample_dpmpp_sde(model, x, sigmas, extra_args=extra_args, callback=callback, disable=disable, eta=eta, s_noise=s_noise, noise_sampler=noise_sampler, r=r)883 884 885def DDPMSampler_step(x, sigma, sigma_prev, noise, noise_sampler):886 alpha_cumprod = 1 / ((sigma * sigma) + 1)887 alpha_cumprod_prev = 1 / ((sigma_prev * sigma_prev) + 1)888 alpha = (alpha_cumprod / alpha_cumprod_prev)889 890 mu = (1.0 / alpha).sqrt() * (x - (1 - alpha) * noise / (1 - alpha_cumprod).sqrt())891 if sigma_prev > 0:892 mu += ((1 - alpha) * (1. - alpha_cumprod_prev) / (1. - alpha_cumprod)).sqrt() * noise_sampler(sigma, sigma_prev)893 return mu894 895def generic_step_sampler(model, x, sigmas, extra_args=None, callback=None, disable=None, noise_sampler=None, step_function=None):896 extra_args = {} if extra_args is None else extra_args897 seed = extra_args.get("seed", None)898 noise_sampler = default_noise_sampler(x, seed=seed) if noise_sampler is None else noise_sampler899 s_in = x.new_ones([x.shape[0]])900 901 for i in trange(len(sigmas) - 1, disable=disable):902 denoised = model(x, sigmas[i] * s_in, **extra_args)903 if callback is not None:904 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})905 x = step_function(x / torch.sqrt(1.0 + sigmas[i] ** 2.0), sigmas[i], sigmas[i + 1], (x - denoised) / sigmas[i], noise_sampler)906 if sigmas[i + 1] != 0:907 x *= torch.sqrt(1.0 + sigmas[i + 1] ** 2.0)908 return x909 910 911@torch.no_grad()912def sample_ddpm(model, x, sigmas, extra_args=None, callback=None, disable=None, noise_sampler=None):913 return generic_step_sampler(model, x, sigmas, extra_args, callback, disable, noise_sampler, DDPMSampler_step)914 915@torch.no_grad()916def sample_lcm(model, x, sigmas, extra_args=None, callback=None, disable=None, noise_sampler=None):917 extra_args = {} if extra_args is None else extra_args918 seed = extra_args.get("seed", None)919 noise_sampler = default_noise_sampler(x, seed=seed) if noise_sampler is None else noise_sampler920 s_in = x.new_ones([x.shape[0]])921 for i in trange(len(sigmas) - 1, disable=disable):922 denoised = model(x, sigmas[i] * s_in, **extra_args)923 if callback is not None:924 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})925 926 x = denoised927 if sigmas[i + 1] > 0:928 x = model.inner_model.inner_model.model_sampling.noise_scaling(sigmas[i + 1], noise_sampler(sigmas[i], sigmas[i + 1]), x)929 return x930 931 932 933@torch.no_grad()934def sample_heunpp2(model, x, sigmas, extra_args=None, callback=None, disable=None, s_churn=0., s_tmin=0., s_tmax=float('inf'), s_noise=1.):935 # From MIT licensed: https://github.com/Carzit/sd-webui-samplers-scheduler/936 extra_args = {} if extra_args is None else extra_args937 s_in = x.new_ones([x.shape[0]])938 s_end = sigmas[-1]939 for i in trange(len(sigmas) - 1, disable=disable):940 gamma = min(s_churn / (len(sigmas) - 1), 2 ** 0.5 - 1) if s_tmin <= sigmas[i] <= s_tmax else 0.941 eps = torch.randn_like(x) * s_noise942 sigma_hat = sigmas[i] * (gamma + 1)943 if gamma > 0:944 x = x + eps * (sigma_hat ** 2 - sigmas[i] ** 2) ** 0.5945 denoised = model(x, sigma_hat * s_in, **extra_args)946 d = to_d(x, sigma_hat, denoised)947 if callback is not None:948 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigma_hat, 'denoised': denoised})949 dt = sigmas[i + 1] - sigma_hat950 if sigmas[i + 1] == s_end:951 # Euler method952 x = x + d * dt953 elif sigmas[i + 2] == s_end:954 955 # Heun's method956 x_2 = x + d * dt957 denoised_2 = model(x_2, sigmas[i + 1] * s_in, **extra_args)958 d_2 = to_d(x_2, sigmas[i + 1], denoised_2)959 960 w = 2 * sigmas[0]961 w2 = sigmas[i+1]/w962 w1 = 1 - w2963 964 d_prime = d * w1 + d_2 * w2965 966 967 x = x + d_prime * dt968 969 else:970 # Heun++971 x_2 = x + d * dt972 denoised_2 = model(x_2, sigmas[i + 1] * s_in, **extra_args)973 d_2 = to_d(x_2, sigmas[i + 1], denoised_2)974 dt_2 = sigmas[i + 2] - sigmas[i + 1]975 976 x_3 = x_2 + d_2 * dt_2977 denoised_3 = model(x_3, sigmas[i + 2] * s_in, **extra_args)978 d_3 = to_d(x_3, sigmas[i + 2], denoised_3)979 980 w = 3 * sigmas[0]981 w2 = sigmas[i + 1] / w982 w3 = sigmas[i + 2] / w983 w1 = 1 - w2 - w3984 985 d_prime = w1 * d + w2 * d_2 + w3 * d_3986 x = x + d_prime * dt987 return x988 989 990#From https://github.com/zju-pi/diff-sampler/blob/main/diff-solvers-main/solvers.py991#under Apache 2 license992def sample_ipndm(model, x, sigmas, extra_args=None, callback=None, disable=None, max_order=4):993 extra_args = {} if extra_args is None else extra_args994 s_in = x.new_ones([x.shape[0]])995 996 x_next = x997 998 buffer_model = []999 for i in trange(len(sigmas) - 1, disable=disable):1000 t_cur = sigmas[i]1001 t_next = sigmas[i + 1]1002 1003 x_cur = x_next1004 1005 denoised = model(x_cur, t_cur * s_in, **extra_args)1006 if callback is not None:1007 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})1008 1009 d_cur = (x_cur - denoised) / t_cur1010 1011 order = min(max_order, i+1)1012 if order == 1: # First Euler step.1013 x_next = x_cur + (t_next - t_cur) * d_cur1014 elif order == 2: # Use one history point.1015 x_next = x_cur + (t_next - t_cur) * (3 * d_cur - buffer_model[-1]) / 21016 elif order == 3: # Use two history points.1017 x_next = x_cur + (t_next - t_cur) * (23 * d_cur - 16 * buffer_model[-1] + 5 * buffer_model[-2]) / 121018 elif order == 4: # Use three history points.1019 x_next = x_cur + (t_next - t_cur) * (55 * d_cur - 59 * buffer_model[-1] + 37 * buffer_model[-2] - 9 * buffer_model[-3]) / 241020 1021 if len(buffer_model) == max_order - 1:1022 for k in range(max_order - 2):1023 buffer_model[k] = buffer_model[k+1]1024 buffer_model[-1] = d_cur1025 else:1026 buffer_model.append(d_cur)1027 1028 return x_next1029 1030#From https://github.com/zju-pi/diff-sampler/blob/main/diff-solvers-main/solvers.py1031#under Apache 2 license1032def sample_ipndm_v(model, x, sigmas, extra_args=None, callback=None, disable=None, max_order=4):1033 extra_args = {} if extra_args is None else extra_args1034 s_in = x.new_ones([x.shape[0]])1035 1036 x_next = x1037 t_steps = sigmas1038 1039 buffer_model = []1040 for i in trange(len(sigmas) - 1, disable=disable):1041 t_cur = sigmas[i]1042 t_next = sigmas[i + 1]1043 1044 x_cur = x_next1045 1046 denoised = model(x_cur, t_cur * s_in, **extra_args)1047 if callback is not None:1048 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})1049 1050 d_cur = (x_cur - denoised) / t_cur1051 1052 order = min(max_order, i+1)1053 if order == 1: # First Euler step.1054 x_next = x_cur + (t_next - t_cur) * d_cur1055 elif order == 2: # Use one history point.1056 h_n = (t_next - t_cur)1057 h_n_1 = (t_cur - t_steps[i-1])1058 coeff1 = (2 + (h_n / h_n_1)) / 21059 coeff2 = -(h_n / h_n_1) / 21060 x_next = x_cur + (t_next - t_cur) * (coeff1 * d_cur + coeff2 * buffer_model[-1])1061 elif order == 3: # Use two history points.1062 h_n = (t_next - t_cur)1063 h_n_1 = (t_cur - t_steps[i-1])1064 h_n_2 = (t_steps[i-1] - t_steps[i-2])1065 temp = (1 - h_n / (3 * (h_n + h_n_1)) * (h_n * (h_n + h_n_1)) / (h_n_1 * (h_n_1 + h_n_2))) / 21066 coeff1 = (2 + (h_n / h_n_1)) / 2 + temp1067 coeff2 = -(h_n / h_n_1) / 2 - (1 + h_n_1 / h_n_2) * temp1068 coeff3 = temp * h_n_1 / h_n_21069 x_next = x_cur + (t_next - t_cur) * (coeff1 * d_cur + coeff2 * buffer_model[-1] + coeff3 * buffer_model[-2])1070 elif order == 4: # Use three history points.1071 h_n = (t_next - t_cur)1072 h_n_1 = (t_cur - t_steps[i-1])1073 h_n_2 = (t_steps[i-1] - t_steps[i-2])1074 h_n_3 = (t_steps[i-2] - t_steps[i-3])1075 temp1 = (1 - h_n / (3 * (h_n + h_n_1)) * (h_n * (h_n + h_n_1)) / (h_n_1 * (h_n_1 + h_n_2))) / 21076 temp2 = ((1 - h_n / (3 * (h_n + h_n_1))) / 2 + (1 - h_n / (2 * (h_n + h_n_1))) * h_n / (6 * (h_n + h_n_1 + h_n_2))) \1077 * (h_n * (h_n + h_n_1) * (h_n + h_n_1 + h_n_2)) / (h_n_1 * (h_n_1 + h_n_2) * (h_n_1 + h_n_2 + h_n_3))1078 coeff1 = (2 + (h_n / h_n_1)) / 2 + temp1 + temp21079 coeff2 = -(h_n / h_n_1) / 2 - (1 + h_n_1 / h_n_2) * temp1 - (1 + (h_n_1 / h_n_2) + (h_n_1 * (h_n_1 + h_n_2) / (h_n_2 * (h_n_2 + h_n_3)))) * temp21080 coeff3 = temp1 * h_n_1 / h_n_2 + ((h_n_1 / h_n_2) + (h_n_1 * (h_n_1 + h_n_2) / (h_n_2 * (h_n_2 + h_n_3))) * (1 + h_n_2 / h_n_3)) * temp21081 coeff4 = -temp2 * (h_n_1 * (h_n_1 + h_n_2) / (h_n_2 * (h_n_2 + h_n_3))) * h_n_1 / h_n_21082 x_next = x_cur + (t_next - t_cur) * (coeff1 * d_cur + coeff2 * buffer_model[-1] + coeff3 * buffer_model[-2] + coeff4 * buffer_model[-3])1083 1084 if len(buffer_model) == max_order - 1:1085 for k in range(max_order - 2):1086 buffer_model[k] = buffer_model[k+1]1087 buffer_model[-1] = d_cur.detach()1088 else:1089 buffer_model.append(d_cur.detach())1090 1091 return x_next1092 1093#From https://github.com/zju-pi/diff-sampler/blob/main/diff-solvers-main/solvers.py1094#under Apache 2 license1095@torch.no_grad()1096def sample_deis(model, x, sigmas, extra_args=None, callback=None, disable=None, max_order=3, deis_mode='tab'):1097 extra_args = {} if extra_args is None else extra_args1098 s_in = x.new_ones([x.shape[0]])1099 1100 x_next = x1101 t_steps = sigmas1102 1103 coeff_list = deis.get_deis_coeff_list(t_steps, max_order, deis_mode=deis_mode)1104 1105 buffer_model = []1106 for i in trange(len(sigmas) - 1, disable=disable):1107 t_cur = sigmas[i]1108 t_next = sigmas[i + 1]1109 1110 x_cur = x_next1111 1112 denoised = model(x_cur, t_cur * s_in, **extra_args)1113 if callback is not None:1114 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})1115 1116 d_cur = (x_cur - denoised) / t_cur1117 1118 order = min(max_order, i+1)1119 if t_next <= 0:1120 order = 11121 1122 if order == 1: # First Euler step.1123 x_next = x_cur + (t_next - t_cur) * d_cur1124 elif order == 2: # Use one history point.1125 coeff_cur, coeff_prev1 = coeff_list[i]1126 x_next = x_cur + coeff_cur * d_cur + coeff_prev1 * buffer_model[-1]1127 elif order == 3: # Use two history points.1128 coeff_cur, coeff_prev1, coeff_prev2 = coeff_list[i]1129 x_next = x_cur + coeff_cur * d_cur + coeff_prev1 * buffer_model[-1] + coeff_prev2 * buffer_model[-2]1130 elif order == 4: # Use three history points.1131 coeff_cur, coeff_prev1, coeff_prev2, coeff_prev3 = coeff_list[i]1132 x_next = x_cur + coeff_cur * d_cur + coeff_prev1 * buffer_model[-1] + coeff_prev2 * buffer_model[-2] + coeff_prev3 * buffer_model[-3]1133 1134 if len(buffer_model) == max_order - 1:1135 for k in range(max_order - 2):1136 buffer_model[k] = buffer_model[k+1]1137 buffer_model[-1] = d_cur.detach()1138 else:1139 buffer_model.append(d_cur.detach())1140 1141 return x_next1142 1143@torch.no_grad()1144def sample_euler_cfg_pp(model, x, sigmas, extra_args=None, callback=None, disable=None):1145 extra_args = {} if extra_args is None else extra_args1146 1147 temp = [0]1148 def post_cfg_function(args):1149 temp[0] = args["uncond_denoised"]1150 return args["denoised"]1151 1152 model_options = extra_args.get("model_options", {}).copy()1153 extra_args["model_options"] = comfy.model_patcher.set_model_options_post_cfg_function(model_options, post_cfg_function, disable_cfg1_optimization=True)1154 1155 s_in = x.new_ones([x.shape[0]])1156 for i in trange(len(sigmas) - 1, disable=disable):1157 sigma_hat = sigmas[i]1158 denoised = model(x, sigma_hat * s_in, **extra_args)1159 d = to_d(x, sigma_hat, temp[0])1160 if callback is not None:1161 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigma_hat, 'denoised': denoised})1162 # Euler method1163 x = denoised + d * sigmas[i + 1]1164 return x1165 1166@torch.no_grad()1167def sample_euler_ancestral_cfg_pp(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):1168 """Ancestral sampling with Euler method steps."""1169 extra_args = {} if extra_args is None else extra_args1170 seed = extra_args.get("seed", None)1171 noise_sampler = default_noise_sampler(x, seed=seed) if noise_sampler is None else noise_sampler1172 1173 temp = [0]1174 def post_cfg_function(args):1175 temp[0] = args["uncond_denoised"]1176 return args["denoised"]1177 1178 model_options = extra_args.get("model_options", {}).copy()1179 extra_args["model_options"] = comfy.model_patcher.set_model_options_post_cfg_function(model_options, post_cfg_function, disable_cfg1_optimization=True)1180 1181 s_in = x.new_ones([x.shape[0]])1182 for i in trange(len(sigmas) - 1, disable=disable):1183 denoised = model(x, sigmas[i] * s_in, **extra_args)1184 sigma_down, sigma_up = get_ancestral_step(sigmas[i], sigmas[i + 1], eta=eta)1185 if callback is not None:1186 callback({'x': x, 'i': i, 'sigma': sigmas[i], 'sigma_hat': sigmas[i], 'denoised': denoised})1187 d = to_d(x, sigmas[i], temp[0])1188 # Euler method1189 x = denoised + d * sigma_down1190 if sigmas[i + 1] > 0:1191 x = x + noise_sampler(sigmas[i], sigmas[i + 1]) * s_noise * sigma_up1192 return x1193@torch.no_grad()1194def sample_dpmpp_2s_ancestral_cfg_pp(model, x, sigmas, extra_args=None, callback=None, disable=None, eta=1., s_noise=1., noise_sampler=None):1195 """Ancestral sampling with DPM-Solver++(2S) second-order steps."""1196 extra_args = {} if extra_args is None else extra_args1197 seed = extra_args.get("seed", None)1198 noise_sampler = default_noise_sampler(x, seed=seed) if noise_sampler is None else noise_sampler1199 1200 temp = [0]