CoolFace
Apppublic

parson/audioEditing

sourceHugging Faceupdated 3y agoView on Hugging Face
1likes
inversion_utils.py451 linesDownload Raw Back to root
1import torch2from tqdm import tqdm3# from torchvision import transforms as T4from typing import List, Optional, Dict, Union5from models import PipelineWrapper6 7 8def mu_tilde(model, xt, x0, timestep):9    "mu_tilde(x_t, x_0) DDPM paper eq. 7"10    prev_timestep = timestep - model.scheduler.config.num_train_timesteps // model.scheduler.num_inference_steps11    alpha_prod_t_prev = model.scheduler.alphas_cumprod[prev_timestep] if prev_timestep >= 0 \12        else model.scheduler.final_alpha_cumprod13    alpha_t = model.scheduler.alphas[timestep]14    beta_t = 1 - alpha_t15    alpha_bar = model.scheduler.alphas_cumprod[timestep]16    return ((alpha_prod_t_prev ** 0.5 * beta_t) / (1-alpha_bar)) * x0 + \17        ((alpha_t**0.5 * (1-alpha_prod_t_prev)) / (1 - alpha_bar)) * xt18 19 20def sample_xts_from_x0(model, x0, num_inference_steps=50, x_prev_mode=False):21    """22    Samples from P(x_1:T|x_0)23    """24    # torch.manual_seed(43256465436)25    alpha_bar = model.model.scheduler.alphas_cumprod26    sqrt_one_minus_alpha_bar = (1-alpha_bar) ** 0.527    alphas = model.model.scheduler.alphas28    # betas = 1 - alphas29    variance_noise_shape = (30            num_inference_steps + 1,31            model.model.unet.config.in_channels,32            # model.unet.sample_size,33            # model.unet.sample_size)34            x0.shape[-2],35            x0.shape[-1])36 37    timesteps = model.model.scheduler.timesteps.to(model.device)38    t_to_idx = {int(v): k for k, v in enumerate(timesteps)}39    xts = torch.zeros(variance_noise_shape).to(x0.device)40    xts[0] = x041    x_prev = x042    for t in reversed(timesteps):43        # idx = t_to_idx[int(t)]44        idx = num_inference_steps-t_to_idx[int(t)]45        if x_prev_mode:46            xts[idx] = x_prev * (alphas[t] ** 0.5) + torch.randn_like(x0) * ((1-alphas[t]) ** 0.5)47            x_prev = xts[idx].clone()48        else:49            xts[idx] = x0 * (alpha_bar[t] ** 0.5) + torch.randn_like(x0) * sqrt_one_minus_alpha_bar[t]50    # xts = torch.cat([xts, x0 ],dim = 0)51 52    return xts53 54 55def forward_step(model, model_output, timestep, sample):56    next_timestep = min(model.scheduler.config.num_train_timesteps - 2,57                        timestep + model.scheduler.config.num_train_timesteps // model.scheduler.num_inference_steps)58 59    # 2. compute alphas, betas60    alpha_prod_t = model.scheduler.alphas_cumprod[timestep]61    # alpha_prod_t_next = self.scheduler.alphas_cumprod[next_timestep] if next_ltimestep >= 0 \62    #     else self.scheduler.final_alpha_cumprod63 64    beta_prod_t = 1 - alpha_prod_t65 66    # 3. compute predicted original sample from predicted noise also called67    # "predicted x_0" of formula (12) from https://arxiv.org/pdf/2010.02502.pdf68    pred_original_sample = (sample - beta_prod_t ** (0.5) * model_output) / alpha_prod_t ** (0.5)69 70    # 5. TODO: simple noising implementatiom71    next_sample = model.scheduler.add_noise(pred_original_sample, model_output, torch.LongTensor([next_timestep]))72    return next_sample73 74 75def inversion_forward_process(model: PipelineWrapper,76                              x0: torch.Tensor,77                              etas: Optional[float] = None,78                              prog_bar: bool = False,79                              prompts: List[str] = [""],80                              cfg_scales: List[float] = [3.5],81                              num_inference_steps: int = 50,82                              eps: Optional[float] = None,83                              cutoff_points: Optional[List[float]] = None,84                              numerical_fix: bool = False,85                              extract_h_space: bool = False,86                              extract_skipconns: bool = False,87                              x_prev_mode: bool = False):88    if len(prompts) > 1 and extract_h_space:89        raise NotImplementedError("How do you split cfg_scales for hspace? TODO")90 91    if len(prompts) > 1 or prompts[0] != "":92        text_embeddings_hidden_states, text_embeddings_class_labels, \93            text_embeddings_boolean_prompt_mask = model.encode_text(prompts)94        # text_embeddings = encode_text(model, prompt)95 96        # # classifier free guidance97        batch_size = len(prompts)98        cfg_scales_tensor = torch.ones((batch_size, *x0.shape[1:]), device=model.device, dtype=x0.dtype)99 100        # if len(prompts) > 1:101        #     if cutoff_points is None:102        #         cutoff_points = [i * 1 / batch_size for i in range(1, batch_size)]103        #     if len(cfg_scales) == 1:104        #         cfg_scales *= batch_size105        #     elif len(cfg_scales) < batch_size:106        #         raise ValueError("Not enough target CFG scales")107 108        #     cutoff_points = [int(x * cfg_scales_tensor.shape[2]) for x in cutoff_points]109        #     cutoff_points = [0, *cutoff_points, cfg_scales_tensor.shape[2]]110 111        #     for i, (start, end) in enumerate(zip(cutoff_points[:-1], cutoff_points[1:])):112        #         cfg_scales_tensor[i, :, end:] = 0113        #         cfg_scales_tensor[i, :, :start] = 0114        #         cfg_scales_tensor[i] *= cfg_scales[i]115        #         if prompts[i] == "":116        #             cfg_scales_tensor[i] = 0117        #     cfg_scales_tensor = T.functional.gaussian_blur(cfg_scales_tensor, kernel_size=15, sigma=1)118        # else:119        cfg_scales_tensor *= cfg_scales[0]120 121    uncond_embedding_hidden_states, uncond_embedding_class_lables, uncond_boolean_prompt_mask = model.encode_text([""])122    # uncond_embedding = encode_text(model, "")123    timesteps = model.model.scheduler.timesteps.to(model.device)124    variance_noise_shape = (125        num_inference_steps,126        model.model.unet.config.in_channels,127        # model.unet.sample_size,128        # model.unet.sample_size)129        x0.shape[-2],130        x0.shape[-1])131 132    if etas is None or (type(etas) in [int, float] and etas == 0):133        eta_is_zero = True134        zs = None135    else:136        eta_is_zero = False137        if type(etas) in [int, float]:138            etas = [etas]*model.model.scheduler.num_inference_steps139        xts = sample_xts_from_x0(model, x0, num_inference_steps=num_inference_steps, x_prev_mode=x_prev_mode)140        alpha_bar = model.model.scheduler.alphas_cumprod141        zs = torch.zeros(size=variance_noise_shape, device=model.device)142    hspaces = []143    skipconns = []144    t_to_idx = {int(v): k for k, v in enumerate(timesteps)}145    xt = x0146    # op = tqdm(reversed(timesteps)) if prog_bar else reversed(timesteps)147    op = tqdm(timesteps) if prog_bar else timesteps148 149    for t in op:150        # idx = t_to_idx[int(t)]151        idx = num_inference_steps - t_to_idx[int(t)] - 1152        # 1. predict noise residual153        if not eta_is_zero:154            xt = xts[idx+1][None]155 156        with torch.no_grad():157            out, out_hspace, out_skipconns = model.unet_forward(xt, timestep=t,158                                                                encoder_hidden_states=uncond_embedding_hidden_states,159                                                                class_labels=uncond_embedding_class_lables,160                                                                encoder_attention_mask=uncond_boolean_prompt_mask)161            # out = model.unet.forward(xt, timestep= t, encoder_hidden_states=uncond_embedding)162            if len(prompts) > 1 or prompts[0] != "":163                cond_out, cond_out_hspace, cond_out_skipconns = model.unet_forward(164                    xt.expand(len(prompts), -1, -1, -1), timestep=t,165                    encoder_hidden_states=text_embeddings_hidden_states,166                    class_labels=text_embeddings_class_labels,167                    encoder_attention_mask=text_embeddings_boolean_prompt_mask)168                # cond_out = model.unet.forward(xt, timestep=t, encoder_hidden_states = text_embeddings)169 170        if len(prompts) > 1 or prompts[0] != "":171            # # classifier free guidance172            noise_pred = out.sample + \173                (cfg_scales_tensor * (cond_out.sample - out.sample.expand(batch_size, -1, -1, -1))174                 ).sum(axis=0).unsqueeze(0)175            if extract_h_space or extract_skipconns:176                noise_h_space = out_hspace + cfg_scales[0] * (cond_out_hspace - out_hspace)177            if extract_skipconns:178                noise_skipconns = {k: [out_skipconns[k][j] + cfg_scales[0] *179                                       (cond_out_skipconns[k][j] - out_skipconns[k][j])180                                       for j in range(len(out_skipconns[k]))]181                                   for k in out_skipconns}182        else:183            noise_pred = out.sample184            if extract_h_space or extract_skipconns:185                noise_h_space = out_hspace186            if extract_skipconns:187                noise_skipconns = out_skipconns188        if extract_h_space or extract_skipconns:189            hspaces.append(noise_h_space)190        if extract_skipconns:191            skipconns.append(noise_skipconns)192 193        if eta_is_zero:194            # 2. compute more noisy image and set x_t -> x_t+1195            xt = forward_step(model.model, noise_pred, t, xt)196        else:197            # xtm1 =  xts[idx+1][None]198            xtm1 = xts[idx][None]199            # pred of x0200            if model.model.scheduler.config.prediction_type == 'epsilon':201                pred_original_sample = (xt - (1 - alpha_bar[t]) ** 0.5 * noise_pred) / alpha_bar[t] ** 0.5202            elif model.model.scheduler.config.prediction_type == 'v_prediction':203                pred_original_sample = (alpha_bar[t] ** 0.5) * xt - ((1 - alpha_bar[t]) ** 0.5) * noise_pred204 205            # direction to xt206            prev_timestep = t - model.model.scheduler.config.num_train_timesteps // \207                model.model.scheduler.num_inference_steps208 209            alpha_prod_t_prev = model.get_alpha_prod_t_prev(prev_timestep)210            variance = model.get_variance(t, prev_timestep)211 212            if model.model.scheduler.config.prediction_type == 'epsilon':213                radom_noise_pred = noise_pred214            elif model.model.scheduler.config.prediction_type == 'v_prediction':215                radom_noise_pred = (alpha_bar[t] ** 0.5) * noise_pred + ((1 - alpha_bar[t]) ** 0.5) * xt216 217            pred_sample_direction = (1 - alpha_prod_t_prev - etas[idx] * variance) ** (0.5) * radom_noise_pred218 219            mu_xt = alpha_prod_t_prev ** (0.5) * pred_original_sample + pred_sample_direction220 221            z = (xtm1 - mu_xt) / (etas[idx] * variance ** 0.5)222 223            zs[idx] = z224 225            # correction to avoid error accumulation226            if numerical_fix:227                xtm1 = mu_xt + (etas[idx] * variance ** 0.5)*z228            xts[idx] = xtm1229 230    if zs is not None:231        # zs[-1] = torch.zeros_like(zs[-1])232        zs[0] = torch.zeros_like(zs[0])233        # zs_cycle[0] = torch.zeros_like(zs[0])234 235    if extract_h_space:236        hspaces = torch.concat(hspaces, axis=0)237        return xt, zs, xts, hspaces238 239    if extract_skipconns:240        hspaces = torch.concat(hspaces, axis=0)241        return xt, zs, xts, hspaces, skipconns242 243    return xt, zs, xts244 245 246def reverse_step(model, model_output, timestep, sample, eta=0, variance_noise=None):247    # 1. get previous step value (=t-1)248    prev_timestep = timestep - model.model.scheduler.config.num_train_timesteps // \249        model.model.scheduler.num_inference_steps250    # 2. compute alphas, betas251    alpha_prod_t = model.model.scheduler.alphas_cumprod[timestep]252    alpha_prod_t_prev = model.get_alpha_prod_t_prev(prev_timestep)253    beta_prod_t = 1 - alpha_prod_t254    # 3. compute predicted original sample from predicted noise also called255    # "predicted x_0" of formula (12) from https://arxiv.org/pdf/2010.02502.pdf256    if model.model.scheduler.config.prediction_type == 'epsilon':257        pred_original_sample = (sample - beta_prod_t ** (0.5) * model_output) / alpha_prod_t ** (0.5)258    elif model.model.scheduler.config.prediction_type == 'v_prediction':259        pred_original_sample = (alpha_prod_t ** 0.5) * sample - (beta_prod_t ** 0.5) * model_output260 261    # 5. compute variance: "sigma_t(η)" -> see formula (16)262    # σ_t = sqrt((1 − α_t−1)/(1 − α_t)) * sqrt(1 − α_t/α_t−1)263    # variance = self.scheduler._get_variance(timestep, prev_timestep)264    variance = model.get_variance(timestep, prev_timestep)265    # std_dev_t = eta * variance ** (0.5)266    # Take care of asymetric reverse process (asyrp)267    if model.model.scheduler.config.prediction_type == 'epsilon':268        model_output_direction = model_output269    elif model.model.scheduler.config.prediction_type == 'v_prediction':270        model_output_direction = (alpha_prod_t**0.5) * model_output + (beta_prod_t**0.5) * sample271    # 6. compute "direction pointing to x_t" of formula (12) from https://arxiv.org/pdf/2010.02502.pdf272    # pred_sample_direction = (1 - alpha_prod_t_prev - std_dev_t**2) ** (0.5) * model_output_direction273    pred_sample_direction = (1 - alpha_prod_t_prev - eta * variance) ** (0.5) * model_output_direction274    # 7. compute x_t without "random noise" of formula (12) from https://arxiv.org/pdf/2010.02502.pdf275    prev_sample = alpha_prod_t_prev ** (0.5) * pred_original_sample + pred_sample_direction276    # 8. Add noice if eta > 0277    if eta > 0:278        if variance_noise is None:279            variance_noise = torch.randn(model_output.shape, device=model.device)280        sigma_z = eta * variance ** (0.5) * variance_noise281        prev_sample = prev_sample + sigma_z282 283    return prev_sample284 285 286def inversion_reverse_process(model: PipelineWrapper,287                              xT: torch.Tensor,288                              skips: torch.Tensor,289                              fix_alpha: float = 0.1,290                              etas: float = 0,291                              prompts: List[str] = [""],292                              neg_prompts: List[str] = [""],293                              cfg_scales: Optional[List[float]] = None,294                              prog_bar: bool = False,295                              zs: Optional[List[torch.Tensor]] = None,296                            #   controller=None,297                              cutoff_points: Optional[List[float]] = None,298                              hspace_add: Optional[torch.Tensor] = None,299                              hspace_replace: Optional[torch.Tensor] = None,300                              skipconns_replace: Optional[Dict[int, torch.Tensor]] = None,301                              zero_out_resconns: Optional[Union[int, List]] = None,302                              asyrp: bool = False,303                              extract_h_space: bool = False,304                              extract_skipconns: bool = False):305 306    batch_size = len(prompts)307 308    text_embeddings_hidden_states, text_embeddings_class_labels, \309        text_embeddings_boolean_prompt_mask = model.encode_text(prompts)310    uncond_embedding_hidden_states, uncond_embedding_class_lables, \311        uncond_boolean_prompt_mask = model.encode_text(neg_prompts)312    # text_embeddings = encode_text(model, prompts)313    # uncond_embedding = encode_text(model, [""] * batch_size)314 315    masks = torch.ones((batch_size, *xT.shape[1:]), device=model.device, dtype=xT.dtype)316    cfg_scales_tensor = torch.ones((batch_size, *xT.shape[1:]), device=model.device, dtype=xT.dtype)317 318    # if batch_size > 1:319    #     if cutoff_points is None:320    #         cutoff_points = [i * 1 / batch_size for i in range(1, batch_size)]321    #     if len(cfg_scales) == 1:322    #         cfg_scales *= batch_size323    #     elif len(cfg_scales) < batch_size:324    #         raise ValueError("Not enough target CFG scales")325 326    #     cutoff_points = [int(x * cfg_scales_tensor.shape[2]) for x in cutoff_points]327    #     cutoff_points = [0, *cutoff_points, cfg_scales_tensor.shape[2]]328 329    #     for i, (start, end) in enumerate(zip(cutoff_points[:-1], cutoff_points[1:])):330    #         cfg_scales_tensor[i, :, end:] = 0331    #         cfg_scales_tensor[i, :, :start] = 0332    #         masks[i, :, end:] = 0333    #         masks[i, :, :start] = 0334    #         cfg_scales_tensor[i] *= cfg_scales[i]335    #     cfg_scales_tensor = T.functional.gaussian_blur(cfg_scales_tensor, kernel_size=15, sigma=1)336    #     masks = T.functional.gaussian_blur(masks, kernel_size=15, sigma=1)337    # else:338    cfg_scales_tensor *= cfg_scales[0]339 340    if etas is None:341        etas = 0342    if type(etas) in [int, float]:343        etas = [etas]*model.model.scheduler.num_inference_steps344    assert len(etas) == model.model.scheduler.num_inference_steps345    timesteps = model.model.scheduler.timesteps.to(model.device)346 347    # xt = xT.expand(1, -1, -1, -1)348    xt = xT[skips.max()].unsqueeze(0)349    op = tqdm(timesteps[-zs.shape[0]:]) if prog_bar else timesteps[-zs.shape[0]:]350 351    t_to_idx = {int(v): k for k, v in enumerate(timesteps[-zs.shape[0]:])}352    hspaces = []353    skipconns = []354 355    for it, t in enumerate(op):356        # idx = t_to_idx[int(t)]357        idx = model.model.scheduler.num_inference_steps - t_to_idx[int(t)] - \358            (model.model.scheduler.num_inference_steps - zs.shape[0] + 1)359        # # Unconditional embedding360        with torch.no_grad():361            uncond_out, out_hspace, out_skipconns = model.unet_forward(362                xt, timestep=t,363                encoder_hidden_states=uncond_embedding_hidden_states,364                class_labels=uncond_embedding_class_lables,365                encoder_attention_mask=uncond_boolean_prompt_mask,366                mid_block_additional_residual=(None if hspace_add is None else367                                               (1 / (cfg_scales[0] + 1)) *368                                               (hspace_add[-zs.shape[0]:][it] if hspace_add.shape[0] > 1369                                                else hspace_add)),370                replace_h_space=(None if hspace_replace is None else371                                 (hspace_replace[-zs.shape[0]:][it].unsqueeze(0) if hspace_replace.shape[0] > 1372                                  else hspace_replace)),373                zero_out_resconns=zero_out_resconns,374                replace_skip_conns=(None if skipconns_replace is None else375                                    (skipconns_replace[-zs.shape[0]:][it] if len(skipconns_replace) > 1376                                     else skipconns_replace))377                )  # encoder_hidden_states = uncond_embedding)378 379        # # Conditional embedding380        if prompts:381            with torch.no_grad():382                cond_out, cond_out_hspace, cond_out_skipconns = model.unet_forward(383                    xt.expand(batch_size, -1, -1, -1),384                    timestep=t,385                    encoder_hidden_states=text_embeddings_hidden_states,386                    class_labels=text_embeddings_class_labels,387                    encoder_attention_mask=text_embeddings_boolean_prompt_mask,388                    mid_block_additional_residual=(None if hspace_add is None else389                                                   (cfg_scales[0] / (cfg_scales[0] + 1)) *390                                                   (hspace_add[-zs.shape[0]:][it] if hspace_add.shape[0] > 1391                                                    else hspace_add)),392                    replace_h_space=(None if hspace_replace is None else393                                     (hspace_replace[-zs.shape[0]:][it].unsqueeze(0) if hspace_replace.shape[0] > 1394                                      else hspace_replace)),395                    zero_out_resconns=zero_out_resconns,396                    replace_skip_conns=(None if skipconns_replace is None else397                                        (skipconns_replace[-zs.shape[0]:][it] if len(skipconns_replace) > 1398                                         else skipconns_replace))399                    )  # encoder_hidden_states = text_embeddings)400 401        z = zs[idx] if zs is not None else None402        # print(f'idx: {idx}')403        # print(f't: {t}')404        z = z.unsqueeze(0)405        # z = z.expand(batch_size, -1, -1, -1)406        if prompts:407            # # classifier free guidance408            # noise_pred = uncond_out.sample + cfg_scales_tensor * (cond_out.sample - uncond_out.sample)409            noise_pred = uncond_out.sample + \410                (cfg_scales_tensor * (cond_out.sample - uncond_out.sample.expand(batch_size, -1, -1, -1))411                 ).sum(axis=0).unsqueeze(0)412            if extract_h_space or extract_skipconns:413                noise_h_space = out_hspace + cfg_scales[0] * (cond_out_hspace - out_hspace)414            if extract_skipconns:415                noise_skipconns = {k: [out_skipconns[k][j] + cfg_scales[0] *416                                       (cond_out_skipconns[k][j] - out_skipconns[k][j])417                                       for j in range(len(out_skipconns[k]))]418                                   for k in out_skipconns}419        else:420            noise_pred = uncond_out.sample421            if extract_h_space or extract_skipconns:422                noise_h_space = out_hspace423            if extract_skipconns:424                noise_skipconns = out_skipconns425 426        if extract_h_space or extract_skipconns:427            hspaces.append(noise_h_space)428        if extract_skipconns:429            skipconns.append(noise_skipconns)430 431        # 2. compute less noisy image and set x_t -> x_t-1432        xt = reverse_step(model, noise_pred, t, xt, eta=etas[idx], variance_noise=z)433        # if controller is not None:434            # xt = controller.step_callback(xt)435 436        # "fix" xt437        apply_fix = ((skips.max() - skips) > it)438        if apply_fix.any():439            apply_fix = (apply_fix * fix_alpha).unsqueeze(1).unsqueeze(2).unsqueeze(3).to(xT.device)440            xt = (masks * (xt.expand(batch_size, -1, -1, -1) * (1 - apply_fix) +441                           apply_fix * xT[skips.max() - it - 1].expand(batch_size, -1, -1, -1))442                  ).sum(axis=0).unsqueeze(0)443 444    if extract_h_space:445        return xt, zs, torch.concat(hspaces, axis=0)446 447    if extract_skipconns:448        return xt, zs, torch.concat(hspaces, axis=0), skipconns449 450    return xt, zs451