Team Ai
Apppublic

codejin/diffsingerkr

sourceHugging Facemitupdated 4y agoView on Hugging Face
6likes
Layer.py318 linesDownload Raw Back to Modules
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