BABE-2 / model /edm.py
Vansh Chugh
initial deploy
6dea0da
Raw
History Blame Contribute Delete
2.86 kB
import torch
class EDM():
"""
Definition of most of the diffusion parameterization, following (Karras et al., "Elucidating...", 2022)
"""
def __init__(self, args):
self.args = args
self.sigma_min = args.diff_params.sigma_min
self.sigma_max = args.diff_params.sigma_max
self.P_mean = args.diff_params.P_mean
self.P_std = args.diff_params.P_std
self.ro = args.diff_params.ro
self.ro_train = args.diff_params.ro_train
self.sigma_data = args.diff_params.sigma_data
self.Schurn = args.diff_params.Schurn
self.Stmin = args.diff_params.Stmin
self.Stmax = args.diff_params.Stmax
self.Snoise = args.diff_params.Snoise
def get_gamma(self, t):
"""
Get the parameter gamma that defines the stochasticity of the sampler
Args
t (Tensor): shape: (N_steps, ) Tensor of timesteps, from which we will compute gamma
"""
N = t.shape[0]
gamma = torch.zeros(t.shape).to(t.device)
indexes = torch.logical_and(t > self.Stmin, t < self.Stmax)
gamma[indexes] = gamma[indexes] + torch.min(torch.Tensor([self.Schurn / N, 2**(1 / 2) - 1]))
return gamma
def create_schedule(self, nb_steps):
i = torch.arange(0, nb_steps + 1)
t = (self.sigma_max**(1 / self.ro) + i / (nb_steps - 1) * (self.sigma_min**(1 / self.ro) - self.sigma_max**(1 / self.ro)))**self.ro
t[-1] = 0
return t
def create_schedule_from_initial_t(self, initial_t, nb_steps):
i = torch.arange(0, nb_steps + 1)
t = (initial_t**(1 / self.ro) + i / (nb_steps - 1) * (self.sigma_min**(1 / self.ro) - initial_t**(1 / self.ro)))**self.ro
t[-1] = 0
return t
def sample_prior(self, shape, sigma):
n = torch.randn(shape).to(sigma.device) * sigma
return n
def cskip(self, sigma):
return self.sigma_data**2 * (sigma**2 + self.sigma_data**2)**-1
def cout(self, sigma):
return sigma * self.sigma_data * (self.sigma_data**2 + sigma**2)**(-0.5)
def cin(self, sigma):
return (self.sigma_data**2 + sigma**2)**(-0.5)
def cnoise(self, sigma):
return (1 / 4) * torch.log(sigma)
def denoiser(self, xn, net, sigma):
"""
Applies the model and its EDM preconditioning to produce a denoised estimate.
Args:
xn (Tensor): shape (B,T) noisy latent to denoise
net (nn.Module): the diffusion prior network
sigma (Tensor): noise level (equal to timestep t)
"""
if len(sigma.shape) == 1:
sigma = sigma.unsqueeze(-1)
cskip = self.cskip(sigma)
cout = self.cout(sigma)
cin = self.cin(sigma)
cnoise = self.cnoise(sigma)
return cskip * xn + cout * net(cin * xn, cnoise)