CoolFace
Apppublic

Kafke/Code-Realize-TTS

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
models.py543 linesDownload Raw Back to root
1import math2import torch3from torch import nn4from torch.nn import functional as F5 6import commons7import modules8import attentions9 10from torch.nn import Conv1d, ConvTranspose1d, Conv2d11from torch.nn.utils import weight_norm, remove_weight_norm, spectral_norm12from commons import init_weights, get_padding13 14 15class StochasticDurationPredictor(nn.Module):16  def __init__(self, in_channels, filter_channels, kernel_size, p_dropout, n_flows=4, gin_channels=0):17    super().__init__()18    filter_channels = in_channels # it needs to be removed from future version.19    self.in_channels = in_channels20    self.filter_channels = filter_channels21    self.kernel_size = kernel_size22    self.p_dropout = p_dropout23    self.n_flows = n_flows24    self.gin_channels = gin_channels25 26    self.log_flow = modules.Log()27    self.flows = nn.ModuleList()28    self.flows.append(modules.ElementwiseAffine(2))29    for i in range(n_flows):30      self.flows.append(modules.ConvFlow(2, filter_channels, kernel_size, n_layers=3))31      self.flows.append(modules.Flip())32 33    self.post_pre = nn.Conv1d(1, filter_channels, 1)34    self.post_proj = nn.Conv1d(filter_channels, filter_channels, 1)35    self.post_convs = modules.DDSConv(filter_channels, kernel_size, n_layers=3, p_dropout=p_dropout)36    self.post_flows = nn.ModuleList()37    self.post_flows.append(modules.ElementwiseAffine(2))38    for i in range(4):39      self.post_flows.append(modules.ConvFlow(2, filter_channels, kernel_size, n_layers=3))40      self.post_flows.append(modules.Flip())41 42    self.pre = nn.Conv1d(in_channels, filter_channels, 1)43    self.proj = nn.Conv1d(filter_channels, filter_channels, 1)44    self.convs = modules.DDSConv(filter_channels, kernel_size, n_layers=3, p_dropout=p_dropout)45    if gin_channels != 0:46      self.cond = nn.Conv1d(gin_channels, filter_channels, 1)47 48  def forward(self, x, x_mask, w=None, g=None, reverse=False, noise_scale=1.0):49    x = torch.detach(x)50    x = self.pre(x)51    if g is not None:52      g = torch.detach(g)53      x = x + self.cond(g)54    x = self.convs(x, x_mask)55    x = self.proj(x) * x_mask56 57    if not reverse:58      flows = self.flows59      assert w is not None60 61      logdet_tot_q = 0 62      h_w = self.post_pre(w)63      h_w = self.post_convs(h_w, x_mask)64      h_w = self.post_proj(h_w) * x_mask65      e_q = torch.randn(w.size(0), 2, w.size(2)).to(device=x.device, dtype=x.dtype) * x_mask66      z_q = e_q67      for flow in self.post_flows:68        z_q, logdet_q = flow(z_q, x_mask, g=(x + h_w))69        logdet_tot_q += logdet_q70      z_u, z1 = torch.split(z_q, [1, 1], 1) 71      u = torch.sigmoid(z_u) * x_mask72      z0 = (w - u) * x_mask73      logdet_tot_q += torch.sum((F.logsigmoid(z_u) + F.logsigmoid(-z_u)) * x_mask, [1,2])74      logq = torch.sum(-0.5 * (math.log(2*math.pi) + (e_q**2)) * x_mask, [1,2]) - logdet_tot_q75 76      logdet_tot = 077      z0, logdet = self.log_flow(z0, x_mask)78      logdet_tot += logdet79      z = torch.cat([z0, z1], 1)80      for flow in flows:81        z, logdet = flow(z, x_mask, g=x, reverse=reverse)82        logdet_tot = logdet_tot + logdet83      nll = torch.sum(0.5 * (math.log(2*math.pi) + (z**2)) * x_mask, [1,2]) - logdet_tot84      return nll + logq # [b]85    else:86      flows = list(reversed(self.flows))87      flows = flows[:-2] + [flows[-1]] # remove a useless vflow88      z = torch.randn(x.size(0), 2, x.size(2)).to(device=x.device, dtype=x.dtype) * noise_scale89      for flow in flows:90        z = flow(z, x_mask, g=x, reverse=reverse)91      z0, z1 = torch.split(z, [1, 1], 1)92      logw = z093      return logw94 95 96class DurationPredictor(nn.Module):97  def __init__(self, in_channels, filter_channels, kernel_size, p_dropout, gin_channels=0):98    super().__init__()99 100    self.in_channels = in_channels101    self.filter_channels = filter_channels102    self.kernel_size = kernel_size103    self.p_dropout = p_dropout104    self.gin_channels = gin_channels105 106    self.drop = nn.Dropout(p_dropout)107    self.conv_1 = nn.Conv1d(in_channels, filter_channels, kernel_size, padding=kernel_size//2)108    self.norm_1 = modules.LayerNorm(filter_channels)109    self.conv_2 = nn.Conv1d(filter_channels, filter_channels, kernel_size, padding=kernel_size//2)110    self.norm_2 = modules.LayerNorm(filter_channels)111    self.proj = nn.Conv1d(filter_channels, 1, 1)112 113    if gin_channels != 0:114      self.cond = nn.Conv1d(gin_channels, in_channels, 1)115 116  def forward(self, x, x_mask, g=None):117    x = torch.detach(x)118    if g is not None:119      g = torch.detach(g)120      x = x + self.cond(g)121    x = self.conv_1(x * x_mask)122    x = torch.relu(x)123    x = self.norm_1(x)124    x = self.drop(x)125    x = self.conv_2(x * x_mask)126    x = torch.relu(x)127    x = self.norm_2(x)128    x = self.drop(x)129    x = self.proj(x * x_mask)130    return x * x_mask131 132 133class TextEncoder(nn.Module):134  def __init__(self,135      n_vocab,136      out_channels,137      hidden_channels,138      filter_channels,139      n_heads,140      n_layers,141      kernel_size,142      p_dropout,143      emotion_embedding):144    super().__init__()145    self.n_vocab = n_vocab146    self.out_channels = out_channels147    self.hidden_channels = hidden_channels148    self.filter_channels = filter_channels149    self.n_heads = n_heads150    self.n_layers = n_layers151    self.kernel_size = kernel_size152    self.p_dropout = p_dropout153    self.emotion_embedding = emotion_embedding154    155    if self.n_vocab!=0:156      self.emb = nn.Embedding(n_vocab, hidden_channels)157      if emotion_embedding:158        self.emotion_emb = nn.Linear(1024, hidden_channels)159      nn.init.normal_(self.emb.weight, 0.0, hidden_channels**-0.5)160 161    self.encoder = attentions.Encoder(162      hidden_channels,163      filter_channels,164      n_heads,165      n_layers,166      kernel_size,167      p_dropout)168    self.proj= nn.Conv1d(hidden_channels, out_channels * 2, 1)169 170  def forward(self, x, x_lengths, emotion_embedding=None):171    if self.n_vocab!=0:172      x = self.emb(x) * math.sqrt(self.hidden_channels) # [b, t, h]173    if emotion_embedding is not None:174      x = x + self.emotion_emb(emotion_embedding.unsqueeze(1))175    x = torch.transpose(x, 1, -1) # [b, h, t]176    x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(x.dtype)177 178    x = self.encoder(x * x_mask, x_mask)179    stats = self.proj(x) * x_mask180 181    m, logs = torch.split(stats, self.out_channels, dim=1)182    return x, m, logs, x_mask183 184 185class ResidualCouplingBlock(nn.Module):186  def __init__(self,187      channels,188      hidden_channels,189      kernel_size,190      dilation_rate,191      n_layers,192      n_flows=4,193      gin_channels=0):194    super().__init__()195    self.channels = channels196    self.hidden_channels = hidden_channels197    self.kernel_size = kernel_size198    self.dilation_rate = dilation_rate199    self.n_layers = n_layers200    self.n_flows = n_flows201    self.gin_channels = gin_channels202 203    self.flows = nn.ModuleList()204    for i in range(n_flows):205      self.flows.append(modules.ResidualCouplingLayer(channels, hidden_channels, kernel_size, dilation_rate, n_layers, gin_channels=gin_channels, mean_only=True))206      self.flows.append(modules.Flip())207 208  def forward(self, x, x_mask, g=None, reverse=False):209    if not reverse:210      for flow in self.flows:211        x, _ = flow(x, x_mask, g=g, reverse=reverse)212    else:213      for flow in reversed(self.flows):214        x = flow(x, x_mask, g=g, reverse=reverse)215    return x216 217 218class PosteriorEncoder(nn.Module):219  def __init__(self,220      in_channels,221      out_channels,222      hidden_channels,223      kernel_size,224      dilation_rate,225      n_layers,226      gin_channels=0):227    super().__init__()228    self.in_channels = in_channels229    self.out_channels = out_channels230    self.hidden_channels = hidden_channels231    self.kernel_size = kernel_size232    self.dilation_rate = dilation_rate233    self.n_layers = n_layers234    self.gin_channels = gin_channels235 236    self.pre = nn.Conv1d(in_channels, hidden_channels, 1)237    self.enc = modules.WN(hidden_channels, kernel_size, dilation_rate, n_layers, gin_channels=gin_channels)238    self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)239 240  def forward(self, x, x_lengths, g=None):241    x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(x.dtype)242    x = self.pre(x) * x_mask243    x = self.enc(x, x_mask, g=g)244    stats = self.proj(x) * x_mask245    m, logs = torch.split(stats, self.out_channels, dim=1)246    z = (m + torch.randn_like(m) * torch.exp(logs)) * x_mask247    return z, m, logs, x_mask248 249 250class Generator(torch.nn.Module):251    def __init__(self, initial_channel, resblock, resblock_kernel_sizes, resblock_dilation_sizes, upsample_rates, upsample_initial_channel, upsample_kernel_sizes, gin_channels=0):252        super(Generator, self).__init__()253        self.num_kernels = len(resblock_kernel_sizes)254        self.num_upsamples = len(upsample_rates)255        self.conv_pre = Conv1d(initial_channel, upsample_initial_channel, 7, 1, padding=3)256        resblock = modules.ResBlock1 if resblock == '1' else modules.ResBlock2257 258        self.ups = nn.ModuleList()259        for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)):260            self.ups.append(weight_norm(261                ConvTranspose1d(upsample_initial_channel//(2**i), upsample_initial_channel//(2**(i+1)),262                                k, u, padding=(k-u)//2)))263 264        self.resblocks = nn.ModuleList()265        for i in range(len(self.ups)):266            ch = upsample_initial_channel//(2**(i+1))267            for j, (k, d) in enumerate(zip(resblock_kernel_sizes, resblock_dilation_sizes)):268                self.resblocks.append(resblock(ch, k, d))269 270        self.conv_post = Conv1d(ch, 1, 7, 1, padding=3, bias=False)271        self.ups.apply(init_weights)272 273        if gin_channels != 0:274            self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1)275 276    def forward(self, x, g=None):277        x = self.conv_pre(x)278        if g is not None:279          x = x + self.cond(g)280 281        for i in range(self.num_upsamples):282            x = F.leaky_relu(x, modules.LRELU_SLOPE)283            x = self.ups[i](x)284            xs = None285            for j in range(self.num_kernels):286                if xs is None:287                    xs = self.resblocks[i*self.num_kernels+j](x)288                else:289                    xs += self.resblocks[i*self.num_kernels+j](x)290            x = xs / self.num_kernels291        x = F.leaky_relu(x)292        x = self.conv_post(x)293        x = torch.tanh(x)294 295        return x296 297    def remove_weight_norm(self):298        print('Removing weight norm...')299        for l in self.ups:300            remove_weight_norm(l)301        for l in self.resblocks:302            l.remove_weight_norm()303 304 305class DiscriminatorP(torch.nn.Module):306    def __init__(self, period, kernel_size=5, stride=3, use_spectral_norm=False):307        super(DiscriminatorP, self).__init__()308        self.period = period309        self.use_spectral_norm = use_spectral_norm310        norm_f = weight_norm if use_spectral_norm == False else spectral_norm311        self.convs = nn.ModuleList([312            norm_f(Conv2d(1, 32, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),313            norm_f(Conv2d(32, 128, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),314            norm_f(Conv2d(128, 512, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),315            norm_f(Conv2d(512, 1024, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),316            norm_f(Conv2d(1024, 1024, (kernel_size, 1), 1, padding=(get_padding(kernel_size, 1), 0))),317        ])318        self.conv_post = norm_f(Conv2d(1024, 1, (3, 1), 1, padding=(1, 0)))319 320    def forward(self, x):321        fmap = []322 323        # 1d to 2d324        b, c, t = x.shape325        if t % self.period != 0: # pad first326            n_pad = self.period - (t % self.period)327            x = F.pad(x, (0, n_pad), "reflect")328            t = t + n_pad329        x = x.view(b, c, t // self.period, self.period)330 331        for l in self.convs:332            x = l(x)333            x = F.leaky_relu(x, modules.LRELU_SLOPE)334            fmap.append(x)335        x = self.conv_post(x)336        fmap.append(x)337        x = torch.flatten(x, 1, -1)338 339        return x, fmap340 341 342class DiscriminatorS(torch.nn.Module):343    def __init__(self, use_spectral_norm=False):344        super(DiscriminatorS, self).__init__()345        norm_f = weight_norm if use_spectral_norm == False else spectral_norm346        self.convs = nn.ModuleList([347            norm_f(Conv1d(1, 16, 15, 1, padding=7)),348            norm_f(Conv1d(16, 64, 41, 4, groups=4, padding=20)),349            norm_f(Conv1d(64, 256, 41, 4, groups=16, padding=20)),350            norm_f(Conv1d(256, 1024, 41, 4, groups=64, padding=20)),351            norm_f(Conv1d(1024, 1024, 41, 4, groups=256, padding=20)),352            norm_f(Conv1d(1024, 1024, 5, 1, padding=2)),353        ])354        self.conv_post = norm_f(Conv1d(1024, 1, 3, 1, padding=1))355 356    def forward(self, x):357        fmap = []358 359        for l in self.convs:360            x = l(x)361            x = F.leaky_relu(x, modules.LRELU_SLOPE)362            fmap.append(x)363        x = self.conv_post(x)364        fmap.append(x)365        x = torch.flatten(x, 1, -1)366 367        return x, fmap368 369 370class MultiPeriodDiscriminator(torch.nn.Module):371    def __init__(self, use_spectral_norm=False):372        super(MultiPeriodDiscriminator, self).__init__()373        periods = [2,3,5,7,11]374 375        discs = [DiscriminatorS(use_spectral_norm=use_spectral_norm)]376        discs = discs + [DiscriminatorP(i, use_spectral_norm=use_spectral_norm) for i in periods]377        self.discriminators = nn.ModuleList(discs)378 379    def forward(self, y, y_hat):380        y_d_rs = []381        y_d_gs = []382        fmap_rs = []383        fmap_gs = []384        for i, d in enumerate(self.discriminators):385            y_d_r, fmap_r = d(y)386            y_d_g, fmap_g = d(y_hat)387            y_d_rs.append(y_d_r)388            y_d_gs.append(y_d_g)389            fmap_rs.append(fmap_r)390            fmap_gs.append(fmap_g)391 392        return y_d_rs, y_d_gs, fmap_rs, fmap_gs393 394 395 396class SynthesizerTrn(nn.Module):397  """398  Synthesizer for Training399  """400 401  def __init__(self, 402    n_vocab,403    spec_channels,404    segment_size,405    inter_channels,406    hidden_channels,407    filter_channels,408    n_heads,409    n_layers,410    kernel_size,411    p_dropout,412    resblock, 413    resblock_kernel_sizes, 414    resblock_dilation_sizes, 415    upsample_rates, 416    upsample_initial_channel, 417    upsample_kernel_sizes,418    n_speakers=0,419    gin_channels=0,420    use_sdp=True,421    emotion_embedding=False,422    **kwargs):423 424    super().__init__()425    self.n_vocab = n_vocab426    self.spec_channels = spec_channels427    self.inter_channels = inter_channels428    self.hidden_channels = hidden_channels429    self.filter_channels = filter_channels430    self.n_heads = n_heads431    self.n_layers = n_layers432    self.kernel_size = kernel_size433    self.p_dropout = p_dropout434    self.resblock = resblock435    self.resblock_kernel_sizes = resblock_kernel_sizes436    self.resblock_dilation_sizes = resblock_dilation_sizes437    self.upsample_rates = upsample_rates438    self.upsample_initial_channel = upsample_initial_channel439    self.upsample_kernel_sizes = upsample_kernel_sizes440    self.segment_size = segment_size441    self.n_speakers = n_speakers442    self.gin_channels = gin_channels443 444    self.use_sdp = use_sdp445 446    self.enc_p = TextEncoder(n_vocab,447        inter_channels,448        hidden_channels,449        filter_channels,450        n_heads,451        n_layers,452        kernel_size,453        p_dropout,454        emotion_embedding)455    self.dec = Generator(inter_channels, resblock, resblock_kernel_sizes, resblock_dilation_sizes, upsample_rates, upsample_initial_channel, upsample_kernel_sizes, gin_channels=gin_channels)456    self.enc_q = PosteriorEncoder(spec_channels, inter_channels, hidden_channels, 5, 1, 16, gin_channels=gin_channels)457    self.flow = ResidualCouplingBlock(inter_channels, hidden_channels, 5, 1, 4, gin_channels=gin_channels)458 459    if use_sdp:460      self.dp = StochasticDurationPredictor(hidden_channels, 192, 3, 0.5, 4, gin_channels=gin_channels)461    else:462      self.dp = DurationPredictor(hidden_channels, 256, 3, 0.5, gin_channels=gin_channels)463 464    if n_speakers > 1:465      self.emb_g = nn.Embedding(n_speakers, gin_channels)466 467  def forward(self, x, x_lengths, y, y_lengths, sid=None):468 469    x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths)470    if self.n_speakers > 0:471      g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1]472    else:473      g = None474 475    z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g)476    z_p = self.flow(z, y_mask, g=g)477 478    with torch.no_grad():479      # negative cross-entropy480      s_p_sq_r = torch.exp(-2 * logs_p) # [b, d, t]481      neg_cent1 = torch.sum(-0.5 * math.log(2 * math.pi) - logs_p, [1], keepdim=True) # [b, 1, t_s]482      neg_cent2 = torch.matmul(-0.5 * (z_p ** 2).transpose(1, 2), s_p_sq_r) # [b, t_t, d] x [b, d, t_s] = [b, t_t, t_s]483      neg_cent3 = torch.matmul(z_p.transpose(1, 2), (m_p * s_p_sq_r)) # [b, t_t, d] x [b, d, t_s] = [b, t_t, t_s]484      neg_cent4 = torch.sum(-0.5 * (m_p ** 2) * s_p_sq_r, [1], keepdim=True) # [b, 1, t_s]485      neg_cent = neg_cent1 + neg_cent2 + neg_cent3 + neg_cent4486 487      attn_mask = torch.unsqueeze(x_mask, 2) * torch.unsqueeze(y_mask, -1)488      attn = monotonic_align.maximum_path(neg_cent, attn_mask.squeeze(1)).unsqueeze(1).detach()489 490    w = attn.sum(2)491    if self.use_sdp:492      l_length = self.dp(x, x_mask, w, g=g)493      l_length = l_length / torch.sum(x_mask)494    else:495      logw_ = torch.log(w + 1e-6) * x_mask496      logw = self.dp(x, x_mask, g=g)497      l_length = torch.sum((logw - logw_)**2, [1,2]) / torch.sum(x_mask) # for averaging 498 499    # expand prior500    m_p = torch.matmul(attn.squeeze(1), m_p.transpose(1, 2)).transpose(1, 2)501    logs_p = torch.matmul(attn.squeeze(1), logs_p.transpose(1, 2)).transpose(1, 2)502 503    z_slice, ids_slice = commons.rand_slice_segments(z, y_lengths, self.segment_size)504    o = self.dec(z_slice, g=g)505    return o, l_length, attn, ids_slice, x_mask, y_mask, (z, z_p, m_p, logs_p, m_q, logs_q)506 507  def infer(self, x, x_lengths, sid=None, noise_scale=1, length_scale=1, noise_scale_w=1., max_len=None, emotion_embedding=None):508    x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths, emotion_embedding)509    if self.n_speakers > 0:510      g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1]511    else:512      g = None513 514    if self.use_sdp:515      logw = self.dp(x, x_mask, g=g, reverse=True, noise_scale=noise_scale_w)516    else:517      logw = self.dp(x, x_mask, g=g)518    w = torch.exp(logw) * x_mask * length_scale519    w_ceil = torch.ceil(w)520    y_lengths = torch.clamp_min(torch.sum(w_ceil, [1, 2]), 1).long()521    y_mask = torch.unsqueeze(commons.sequence_mask(y_lengths, None), 1).to(x_mask.dtype)522    attn_mask = torch.unsqueeze(x_mask, 2) * torch.unsqueeze(y_mask, -1)523    attn = commons.generate_path(w_ceil, attn_mask)524 525    m_p = torch.matmul(attn.squeeze(1), m_p.transpose(1, 2)).transpose(1, 2) # [b, t', t], [b, t, d] -> [b, d, t']526    logs_p = torch.matmul(attn.squeeze(1), logs_p.transpose(1, 2)).transpose(1, 2) # [b, t', t], [b, t, d] -> [b, d, t']527 528    z_p = m_p + torch.randn_like(m_p) * torch.exp(logs_p) * noise_scale529    z = self.flow(z_p, y_mask, g=g, reverse=True)530    o = self.dec((z * y_mask)[:,:,:max_len], g=g)531    return o, attn, y_mask, (z, z_p, m_p, logs_p)532 533  def voice_conversion(self, y, y_lengths, sid_src, sid_tgt):534    assert self.n_speakers > 0, "n_speakers have to be larger than 0."535    g_src = self.emb_g(sid_src).unsqueeze(-1)536    g_tgt = self.emb_g(sid_tgt).unsqueeze(-1)537    z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g_src)538    z_p = self.flow(z, y_mask, g=g_src)539    z_hat = self.flow(z_p, y_mask, g=g_tgt, reverse=True)540    o_hat = self.dec(z_hat * y_mask, g=g_tgt)541    return o_hat, y_mask, (z, z_p, z_hat)542 543