Spaces:
Running on Zero
Running on Zero
| 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) | |