codejin/diffsingerkr
6
1import torch2 3class Conv1d(torch.nn.Conv1d):4 def __init__(self, w_init_gain= 'linear', *args, **kwargs):5 self.w_init_gain = w_init_gain6 super().__init__(*args, **kwargs)7 8 def reset_parameters(self):9 if self.w_init_gain in ['zero']:10 torch.nn.init.zeros_(self.weight)11 elif self.w_init_gain is None:12 pass13 elif self.w_init_gain in ['relu', 'leaky_relu']:14 torch.nn.init.kaiming_uniform_(self.weight, nonlinearity= self.w_init_gain)15 elif self.w_init_gain == 'glu':16 assert self.out_channels % 2 == 0, 'The out_channels of GLU requires even number.'17 torch.nn.init.kaiming_uniform_(self.weight[:self.out_channels // 2], nonlinearity= 'linear')18 torch.nn.init.xavier_uniform_(self.weight[self.out_channels // 2:], gain= torch.nn.init.calculate_gain('sigmoid'))19 elif self.w_init_gain == 'gate':20 assert self.out_channels % 2 == 0, 'The out_channels of GLU requires even number.'21 torch.nn.init.xavier_uniform_(self.weight[:self.out_channels // 2], gain= torch.nn.init.calculate_gain('tanh'))22 torch.nn.init.xavier_uniform_(self.weight[self.out_channels // 2:], gain= torch.nn.init.calculate_gain('sigmoid'))23 else:24 torch.nn.init.xavier_uniform_(self.weight, gain= torch.nn.init.calculate_gain(self.w_init_gain))25 if not self.bias is None:26 torch.nn.init.zeros_(self.bias)27 28class ConvTranspose1d(torch.nn.ConvTranspose1d):29 def __init__(self, w_init_gain= 'linear', *args, **kwargs):30 self.w_init_gain = w_init_gain31 super().__init__(*args, **kwargs)32 33 def reset_parameters(self):34 if self.w_init_gain in ['zero']:35 torch.nn.init.zeros_(self.weight)36 elif self.w_init_gain in ['relu', 'leaky_relu']:37 torch.nn.init.kaiming_uniform_(self.weight, nonlinearity= self.w_init_gain)38 elif self.w_init_gain == 'glu':39 assert self.out_channels % 2 == 0, 'The out_channels of GLU requires even number.'40 torch.nn.init.kaiming_uniform_(self.weight[:self.out_channels // 2], nonlinearity= 'linear')41 torch.nn.init.xavier_uniform_(self.weight[self.out_channels // 2:], gain= torch.nn.init.calculate_gain('sigmoid'))42 elif self.w_init_gain == 'gate':43 assert self.out_channels % 2 == 0, 'The out_channels of GLU requires even number.'44 torch.nn.init.xavier_uniform_(self.weight[:self.out_channels // 2], gain= torch.nn.init.calculate_gain('tanh'))45 torch.nn.init.xavier_uniform_(self.weight[self.out_channels // 2:], gain= torch.nn.init.calculate_gain('sigmoid'))46 else:47 torch.nn.init.xavier_uniform_(self.weight, gain= torch.nn.init.calculate_gain(self.w_init_gain))48 if not self.bias is None:49 torch.nn.init.zeros_(self.bias)50 51class Conv2d(torch.nn.Conv2d):52 def __init__(self, w_init_gain= 'linear', *args, **kwargs):53 self.w_init_gain = w_init_gain54 super().__init__(*args, **kwargs)55 56 def reset_parameters(self):57 if self.w_init_gain in ['zero']:58 torch.nn.init.zeros_(self.weight)59 elif self.w_init_gain in ['relu', 'leaky_relu']:60 torch.nn.init.kaiming_uniform_(self.weight, nonlinearity= self.w_init_gain)61 elif self.w_init_gain == 'glu':62 assert self.out_channels % 2 == 0, 'The out_channels of GLU requires even number.'63 torch.nn.init.kaiming_uniform_(self.weight[:self.out_channels // 2], nonlinearity= 'linear')64 torch.nn.init.xavier_uniform_(self.weight[self.out_channels // 2:], gain= torch.nn.init.calculate_gain('sigmoid'))65 elif self.w_init_gain == 'gate':66 assert self.out_channels % 2 == 0, 'The out_channels of GLU requires even number.'67 torch.nn.init.xavier_uniform_(self.weight[:self.out_channels // 2], gain= torch.nn.init.calculate_gain('tanh'))68 torch.nn.init.xavier_uniform_(self.weight[self.out_channels // 2:], gain= torch.nn.init.calculate_gain('sigmoid'))69 else:70 torch.nn.init.xavier_uniform_(self.weight, gain= torch.nn.init.calculate_gain(self.w_init_gain))71 if not self.bias is None:72 torch.nn.init.zeros_(self.bias)73 74class ConvTranspose2d(torch.nn.ConvTranspose2d):75 def __init__(self, w_init_gain= 'linear', *args, **kwargs):76 self.w_init_gain = w_init_gain77 super().__init__(*args, **kwargs)78 79 def reset_parameters(self):80 if self.w_init_gain in ['zero']:81 torch.nn.init.zeros_(self.weight)82 elif self.w_init_gain in ['relu', 'leaky_relu']:83 torch.nn.init.kaiming_uniform_(self.weight, nonlinearity= self.w_init_gain)84 elif self.w_init_gain == 'glu':85 assert self.out_channels % 2 == 0, 'The out_channels of GLU requires even number.'86 torch.nn.init.kaiming_uniform_(self.weight[:self.out_channels // 2], nonlinearity= 'linear')87 torch.nn.init.xavier_uniform_(self.weight[self.out_channels // 2:], gain= torch.nn.init.calculate_gain('sigmoid'))88 elif self.w_init_gain == 'gate':89 assert self.out_channels % 2 == 0, 'The out_channels of GLU requires even number.'90 torch.nn.init.xavier_uniform_(self.weight[:self.out_channels // 2], gain= torch.nn.init.calculate_gain('tanh'))91 torch.nn.init.xavier_uniform_(self.weight[self.out_channels // 2:], gain= torch.nn.init.calculate_gain('sigmoid'))92 else:93 torch.nn.init.xavier_uniform_(self.weight, gain= torch.nn.init.calculate_gain(self.w_init_gain))94 if not self.bias is None:95 torch.nn.init.zeros_(self.bias)96 97class Linear(torch.nn.Linear):98 def __init__(self, w_init_gain= 'linear', *args, **kwargs):99 self.w_init_gain = w_init_gain100 super().__init__(*args, **kwargs)101 102 def reset_parameters(self):103 if self.w_init_gain in ['zero']:104 torch.nn.init.zeros_(self.weight)105 elif self.w_init_gain in ['relu', 'leaky_relu']:106 torch.nn.init.kaiming_uniform_(self.weight, nonlinearity= self.w_init_gain)107 elif self.w_init_gain == 'glu':108 assert self.out_channels % 2 == 0, 'The out_channels of GLU requires even number.'109 torch.nn.init.kaiming_uniform_(self.weight[:self.out_channels // 2], nonlinearity= 'linear')110 torch.nn.init.xavier_uniform_(self.weight[self.out_channels // 2:], gain= torch.nn.init.calculate_gain('sigmoid'))111 else:112 torch.nn.init.xavier_uniform_(self.weight, gain= torch.nn.init.calculate_gain(self.w_init_gain))113 if not self.bias is None:114 torch.nn.init.zeros_(self.bias)115 116class Lambda(torch.nn.Module):117 def __init__(self, lambd):118 super().__init__()119 self.lambd = lambd120 121 def forward(self, x):122 return self.lambd(x)123 124class Residual(torch.nn.Module):125 def __init__(self, module):126 super().__init__()127 self.module = module128 129 def forward(self, *args, **kwargs):130 return self.module(*args, **kwargs)131 132class LayerNorm(torch.nn.Module):133 def __init__(self, num_features: int, eps: float= 1e-5):134 super().__init__()135 136 self.eps = eps137 self.gamma = torch.nn.Parameter(torch.ones(num_features))138 self.beta = torch.nn.Parameter(torch.zeros(num_features))139 140 141 def forward(self, inputs: torch.Tensor):142 means = inputs.mean(dim= 1, keepdim= True)143 variances = (inputs - means).pow(2.0).mean(dim= 1, keepdim= True)144 145 x = (inputs - means) * (variances + self.eps).rsqrt()146 147 shape = [1, -1] + [1] * (x.ndim - 2)148 149 return x * self.gamma.view(*shape) + self.beta.view(*shape)150 151class LightweightConv1d(torch.nn.Module):152 '''153 Args:154 input_size: # of channels of the input and output155 kernel_size: convolution channels156 padding: padding157 num_heads: number of heads used. The weight is of shape158 `(num_heads, 1, kernel_size)`159 weight_softmax: normalize the weight with softmax before the convolution160 161 Shape:162 Input: BxCxT, i.e. (batch_size, input_size, timesteps)163 Output: BxCxT, i.e. (batch_size, input_size, timesteps)164 165 Attributes:166 weight: the learnable weights of the module of shape167 `(num_heads, 1, kernel_size)`168 bias: the learnable bias of the module of shape `(input_size)`169 '''170 171 def __init__(172 self,173 input_size,174 kernel_size=1,175 padding=0,176 num_heads=1,177 weight_softmax=False,178 bias=False,179 weight_dropout=0.0,180 w_init_gain= 'linear'181 ):182 super().__init__()183 self.input_size = input_size184 self.kernel_size = kernel_size185 self.num_heads = num_heads186 self.padding = padding187 self.weight_softmax = weight_softmax188 self.weight = torch.nn.Parameter(torch.Tensor(num_heads, 1, kernel_size))189 self.w_init_gain = w_init_gain190 191 if bias:192 self.bias = torch.nn.Parameter(torch.Tensor(input_size))193 else:194 self.bias = None195 self.weight_dropout_module = FairseqDropout(196 weight_dropout, module_name=self.__class__.__name__197 )198 self.reset_parameters()199 200 def reset_parameters(self):201 if self.w_init_gain in ['relu', 'leaky_relu']:202 torch.nn.init.kaiming_uniform_(self.weight, nonlinearity= self.w_init_gain)203 elif self.w_init_gain == 'glu':204 assert self.out_channels % 2 == 0, 'The out_channels of GLU requires even number.'205 torch.nn.init.kaiming_uniform_(self.weight[:self.out_channels // 2], nonlinearity= 'linear')206 torch.nn.init.xavier_uniform_(self.weight[self.out_channels // 2:], gain= torch.nn.init.calculate_gain('sigmoid'))207 else:208 torch.nn.init.xavier_uniform_(self.weight, gain= torch.nn.init.calculate_gain(self.w_init_gain))209 if not self.bias is None:210 torch.nn.init.zeros_(self.bias)211 212 def forward(self, input):213 """214 input size: B x C x T215 output size: B x C x T216 """217 B, C, T = input.size()218 H = self.num_heads219 220 weight = self.weight221 if self.weight_softmax:222 weight = weight.softmax(dim=-1)223 224 weight = self.weight_dropout_module(weight)225 # Merge every C/H entries into the batch dimension (C = self.input_size)226 # B x C x T -> (B * C/H) x H x T227 # One can also expand the weight to C x 1 x K by a factor of C/H228 # and do not reshape the input instead, which is slow though229 input = input.view(-1, H, T)230 output = torch.nn.functional.conv1d(input, weight, padding=self.padding, groups=self.num_heads)231 output = output.view(B, C, T)232 if self.bias is not None:233 output = output + self.bias.view(1, -1, 1)234 235 return output236 237class FairseqDropout(torch.nn.Module):238 def __init__(self, p, module_name=None):239 super().__init__()240 self.p = p241 self.module_name = module_name242 self.apply_during_inference = False243 244 def forward(self, x, inplace: bool = False):245 if self.training or self.apply_during_inference:246 return torch.nn.functional.dropout(x, p=self.p, training=True, inplace=inplace)247 else:248 return x249 250class LinearAttention(torch.nn.Module):251 def __init__(252 self,253 channels: int,254 calc_channels: int,255 num_heads: int,256 dropout_rate: float= 0.1,257 use_scale: bool= True,258 use_residual: bool= True,259 use_norm: bool= True260 ):261 super().__init__()262 assert calc_channels % num_heads == 0263 self.calc_channels = calc_channels264 self.num_heads = num_heads265 self.use_scale = use_scale266 self.use_residual = use_residual267 self.use_norm = use_norm268 269 self.prenet = Conv1d(270 in_channels= channels,271 out_channels= calc_channels * 3,272 kernel_size= 1,273 bias=False,274 w_init_gain= 'linear'275 )276 self.projection = Conv1d(277 in_channels= calc_channels,278 out_channels= channels,279 kernel_size= 1,280 w_init_gain= 'linear'281 )282 self.dropout = torch.nn.Dropout(p= dropout_rate)283 284 if use_scale:285 self.scale = torch.nn.Parameter(torch.zeros(1))286 287 if use_norm:288 self.norm = LayerNorm(num_features= channels)289 290 def forward(self, x: torch.Tensor, *args, **kwargs):291 '''292 x: [Batch, Enc_d, Enc_t]293 '''294 residuals = x295 296 x = self.prenet(x) # [Batch, Calc_d * 3, Enc_t] 297 x = x.view(x.size(0), self.num_heads, x.size(1) // self.num_heads, x.size(2)) # [Batch, Head, Calc_d // Head * 3, Enc_t]298 queries, keys, values = x.chunk(chunks= 3, dim= 2) # [Batch, Head, Calc_d // Head, Enc_t] * 3299 keys = (keys + 1e-5).softmax(dim= 3)300 301 contexts = keys @ values.permute(0, 1, 3, 2) # [Batch, Head, Calc_d // Head, Calc_d // Head]302 contexts = contexts.permute(0, 1, 3, 2) @ queries # [Batch, Head, Calc_d // Head, Enc_t]303 contexts = contexts.view(contexts.size(0), contexts.size(1) * contexts.size(2), contexts.size(3)) # [Batch, Calc_d, Enc_t]304 contexts = self.projection(contexts) # [Batch, Enc_d, Enc_t]305 306 if self.use_scale:307 contexts = self.scale * contexts308 309 contexts = self.dropout(contexts)310 311 if self.use_residual:312 contexts = contexts + residuals313 314 if self.use_norm:315 contexts = self.norm(contexts)316 317 return contexts318 