CoolFace
Apppublic

fred-dev/comfy_ui_ali

sourceHugging Facemitupdated 1y agoView on Hugging Face
0likes
deis.py121 linesDownload Raw Back to k_diffusion
1#Taken from: https://github.com/zju-pi/diff-sampler/blob/main/gits-main/solver_utils.py2#under Apache 2 license3import torch4import numpy as np5 6# A pytorch reimplementation of DEIS (https://github.com/qsh-zh/deis).7#############################8### Utils for DEIS solver ###9#############################10#----------------------------------------------------------------------------11# Transfer from the input time (sigma) used in EDM to that (t) used in DEIS.12 13def edm2t(edm_steps, epsilon_s=1e-3, sigma_min=0.002, sigma_max=80):14    vp_sigma_inv = lambda beta_d, beta_min: lambda sigma: ((beta_min ** 2 + 2 * beta_d * (sigma ** 2 + 1).log()).sqrt() - beta_min) / beta_d15    vp_beta_d = 2 * (np.log(torch.tensor(sigma_min).cpu() ** 2 + 1) / epsilon_s - np.log(torch.tensor(sigma_max).cpu() ** 2 + 1)) / (epsilon_s - 1)16    vp_beta_min = np.log(torch.tensor(sigma_max).cpu() ** 2 + 1) - 0.5 * vp_beta_d17    t_steps = vp_sigma_inv(vp_beta_d.clone().detach().cpu(), vp_beta_min.clone().detach().cpu())(edm_steps.clone().detach().cpu())18    return t_steps, vp_beta_min, vp_beta_d + vp_beta_min19 20#----------------------------------------------------------------------------21 22def cal_poly(prev_t, j, taus):23    poly = 124    for k in range(prev_t.shape[0]):25        if k == j:26            continue27        poly *= (taus - prev_t[k]) / (prev_t[j] - prev_t[k])28    return poly29 30#----------------------------------------------------------------------------31# Transfer from t to alpha_t.32 33def t2alpha_fn(beta_0, beta_1, t):34    return torch.exp(-0.5 * t ** 2 * (beta_1 - beta_0) - t * beta_0)35 36#----------------------------------------------------------------------------37 38def cal_intergrand(beta_0, beta_1, taus):39    with torch.inference_mode(mode=False):40        taus = taus.clone()41        beta_0 = beta_0.clone()42        beta_1 = beta_1.clone()43        with torch.enable_grad():44            taus.requires_grad_(True)45            alpha = t2alpha_fn(beta_0, beta_1, taus)46            log_alpha = alpha.log()47            log_alpha.sum().backward()48            d_log_alpha_dtau = taus.grad49    integrand = -0.5 * d_log_alpha_dtau / torch.sqrt(alpha * (1 - alpha))50    return integrand51 52#----------------------------------------------------------------------------53 54def get_deis_coeff_list(t_steps, max_order, N=10000, deis_mode='tab'):55    """56    Get the coefficient list for DEIS sampling.57 58    Args:59        t_steps: A pytorch tensor. The time steps for sampling.60        max_order: A `int`. Maximum order of the solver. 1 <= max_order <= 461        N: A `int`. Use how many points to perform the numerical integration when deis_mode=='tab'.62        deis_mode: A `str`. Select between 'tab' and 'rhoab'. Type of DEIS.63    Returns:64        A pytorch tensor. A batch of generated samples or sampling trajectories if return_inters=True.65    """66    if deis_mode == 'tab':67        t_steps, beta_0, beta_1 = edm2t(t_steps)68        C = []69        for i, (t_cur, t_next) in enumerate(zip(t_steps[:-1], t_steps[1:])):70            order = min(i+1, max_order)71            if order == 1:72                C.append([])73            else:74                taus = torch.linspace(t_cur, t_next, N)   # split the interval for integral appximation75                dtau = (t_next - t_cur) / N76                prev_t = t_steps[[i - k for k in range(order)]]77                coeff_temp = []78                integrand = cal_intergrand(beta_0, beta_1, taus)79                for j in range(order):80                    poly = cal_poly(prev_t, j, taus)81                    coeff_temp.append(torch.sum(integrand * poly) * dtau)82                C.append(coeff_temp)83 84    elif deis_mode == 'rhoab':85        # Analytical solution, second order86        def get_def_intergral_2(a, b, start, end, c):87            coeff = (end**3 - start**3) / 3 - (end**2 - start**2) * (a + b) / 2 + (end - start) * a * b88            return coeff / ((c - a) * (c - b))89 90        # Analytical solution, third order91        def get_def_intergral_3(a, b, c, start, end, d):92            coeff = (end**4 - start**4) / 4 - (end**3 - start**3) * (a + b + c) / 3 \93                    + (end**2 - start**2) * (a*b + a*c + b*c) / 2 - (end - start) * a * b * c94            return coeff / ((d - a) * (d - b) * (d - c))95 96        C = []97        for i, (t_cur, t_next) in enumerate(zip(t_steps[:-1], t_steps[1:])):98            order = min(i, max_order)99            if order == 0:100                C.append([])101            else:102                prev_t = t_steps[[i - k for k in range(order+1)]]103                if order == 1:104                    coeff_cur = ((t_next - prev_t[1])**2 - (t_cur - prev_t[1])**2) / (2 * (t_cur - prev_t[1]))105                    coeff_prev1 = (t_next - t_cur)**2 / (2 * (prev_t[1] - t_cur))106                    coeff_temp = [coeff_cur, coeff_prev1]107                elif order == 2:108                    coeff_cur = get_def_intergral_2(prev_t[1], prev_t[2], t_cur, t_next, t_cur)109                    coeff_prev1 = get_def_intergral_2(t_cur, prev_t[2], t_cur, t_next, prev_t[1])110                    coeff_prev2 = get_def_intergral_2(t_cur, prev_t[1], t_cur, t_next, prev_t[2])111                    coeff_temp = [coeff_cur, coeff_prev1, coeff_prev2]112                elif order == 3:113                    coeff_cur = get_def_intergral_3(prev_t[1], prev_t[2], prev_t[3], t_cur, t_next, t_cur)114                    coeff_prev1 = get_def_intergral_3(t_cur, prev_t[2], prev_t[3], t_cur, t_next, prev_t[1])115                    coeff_prev2 = get_def_intergral_3(t_cur, prev_t[1], prev_t[3], t_cur, t_next, prev_t[2])116                    coeff_prev3 = get_def_intergral_3(t_cur, prev_t[1], prev_t[2], t_cur, t_next, prev_t[3])117                    coeff_temp = [coeff_cur, coeff_prev1, coeff_prev2, coeff_prev3]118                C.append(coeff_temp)119    return C120 121