CoolFace
Apppublic

fred-dev/comfy_ui_ali

sourceHugging Facemitupdated 1y agoView on Hugging Face
0likes
sampling.py1425 linesDownload Raw Back to k_diffusion
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]

Showing the first 1,200 of 1425 lines. Download the file for the rest.