codejin/diffsingerkr
6
1import torch2import math3from argparse import Namespace4from typing import Optional, List, Dict, Union5from tqdm import tqdm6 7from .Layer import Conv1d, Lambda8 9class Diffusion(torch.nn.Module):10 def __init__(11 self,12 hyper_parameters: Namespace13 ):14 super().__init__()15 self.hp = hyper_parameters16 17 if self.hp.Feature_Type == 'Mel':18 self.feature_size = self.hp.Sound.Mel_Dim19 elif self.hp.Feature_Type == 'Spectrogram':20 self.feature_size = self.hp.Sound.N_FFT // 2 + 121 22 self.denoiser = Denoiser(23 hyper_parameters= self.hp24 )25 26 self.timesteps = self.hp.Diffusion.Max_Step27 betas = torch.linspace(1e-4, 0.06, self.timesteps)28 alphas = 1.0 - betas29 alphas_cumprod = torch.cumprod(alphas, axis= 0)30 alphas_cumprod_prev = torch.cat([torch.tensor([1.0]), alphas_cumprod[:-1]])31 32 # calculations for diffusion q(x_t | x_{t-1}) and others33 self.register_buffer('alphas_cumprod', alphas_cumprod) # [Diffusion_t]34 self.register_buffer('alphas_cumprod_prev', alphas_cumprod_prev) # [Diffusion_t]35 self.register_buffer('sqrt_alphas_cumprod', alphas_cumprod.sqrt())36 self.register_buffer('sqrt_one_minus_alphas_cumprod', (1.0 - alphas_cumprod).sqrt())37 self.register_buffer('sqrt_recip_alphas_cumprod', (1.0 / alphas_cumprod).sqrt())38 self.register_buffer('sqrt_recipm1_alphas_cumprod', (1.0 / alphas_cumprod - 1.0).sqrt())39 40 # calculations for posterior q(x_{t-1} | x_t, x_0)41 posterior_variance = betas * (1.0 - alphas_cumprod_prev) / (1.0 - alphas_cumprod)42 43 # below: log calculation clipped because the posterior variance is 0 at the beginning of the diffusion chain44 self.register_buffer('posterior_log_variance', torch.maximum(posterior_variance, torch.tensor([1e-20])).log())45 self.register_buffer('posterior_mean_coef1', betas * alphas_cumprod_prev.sqrt() / (1.0 - alphas_cumprod))46 self.register_buffer('posterior_mean_coef2', (1.0 - alphas_cumprod_prev) * alphas.sqrt() / (1.0 - alphas_cumprod))47 48 def forward(49 self,50 encodings: torch.Tensor,51 features: torch.Tensor= None52 ):53 '''54 encodings: [Batch, Enc_d, Enc_t]55 features: [Batch, Feature_d, Feature_t]56 feature_lengths: [Batch]57 '''58 if not features is None: # train59 diffusion_steps = torch.randint(60 low= 0,61 high= self.timesteps,62 size= (encodings.size(0),),63 dtype= torch.long,64 device= encodings.device65 ) # random single step66 67 noises, epsilons = self.Get_Noise_Epsilon_for_Train(68 features= features,69 encodings= encodings,70 diffusion_steps= diffusion_steps,71 )72 return None, noises, epsilons73 else: # inference74 features = self.Sampling(75 encodings= encodings,76 )77 return features, None, None78 79 def Sampling(80 self,81 encodings: torch.Tensor,82 ):83 features = torch.randn(84 size= (encodings.size(0), self.feature_size, encodings.size(2)),85 device= encodings.device86 )87 for diffusion_step in reversed(range(self.timesteps)):88 features = self.P_Sampling(89 features= features,90 encodings= encodings,91 diffusion_steps= torch.full(92 size= (encodings.size(0), ),93 fill_value= diffusion_step,94 dtype= torch.long,95 device= encodings.device96 ),97 )98 99 return features100 101 def P_Sampling(102 self,103 features: torch.Tensor,104 encodings: torch.Tensor,105 diffusion_steps: torch.Tensor,106 ):107 posterior_means, posterior_log_variances = self.Get_Posterior(108 features= features,109 encodings= encodings,110 diffusion_steps= diffusion_steps,111 )112 113 noises = torch.randn_like(features) # [Batch, Feature_d, Feature_d]114 masks = (diffusion_steps > 0).float().unsqueeze(1).unsqueeze(1) #[Batch, 1, 1]115 116 return posterior_means + masks * (0.5 * posterior_log_variances).exp() * noises117 118 def Get_Posterior(119 self,120 features: torch.Tensor,121 encodings: torch.Tensor,122 diffusion_steps: torch.Tensor123 ):124 noised_predictions = self.denoiser(125 features= features,126 encodings= encodings,127 diffusion_steps= diffusion_steps128 )129 130 epsilons = \131 features * self.sqrt_recip_alphas_cumprod[diffusion_steps][:, None, None] - \132 noised_predictions * self.sqrt_recipm1_alphas_cumprod[diffusion_steps][:, None, None]133 epsilons.clamp_(-1.0, 1.0) # clipped134 135 posterior_means = \136 epsilons * self.posterior_mean_coef1[diffusion_steps][:, None, None] + \137 features * self.posterior_mean_coef2[diffusion_steps][:, None, None]138 posterior_log_variances = \139 self.posterior_log_variance[diffusion_steps][:, None, None]140 141 return posterior_means, posterior_log_variances142 143 def Get_Noise_Epsilon_for_Train(144 self,145 features: torch.Tensor,146 encodings: torch.Tensor,147 diffusion_steps: torch.Tensor,148 ):149 noises = torch.randn_like(features)150 151 noised_features = \152 features * self.sqrt_alphas_cumprod[diffusion_steps][:, None, None] + \153 noises * self.sqrt_one_minus_alphas_cumprod[diffusion_steps][:, None, None]154 155 epsilons = self.denoiser(156 features= noised_features,157 encodings= encodings,158 diffusion_steps= diffusion_steps159 )160 161 return noises, epsilons162 163 def DDIM(164 self,165 encodings: torch.Tensor,166 ddim_steps: int,167 eta: float= 0.0,168 temperature: float= 1.0,169 use_tqdm: bool= False170 ):171 ddim_timesteps = self.Get_DDIM_Steps(172 ddim_steps= ddim_steps173 )174 sigmas, alphas, alphas_prev = self.Get_DDIM_Sampling_Parameters(175 ddim_timesteps= ddim_timesteps,176 eta= eta177 )178 sqrt_one_minus_alphas = (1. - alphas).sqrt()179 180 features = torch.randn(181 size= (encodings.size(0), self.feature_size, encodings.size(2)),182 device= encodings.device183 )184 185 setp_range = reversed(range(ddim_steps))186 if use_tqdm:187 tqdm(188 setp_range,189 desc= '[Diffusion]',190 total= ddim_steps191 )192 193 for diffusion_steps in setp_range:194 noised_predictions = self.denoiser(195 features= features,196 encodings= encodings,197 diffusion_steps= torch.full(198 size= (encodings.size(0), ),199 fill_value= diffusion_steps,200 dtype= torch.long,201 device= encodings.device202 )203 )204 205 feature_starts = (features - sqrt_one_minus_alphas[diffusion_steps] * noised_predictions) / alphas[diffusion_steps].sqrt()206 direction_pointings = (1.0 - alphas_prev[diffusion_steps] - sigmas[diffusion_steps].pow(2.0)) * noised_predictions207 noises = sigmas[diffusion_steps] * torch.randn_like(features) * temperature208 209 features = alphas_prev[diffusion_steps].sqrt() * feature_starts + direction_pointings + noises210 211 return features212 213 # https://github.com/CompVis/stable-diffusion/blob/main/ldm/modules/diffusionmodules/util.py214 def Get_DDIM_Steps(215 self, 216 ddim_steps: int,217 ddim_discr_method: str= 'uniform'218 ):219 if ddim_discr_method == 'uniform': 220 ddim_timesteps = torch.arange(0, self.timesteps, self.timesteps // ddim_steps).long()221 elif ddim_discr_method == 'quad':222 ddim_timesteps = torch.linspace(0, (torch.tensor(self.timesteps) * 0.8).sqrt(), ddim_steps).pow(2.0).long()223 else:224 raise NotImplementedError(f'There is no ddim discretization method called "{ddim_discr_method}"')225 226 ddim_timesteps[-1] = self.timesteps - 1227 228 return ddim_timesteps229 230 def Get_DDIM_Sampling_Parameters(self, ddim_timesteps, eta):231 alphas = self.alphas_cumprod[ddim_timesteps]232 alphas_prev = self.alphas_cumprod_prev[ddim_timesteps]233 sigmas = eta * ((1 - alphas_prev) / (1 - alphas) * (1 - alphas / alphas_prev)).sqrt()234 235 return sigmas, alphas, alphas_prev236 237class Denoiser(torch.nn.Module):238 def __init__(239 self,240 hyper_parameters: Namespace241 ):242 super().__init__()243 self.hp = hyper_parameters244 245 if self.hp.Feature_Type == 'Mel':246 feature_size = self.hp.Sound.Mel_Dim247 elif self.hp.Feature_Type == 'Spectrogram':248 feature_size = self.hp.Sound.N_FFT // 2 + 1249 250 self.prenet = torch.nn.Sequential(251 Conv1d(252 in_channels= feature_size,253 out_channels= self.hp.Diffusion.Size,254 kernel_size= 1,255 w_init_gain= 'relu'256 ),257 torch.nn.Mish()258 )259 260 self.step_ffn = torch.nn.Sequential(261 Diffusion_Embedding(262 channels= self.hp.Diffusion.Size263 ),264 Lambda(lambda x: x.unsqueeze(2)),265 Conv1d(266 in_channels= self.hp.Diffusion.Size,267 out_channels= self.hp.Diffusion.Size * 4,268 kernel_size= 1,269 w_init_gain= 'relu'270 ),271 torch.nn.Mish(),272 Conv1d(273 in_channels= self.hp.Diffusion.Size * 4,274 out_channels= self.hp.Diffusion.Size,275 kernel_size= 1,276 w_init_gain= 'linear'277 )278 )279 280 self.residual_blocks = torch.nn.ModuleList([281 Residual_Block(282 in_channels= self.hp.Diffusion.Size,283 kernel_size= self.hp.Diffusion.Kernel_Size,284 condition_channels= self.hp.Encoder.Size + feature_size285 )286 for _ in range(self.hp.Diffusion.Stack)287 ])288 289 self.projection = torch.nn.Sequential(290 Conv1d(291 in_channels= self.hp.Diffusion.Size,292 out_channels= self.hp.Diffusion.Size,293 kernel_size= 1,294 w_init_gain= 'relu'295 ),296 torch.nn.ReLU(),297 Conv1d(298 in_channels= self.hp.Diffusion.Size,299 out_channels= feature_size,300 kernel_size= 1301 ),302 )303 torch.nn.init.zeros_(self.projection[-1].weight) # This is key factor....304 305 def forward(306 self,307 features: torch.Tensor,308 encodings: torch.Tensor,309 diffusion_steps: torch.Tensor310 ):311 '''312 features: [Batch, Feature_d, Feature_t]313 encodings: [Batch, Enc_d, Feature_t]314 diffusion_steps: [Batch]315 '''316 x = self.prenet(features)317 318 diffusion_steps = self.step_ffn(diffusion_steps) # [Batch, Res_d, 1]319 320 skips_list = []321 for residual_block in self.residual_blocks:322 x, skips = residual_block(323 x= x,324 conditions= encodings,325 diffusion_steps= diffusion_steps326 )327 skips_list.append(skips)328 329 x = torch.stack(skips_list, dim= 0).sum(dim= 0) / math.sqrt(self.hp.Diffusion.Stack)330 x = self.projection(x)331 332 return x333 334class Diffusion_Embedding(torch.nn.Module):335 def __init__(336 self,337 channels: int338 ):339 super().__init__()340 self.channels = channels341 342 def forward(self, x: torch.Tensor):343 half_channels = self.channels // 2 # sine and cosine344 embeddings = math.log(10000.0) / (half_channels - 1)345 embeddings = torch.exp(torch.arange(half_channels, device= x.device) * -embeddings)346 embeddings = x.unsqueeze(1) * embeddings.unsqueeze(0)347 embeddings = torch.cat([embeddings.sin(), embeddings.cos()], dim= -1)348 349 return embeddings350 351class Residual_Block(torch.nn.Module):352 def __init__(353 self,354 in_channels: int,355 kernel_size: int,356 condition_channels: int357 ):358 super().__init__()359 self.in_channels = in_channels360 361 self.condition = Conv1d(362 in_channels= condition_channels,363 out_channels= in_channels * 2,364 kernel_size= 1365 )366 self.diffusion_step = Conv1d(367 in_channels= in_channels,368 out_channels= in_channels,369 kernel_size= 1370 )371 372 self.conv = Conv1d(373 in_channels= in_channels,374 out_channels= in_channels * 2,375 kernel_size= kernel_size,376 padding= kernel_size // 2377 )378 379 self.projection = Conv1d(380 in_channels= in_channels,381 out_channels= in_channels * 2,382 kernel_size= 1383 )384 385 def forward(386 self,387 x: torch.Tensor,388 conditions: torch.Tensor,389 diffusion_steps: torch.Tensor390 ):391 residuals = x392 393 conditions = self.condition(conditions)394 diffusion_steps = self.diffusion_step(diffusion_steps)395 396 x = self.conv(x + diffusion_steps) + conditions397 x_a, x_b = x.chunk(chunks= 2, dim= 1)398 x = x_a.sigmoid() * x_b.tanh()399 400 x = self.projection(x)401 x, skips = x.chunk(chunks= 2, dim= 1)402 403 return (x + residuals) / math.sqrt(2.0), skips