BUDDy / diff_params /edm.py
Vansh Chugh
initial deploy
a95f6c0
Raw
History Blame Contribute Delete
2.93 kB
import torch
from diff_params.shared import SDE
class EDM(SDE):
"""
Definition of the diffusion parameterization, following ( Karras et al., "Elucidating...", 2022).
This includes only the utilities needed for training, not for sampling.
"""
def __init__(self,
type,
sde_hp):
super().__init__(type, sde_hp)
self.sigma_data = self.sde_hp.sigma_data #depends on the training data!! precalculated variance of the dataset
self.sigma_min = self.sde_hp.sigma_min
self.sigma_max = self.sde_hp.sigma_max
self.rho = self.sde_hp.rho
def sample_time_training(self, N):
"""
For training, getting t according to a similar criteria as sampling. Simpler and safer to what Karras et al. did
Args:
N (int): batch size
"""
a = torch.rand(N)
t = (self.sigma_max**(1/self.rho) +a *(self.sigma_min**(1/self.rho) - self.sigma_max**(1/self.rho)))**self.rho
return t
def sample_prior(self, shape):
"""
Just sample some gaussian noise, nothing more
Args:
shape (tuple): shape of the noise to sample, something like (B,T)
"""
n = torch.randn(shape)
return n
def cskip(self, sigma):
"""
Just one of the preconditioning parameters
Args:
sigma (float): noise level (equal to timestep is sigma=t, which is our default)
"""
return self.sigma_data**2 *(sigma**2+self.sigma_data**2)**-1
def cout(self, sigma):
"""
Just one of the preconditioning parameters
Args:
sigma (float): noise level (equal to timestep is sigma=t, which is our default)
"""
return sigma*self.sigma_data* (self.sigma_data**2+sigma**2)**(-0.5)
def cin(self, sigma):
"""
Just one of the preconditioning parameters
Args:
sigma (float): noise level (equal to timestep is sigma=t, which is our default)
"""
return (self.sigma_data**2+sigma**2)**(-0.5)
def cnoise(self, sigma):
"""
preconditioning of the noise embedding
Args:
sigma (float): noise level (equal to timestep is sigma=t, which is our default)
"""
return (1/4)*torch.log(sigma)
def lambda_w(self, sigma):
"""
Score matching loss weighting
"""
return (sigma*self.sigma_data)**(-2) * (self.sigma_data**2+sigma**2)
def Tweedie2score(self, tweedie, xt, t, *args, **kwargs):
return (tweedie - self._mean(xt, t)) / self._std(t)**2
def score2Tweedie(self, score, xt, t, *args, **kwargs):
return self._std(t)**2 * score + self._mean(xt, t)
def _mean(self, x, t):
return x
def _std(self, t):
return t
def _ode_integrand(self, x, t, score):
return -t * score