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)