Team Ai
Apppublic

codejin/diffsingerkr

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