Spaces:
Running on Zero
Running on Zero
File size: 7,299 Bytes
de6f21d | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 | import torch
import numpy as np
import utils.training_utils as utils
class EDM():
"""
Definition of most of the diffusion parameterization, following ( Karras et al., "Elucidating...", 2022)
"""
def __init__(self, args):
"""
Args:
args (dictionary): hydra arguments
sigma_data (float):
"""
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 #depends on the training data!! precalculated variance of the dataset
#parameters stochastic sampling
self.Schurn=args.diff_params.Schurn
self.Stmin=args.diff_params.Stmin
self.Stmax=args.diff_params.Stmax
self.Snoise=args.diff_params.Snoise
#perceptual filter
if self.args.diff_params.aweighting.use_aweighting:
self.AW=utils.FIRFilter(filter_type="aw", fs=args.exp.sample_rate, ntaps=self.args.diff_params.aweighting.ntaps)
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)
#If desired, only apply stochasticity between a certain range of noises Stmin is 0 by default and Stmax is a huge number by default. (Unless these parameters are specified, this does nothing)
indexes=torch.logical_and(t>self.Stmin , t<self.Stmax)
#We use Schurn=5 as the default in our experiments
gamma[indexes]=gamma[indexes]+torch.min(torch.Tensor([self.Schurn/N, 2**(1/2) -1]))
return gamma
def create_schedule(self,nb_steps):
"""
Define the schedule of timesteps
Args:
nb_steps (int): Number of discretized 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 sample_ptrain(self,N):
"""
For training, getting t as a normal distribution, folowing Karras et al.
I'm not using this
Args:
N (int): batch size
"""
lnsigma=np.random.randn(N)*self.P_std +self.P_mean
return np.clip(np.exp(lnsigma),self.sigma_min, self.sigma_max) #not sure if clipping here is necessary, but makes sense to me
def sample_ptrain_safe(self,N):
"""
For training, getting t according to the same criteria as sampling
Args:
N (int): batch size
"""
a=torch.rand(N)
t=(self.sigma_max**(1/self.ro_train) +a *(self.sigma_min**(1/self.ro_train) - self.sigma_max**(1/self.ro_train)))**self.ro_train
return t
def sample_prior(self,shape,sigma):
"""
Just sample some gaussian noise, nothing more
Args:
shape (tuple): shape of the noise to sample, something like (B,T)
sigma (float): noise level of the noise
"""
n=torch.randn(shape).to(sigma.device)*sigma
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):
return (sigma*self.sigma_data)**(-2) * (self.sigma_data**2+sigma**2)
def denoiser(self, xn , net, sigma):
"""
This method does the whole denoising step, which implies applying the model and the preconditioning
Args:
x (Tensor): shape: (B,T) Intermediate noisy latent to denoise
model (nn.Module): Model of the denoiser
sigma (float): noise level (equal to timestep is sigma=t, which is our default)
"""
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) #this will crash because of broadcasting problems, debug later!
def prepare_train_preconditioning(self, x, sigma):
#weight=self.lambda_w(sigma)
#Is calling the denoiser here a good idea? Maybe it would be better to apply directly the preconditioning as in the paper, even though Karras et al seem to do it this way in their code
print(x.shape)
noise=self.sample_prior(x.shape,sigma)
cskip=self.cskip(sigma)
cout=self.cout(sigma)
cin=self.cin(sigma)
cnoise=self.cnoise(sigma)
target=(1/cout)*(x-cskip*(x+noise))
return cin*(x+noise), target, cnoise
def loss_fn(self, net, x):
"""
Loss function, which is the mean squared error between the denoised latent and the clean latent
Args:
net (nn.Module): Model of the denoiser
x (Tensor): shape: (B,T) Intermediate noisy latent to denoise
sigma (float): noise level (equal to timestep is sigma=t, which is our default)
"""
sigma=self.sample_ptrain_safe(x.shape[0]).unsqueeze(-1).to(x.device)
input, target, cnoise= self.prepare_train_preconditioning(x, sigma)
estimate=net(input,cnoise)
error=(estimate-target)
try:
#this will only happen if the model is cqt-based, if it crashes it is normal
if self.args.net.use_cqt_DC_correction:
error=net.CQTransform.apply_hpf_DC(error) #apply the DC correction to the error as we dont want to propagate the DC component of the error as the network is discarding it. It also applies for the nyquit frequency, but this is less critical.
except:
pass
#APPLY A-WEIGHTING
if self.args.diff_params.aweighting.use_aweighting:
error=self.AW(error)
#here we have the chance to apply further emphasis to the error, as some kind of perceptual frequency weighting could be
return error**2, sigma
|