CoolFace
Apppublic

kevinwang676/FreeVC-OpenAI-TTS

sourceHugging Facemitupdated 2y agoView on Hugging Face
0likes
models.py352 linesDownload Raw Back to root
1import copy2import math3import torch4from torch import nn5from torch.nn import functional as F6 7import commons8import modules9 10from torch.nn import Conv1d, ConvTranspose1d, AvgPool1d, Conv2d11from torch.nn.utils import weight_norm, remove_weight_norm, spectral_norm12from commons import init_weights, get_padding13 14 15class ResidualCouplingBlock(nn.Module):16  def __init__(self,17      channels,18      hidden_channels,19      kernel_size,20      dilation_rate,21      n_layers,22      n_flows=4,23      gin_channels=0):24    super().__init__()25    self.channels = channels26    self.hidden_channels = hidden_channels27    self.kernel_size = kernel_size28    self.dilation_rate = dilation_rate29    self.n_layers = n_layers30    self.n_flows = n_flows31    self.gin_channels = gin_channels32 33    self.flows = nn.ModuleList()34    for i in range(n_flows):35      self.flows.append(modules.ResidualCouplingLayer(channels, hidden_channels, kernel_size, dilation_rate, n_layers, gin_channels=gin_channels, mean_only=True))36      self.flows.append(modules.Flip())37 38  def forward(self, x, x_mask, g=None, reverse=False):39    if not reverse:40      for flow in self.flows:41        x, _ = flow(x, x_mask, g=g, reverse=reverse)42    else:43      for flow in reversed(self.flows):44        x = flow(x, x_mask, g=g, reverse=reverse)45    return x46 47 48class Encoder(nn.Module):49  def __init__(self,50      in_channels,51      out_channels,52      hidden_channels,53      kernel_size,54      dilation_rate,55      n_layers,56      gin_channels=0):57    super().__init__()58    self.in_channels = in_channels59    self.out_channels = out_channels60    self.hidden_channels = hidden_channels61    self.kernel_size = kernel_size62    self.dilation_rate = dilation_rate63    self.n_layers = n_layers64    self.gin_channels = gin_channels65 66    self.pre = nn.Conv1d(in_channels, hidden_channels, 1)67    self.enc = modules.WN(hidden_channels, kernel_size, dilation_rate, n_layers, gin_channels=gin_channels)68    self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)69 70  def forward(self, x, x_lengths, g=None):71    x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(x.dtype)72    x = self.pre(x) * x_mask73    x = self.enc(x, x_mask, g=g)74    stats = self.proj(x) * x_mask75    m, logs = torch.split(stats, self.out_channels, dim=1)76    z = (m + torch.randn_like(m) * torch.exp(logs)) * x_mask77    return z, m, logs, x_mask78 79 80class Generator(torch.nn.Module):81    def __init__(self, initial_channel, resblock, resblock_kernel_sizes, resblock_dilation_sizes, upsample_rates, upsample_initial_channel, upsample_kernel_sizes, gin_channels=0):82        super(Generator, self).__init__()83        self.num_kernels = len(resblock_kernel_sizes)84        self.num_upsamples = len(upsample_rates)85        self.conv_pre = Conv1d(initial_channel, upsample_initial_channel, 7, 1, padding=3)86        resblock = modules.ResBlock1 if resblock == '1' else modules.ResBlock287 88        self.ups = nn.ModuleList()89        for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)):90            self.ups.append(weight_norm(91                ConvTranspose1d(upsample_initial_channel//(2**i), upsample_initial_channel//(2**(i+1)),92                                k, u, padding=(k-u)//2)))93 94        self.resblocks = nn.ModuleList()95        for i in range(len(self.ups)):96            ch = upsample_initial_channel//(2**(i+1))97            for j, (k, d) in enumerate(zip(resblock_kernel_sizes, resblock_dilation_sizes)):98                self.resblocks.append(resblock(ch, k, d))99 100        self.conv_post = Conv1d(ch, 1, 7, 1, padding=3, bias=False)101        self.ups.apply(init_weights)102 103        if gin_channels != 0:104            self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1)105 106    def forward(self, x, g=None):107        x = self.conv_pre(x)108        if g is not None:109          x = x + self.cond(g)110 111        for i in range(self.num_upsamples):112            x = F.leaky_relu(x, modules.LRELU_SLOPE)113            x = self.ups[i](x)114            xs = None115            for j in range(self.num_kernels):116                if xs is None:117                    xs = self.resblocks[i*self.num_kernels+j](x)118                else:119                    xs += self.resblocks[i*self.num_kernels+j](x)120            x = xs / self.num_kernels121        x = F.leaky_relu(x)122        x = self.conv_post(x)123        x = torch.tanh(x)124 125        return x126 127    def remove_weight_norm(self):128        print('Removing weight norm...')129        for l in self.ups:130            remove_weight_norm(l)131        for l in self.resblocks:132            l.remove_weight_norm()133 134 135class DiscriminatorP(torch.nn.Module):136    def __init__(self, period, kernel_size=5, stride=3, use_spectral_norm=False):137        super(DiscriminatorP, self).__init__()138        self.period = period139        self.use_spectral_norm = use_spectral_norm140        norm_f = weight_norm if use_spectral_norm == False else spectral_norm141        self.convs = nn.ModuleList([142            norm_f(Conv2d(1, 32, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),143            norm_f(Conv2d(32, 128, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),144            norm_f(Conv2d(128, 512, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),145            norm_f(Conv2d(512, 1024, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),146            norm_f(Conv2d(1024, 1024, (kernel_size, 1), 1, padding=(get_padding(kernel_size, 1), 0))),147        ])148        self.conv_post = norm_f(Conv2d(1024, 1, (3, 1), 1, padding=(1, 0)))149 150    def forward(self, x):151        fmap = []152 153        # 1d to 2d154        b, c, t = x.shape155        if t % self.period != 0: # pad first156            n_pad = self.period - (t % self.period)157            x = F.pad(x, (0, n_pad), "reflect")158            t = t + n_pad159        x = x.view(b, c, t // self.period, self.period)160 161        for l in self.convs:162            x = l(x)163            x = F.leaky_relu(x, modules.LRELU_SLOPE)164            fmap.append(x)165        x = self.conv_post(x)166        fmap.append(x)167        x = torch.flatten(x, 1, -1)168 169        return x, fmap170 171 172class DiscriminatorS(torch.nn.Module):173    def __init__(self, use_spectral_norm=False):174        super(DiscriminatorS, self).__init__()175        norm_f = weight_norm if use_spectral_norm == False else spectral_norm176        self.convs = nn.ModuleList([177            norm_f(Conv1d(1, 16, 15, 1, padding=7)),178            norm_f(Conv1d(16, 64, 41, 4, groups=4, padding=20)),179            norm_f(Conv1d(64, 256, 41, 4, groups=16, padding=20)),180            norm_f(Conv1d(256, 1024, 41, 4, groups=64, padding=20)),181            norm_f(Conv1d(1024, 1024, 41, 4, groups=256, padding=20)),182            norm_f(Conv1d(1024, 1024, 5, 1, padding=2)),183        ])184        self.conv_post = norm_f(Conv1d(1024, 1, 3, 1, padding=1))185 186    def forward(self, x):187        fmap = []188 189        for l in self.convs:190            x = l(x)191            x = F.leaky_relu(x, modules.LRELU_SLOPE)192            fmap.append(x)193        x = self.conv_post(x)194        fmap.append(x)195        x = torch.flatten(x, 1, -1)196 197        return x, fmap198 199 200class MultiPeriodDiscriminator(torch.nn.Module):201    def __init__(self, use_spectral_norm=False):202        super(MultiPeriodDiscriminator, self).__init__()203        periods = [2,3,5,7,11]204 205        discs = [DiscriminatorS(use_spectral_norm=use_spectral_norm)]206        discs = discs + [DiscriminatorP(i, use_spectral_norm=use_spectral_norm) for i in periods]207        self.discriminators = nn.ModuleList(discs)208 209    def forward(self, y, y_hat):210        y_d_rs = []211        y_d_gs = []212        fmap_rs = []213        fmap_gs = []214        for i, d in enumerate(self.discriminators):215            y_d_r, fmap_r = d(y)216            y_d_g, fmap_g = d(y_hat)217            y_d_rs.append(y_d_r)218            y_d_gs.append(y_d_g)219            fmap_rs.append(fmap_r)220            fmap_gs.append(fmap_g)221 222        return y_d_rs, y_d_gs, fmap_rs, fmap_gs223        224        225class SpeakerEncoder(torch.nn.Module):226    def __init__(self, mel_n_channels=80, model_num_layers=3, model_hidden_size=256, model_embedding_size=256):227        super(SpeakerEncoder, self).__init__()228        self.lstm = nn.LSTM(mel_n_channels, model_hidden_size, model_num_layers, batch_first=True)229        self.linear = nn.Linear(model_hidden_size, model_embedding_size)230        self.relu = nn.ReLU()231 232    def forward(self, mels):233        self.lstm.flatten_parameters()234        _, (hidden, _) = self.lstm(mels)235        embeds_raw = self.relu(self.linear(hidden[-1]))236        return embeds_raw / torch.norm(embeds_raw, dim=1, keepdim=True)237        238    def compute_partial_slices(self, total_frames, partial_frames, partial_hop):239        mel_slices = []240        for i in range(0, total_frames-partial_frames, partial_hop):241            mel_range = torch.arange(i, i+partial_frames)242            mel_slices.append(mel_range)243            244        return mel_slices245    246    def embed_utterance(self, mel, partial_frames=128, partial_hop=64):247        mel_len = mel.size(1)248        last_mel = mel[:,-partial_frames:]249        250        if mel_len > partial_frames:251            mel_slices = self.compute_partial_slices(mel_len, partial_frames, partial_hop)252            mels = list(mel[:,s] for s in mel_slices)253            mels.append(last_mel)254            mels = torch.stack(tuple(mels), 0).squeeze(1)255        256            with torch.no_grad():257                partial_embeds = self(mels)258            embed = torch.mean(partial_embeds, axis=0).unsqueeze(0)259            #embed = embed / torch.linalg.norm(embed, 2)260        else:261            with torch.no_grad():262                embed = self(last_mel)263        264        return embed265 266 267class SynthesizerTrn(nn.Module):268  """269  Synthesizer for Training270  """271 272  def __init__(self, 273    spec_channels,274    segment_size,275    inter_channels,276    hidden_channels,277    filter_channels,278    n_heads,279    n_layers,280    kernel_size,281    p_dropout,282    resblock, 283    resblock_kernel_sizes, 284    resblock_dilation_sizes, 285    upsample_rates, 286    upsample_initial_channel, 287    upsample_kernel_sizes,288    gin_channels,289    ssl_dim,290    use_spk,291    **kwargs):292 293    super().__init__()294    self.spec_channels = spec_channels295    self.inter_channels = inter_channels296    self.hidden_channels = hidden_channels297    self.filter_channels = filter_channels298    self.n_heads = n_heads299    self.n_layers = n_layers300    self.kernel_size = kernel_size301    self.p_dropout = p_dropout302    self.resblock = resblock303    self.resblock_kernel_sizes = resblock_kernel_sizes304    self.resblock_dilation_sizes = resblock_dilation_sizes305    self.upsample_rates = upsample_rates306    self.upsample_initial_channel = upsample_initial_channel307    self.upsample_kernel_sizes = upsample_kernel_sizes308    self.segment_size = segment_size309    self.gin_channels = gin_channels310    self.ssl_dim = ssl_dim311    self.use_spk = use_spk312 313    self.enc_p = Encoder(ssl_dim, inter_channels, hidden_channels, 5, 1, 16)314    self.dec = Generator(inter_channels, resblock, resblock_kernel_sizes, resblock_dilation_sizes, upsample_rates, upsample_initial_channel, upsample_kernel_sizes, gin_channels=gin_channels)315    self.enc_q = Encoder(spec_channels, inter_channels, hidden_channels, 5, 1, 16, gin_channels=gin_channels) 316    self.flow = ResidualCouplingBlock(inter_channels, hidden_channels, 5, 1, 4, gin_channels=gin_channels)317    318    if not self.use_spk:319      self.enc_spk = SpeakerEncoder(model_hidden_size=gin_channels, model_embedding_size=gin_channels)320 321  def forward(self, c, spec, g=None, mel=None, c_lengths=None, spec_lengths=None):322    if c_lengths == None:323      c_lengths = (torch.ones(c.size(0)) * c.size(-1)).to(c.device)324    if spec_lengths == None:325      spec_lengths = (torch.ones(spec.size(0)) * spec.size(-1)).to(spec.device)326      327    if not self.use_spk:328      g = self.enc_spk(mel.transpose(1,2))329    g = g.unsqueeze(-1)330      331    _, m_p, logs_p, _ = self.enc_p(c, c_lengths)332    z, m_q, logs_q, spec_mask = self.enc_q(spec, spec_lengths, g=g) 333    z_p = self.flow(z, spec_mask, g=g)334 335    z_slice, ids_slice = commons.rand_slice_segments(z, spec_lengths, self.segment_size)336    o = self.dec(z_slice, g=g)337    338    return o, ids_slice, spec_mask, (z, z_p, m_p, logs_p, m_q, logs_q)339 340  def infer(self, c, g=None, mel=None, c_lengths=None):341    if c_lengths == None:342      c_lengths = (torch.ones(c.size(0)) * c.size(-1)).to(c.device)343    if not self.use_spk:344      g = self.enc_spk.embed_utterance(mel.transpose(1,2))345    g = g.unsqueeze(-1)346 347    z_p, m_p, logs_p, c_mask = self.enc_p(c, c_lengths)348    z = self.flow(z_p, c_mask, g=g, reverse=True)349    o = self.dec(z * c_mask, g=g)350    351    return o352