codejin/diffsingerkr
6
1from argparse import Namespace2import torch3import math4from typing import Union5 6from .Layer import Conv1d, LayerNorm, LinearAttention7from .Diffusion import Diffusion8 9class DiffSinger(torch.nn.Module):10 def __init__(self, hyper_parameters: Namespace):11 super().__init__()12 self.hp = hyper_parameters13 14 self.encoder = Encoder(self.hp)15 self.diffusion = Diffusion(self.hp)16 17 def forward(18 self,19 tokens: torch.LongTensor,20 notes: torch.LongTensor,21 durations: torch.LongTensor,22 lengths: torch.LongTensor,23 genres: torch.LongTensor,24 singers: torch.LongTensor,25 features: Union[torch.FloatTensor, None]= None,26 ddim_steps: Union[int, None]= None27 ):28 encodings, linear_predictions = self.encoder(29 tokens= tokens,30 notes= notes,31 durations= durations,32 lengths= lengths,33 genres= genres,34 singers= singers35 ) # [Batch, Enc_d, Feature_t]36 37 encodings = torch.cat([encodings, linear_predictions], dim= 1) # [Batch, Enc_d + Feature_d, Feature_t]38 39 if not features is None or ddim_steps is None or ddim_steps == self.hp.Diffusion.Max_Step:40 diffusion_predictions, noises, epsilons = self.diffusion(41 encodings= encodings,42 features= features,43 )44 else:45 noises, epsilons = None, None46 diffusion_predictions = self.diffusion.DDIM(47 encodings= encodings,48 ddim_steps= ddim_steps49 )50 51 return linear_predictions, diffusion_predictions, noises, epsilons52 53 54class Encoder(torch.nn.Module): 55 def __init__(56 self,57 hyper_parameters: Namespace58 ):59 super().__init__()60 self.hp = hyper_parameters61 62 if self.hp.Feature_Type == 'Mel':63 self.feature_size = self.hp.Sound.Mel_Dim64 elif self.hp.Feature_Type == 'Spectrogram':65 self.feature_size = self.hp.Sound.N_FFT // 2 + 166 67 self.token_embedding = torch.nn.Embedding(68 num_embeddings= self.hp.Tokens,69 embedding_dim= self.hp.Encoder.Size70 )71 self.note_embedding = torch.nn.Embedding(72 num_embeddings= self.hp.Notes,73 embedding_dim= self.hp.Encoder.Size74 )75 self.duration_embedding = Duration_Positional_Encoding(76 num_embeddings= self.hp.Durations,77 embedding_dim= self.hp.Encoder.Size78 )79 self.genre_embedding = torch.nn.Embedding(80 num_embeddings= self.hp.Genres,81 embedding_dim= self.hp.Encoder.Size,82 )83 self.singer_embedding = torch.nn.Embedding(84 num_embeddings= self.hp.Singers,85 embedding_dim= self.hp.Encoder.Size,86 )87 torch.nn.init.xavier_uniform_(self.token_embedding.weight)88 torch.nn.init.xavier_uniform_(self.note_embedding.weight)89 torch.nn.init.xavier_uniform_(self.genre_embedding.weight)90 torch.nn.init.xavier_uniform_(self.singer_embedding.weight)91 92 self.fft_blocks = torch.nn.ModuleList([93 FFT_Block(94 channels= self.hp.Encoder.Size,95 num_head= self.hp.Encoder.ConvFFT.Head,96 ffn_kernel_size= self.hp.Encoder.ConvFFT.FFN.Kernel_Size,97 dropout_rate= self.hp.Encoder.ConvFFT.Dropout_Rate98 )99 for _ in range(self.hp.Encoder.ConvFFT.Stack) 100 ])101 102 self.linear_projection = Conv1d(103 in_channels= self.hp.Encoder.Size,104 out_channels= self.feature_size,105 kernel_size= 1,106 bias= True,107 w_init_gain= 'linear' 108 )109 110 def forward(111 self,112 tokens: torch.Tensor,113 notes: torch.Tensor,114 durations: torch.Tensor,115 lengths: torch.Tensor,116 genres: torch.Tensor,117 singers: torch.Tensor118 ):119 x = \120 self.token_embedding(tokens) + \121 self.note_embedding(notes) + \122 self.duration_embedding(durations) + \123 self.genre_embedding(genres).unsqueeze(1) + \124 self.singer_embedding(singers).unsqueeze(1)125 x = x.permute(0, 2, 1) # [Batch, Enc_d, Enc_t]126 127 for block in self.fft_blocks:128 x = block(x, lengths) # [Batch, Enc_d, Enc_t]129 130 linear_predictions = self.linear_projection(x) # [Batch, Feature_d, Enc_t]131 132 return x, linear_predictions133 134class FFT_Block(torch.nn.Module):135 def __init__(136 self,137 channels: int,138 num_head: int,139 ffn_kernel_size: int,140 dropout_rate: float= 0.1,141 ) -> None:142 super().__init__()143 144 self.attention = LinearAttention(145 channels= channels,146 calc_channels= channels,147 num_heads= num_head,148 dropout_rate= dropout_rate149 )150 151 self.ffn = FFN(152 channels= channels,153 kernel_size= ffn_kernel_size,154 dropout_rate= dropout_rate155 )156 157 def forward(158 self,159 x: torch.Tensor,160 lengths: torch.Tensor161 ) -> torch.Tensor:162 '''163 x: [Batch, Dim, Time]164 '''165 masks = (~Mask_Generate(lengths= lengths, max_length= torch.ones_like(x[0, 0]).sum())).unsqueeze(1).float() # float mask166 167 # Attention + Dropout + LayerNorm168 x = self.attention(x)169 170 # FFN + Dropout + LayerNorm171 x = self.ffn(x, masks)172 173 return x * masks174 175class FFN(torch.nn.Module):176 def __init__(177 self,178 channels: int,179 kernel_size: int,180 dropout_rate: float= 0.1,181 ) -> None:182 super().__init__()183 self.conv_0 = Conv1d(184 in_channels= channels,185 out_channels= channels,186 kernel_size= kernel_size,187 padding= (kernel_size - 1) // 2,188 w_init_gain= 'relu'189 )190 self.relu = torch.nn.ReLU()191 self.dropout = torch.nn.Dropout(p= dropout_rate)192 self.conv_1 = Conv1d(193 in_channels= channels,194 out_channels= channels,195 kernel_size= kernel_size,196 padding= (kernel_size - 1) // 2,197 w_init_gain= 'linear'198 )199 self.norm = LayerNorm(200 num_features= channels,201 )202 203 def forward(204 self,205 x: torch.Tensor,206 masks: torch.Tensor207 ) -> torch.Tensor:208 '''209 x: [Batch, Dim, Time]210 '''211 residuals = x212 213 x = self.conv_0(x * masks)214 x = self.relu(x)215 x = self.dropout(x)216 x = self.conv_1(x * masks)217 x = self.dropout(x)218 x = self.norm(x + residuals)219 220 return x * masks221 222# https://pytorch.org/tutorials/beginner/transformer_tutorial.html223# https://github.com/soobinseo/Transformer-TTS/blob/master/network.py224class Duration_Positional_Encoding(torch.nn.Embedding):225 def __init__(226 self, 227 num_embeddings: int,228 embedding_dim: int,229 ): 230 positional_embedding = torch.zeros(num_embeddings, embedding_dim)231 position = torch.arange(0, num_embeddings, dtype=torch.float).unsqueeze(1)232 div_term = torch.exp(torch.arange(0, embedding_dim, 2).float() * (-math.log(10000.0) / embedding_dim))233 positional_embedding[:, 0::2] = torch.sin(position * div_term)234 positional_embedding[:, 1::2] = torch.cos(position * div_term)235 super().__init__(236 num_embeddings= num_embeddings,237 embedding_dim= embedding_dim,238 _weight= positional_embedding239 )240 self.weight.requires_grad = False241 242 self.alpha = torch.nn.Parameter(243 data= torch.ones(1) * 0.01,244 requires_grad= True245 )246 247 def forward(self, durations):248 '''249 durations: [Batch, Length]250 '''251 return self.alpha * super().forward(durations) # [Batch, Dim, Length]252 253 @torch.jit.script254 def get_pe(x: torch.Tensor, pe: torch.Tensor):255 pe = pe.repeat(1, 1, math.ceil(x.size(2) / pe.size(2))) 256 return pe[:, :, :x.size(2)]257 258def Mask_Generate(lengths: torch.Tensor, max_length: Union[torch.Tensor, int, None]= None):259 '''260 lengths: [Batch]261 max_lengths: an int value. If None, max_lengths == max(lengths)262 '''263 max_length = max_length or torch.max(lengths)264 sequence = torch.arange(max_length)[None, :].to(lengths.device)265 return sequence >= lengths[:, None] # [Batch, Time]