47 lines
2.0 KiB
Python
47 lines
2.0 KiB
Python
import numpy as np
|
|
from typing import Union
|
|
import torch
|
|
from tqdm import tqdm
|
|
from einops import rearrange
|
|
|
|
|
|
# DDIM Inversion
|
|
def next_step(model_output: Union[torch.FloatTensor, np.ndarray], timestep: int,
|
|
sample: Union[torch.FloatTensor, np.ndarray], ddim_scheduler):
|
|
timestep, next_timestep = min(
|
|
timestep - ddim_scheduler.config.num_train_timesteps // ddim_scheduler.num_inference_steps, 999), timestep
|
|
alpha_prod_t = ddim_scheduler.alphas_cumprod[timestep] if timestep >= 0 else ddim_scheduler.final_alpha_cumprod
|
|
alpha_prod_t_next = ddim_scheduler.alphas_cumprod[next_timestep]
|
|
beta_prod_t = 1 - alpha_prod_t
|
|
next_original_sample = (sample - beta_prod_t ** 0.5 * model_output) / alpha_prod_t ** 0.5
|
|
next_sample_direction = (1 - alpha_prod_t_next) ** 0.5 * model_output
|
|
next_sample = alpha_prod_t_next ** 0.5 * next_original_sample + next_sample_direction
|
|
return next_sample
|
|
|
|
|
|
def get_noise_pred_single(latents, t, context, unet):
|
|
latents = rearrange(latents.squeeze(0), "c f h w -> f c h w")
|
|
noise_pred = unet(latents, t.view(1), context=context)
|
|
noise_pred = rearrange(noise_pred.unsqueeze(0), "b f c h w -> b c f h w")
|
|
return noise_pred
|
|
|
|
|
|
@torch.no_grad()
|
|
def ddim_loop(pipe, ddim_scheduler, latent, num_inv_steps, context):
|
|
uncond_embeddings, cond_embeddings = context.chunk(2)
|
|
all_latent = [latent]
|
|
latent = latent.clone().detach()
|
|
for i in tqdm(range(num_inv_steps)):
|
|
t = ddim_scheduler.timesteps[len(ddim_scheduler.timesteps) - i - 1]
|
|
noise_pred = get_noise_pred_single(latent, t, cond_embeddings, pipe.model.model.diffusion_model)
|
|
latent = next_step(noise_pred, t, latent, ddim_scheduler)
|
|
all_latent.append(latent)
|
|
return all_latent
|
|
|
|
|
|
@torch.no_grad()
|
|
def ddim_inversion(pipe, ddim_scheduler, video_latent, num_inv_steps, context):
|
|
ddim_latents = ddim_loop(pipe, ddim_scheduler, video_latent, num_inv_steps, context)
|
|
return ddim_latents
|
|
|