CoolFace
Apppublic

everythingfades/vits_personal_exploration

sourceHugging Faceupdated 3y agoView on Hugging Face
0likes
models.py534 linesDownload Raw Back to root
1import math2import torch3from torch import nn4from torch.nn import functional as F5 6import commons7import modules8import attentions9import monotonic_align10 11from torch.nn import Conv1d, ConvTranspose1d, Conv2d12from torch.nn.utils import weight_norm, remove_weight_norm, spectral_norm13from commons import init_weights, get_padding14 15 16class StochasticDurationPredictor(nn.Module):17  def __init__(self, in_channels, filter_channels, kernel_size, p_dropout, n_flows=4, gin_channels=0):18    super().__init__()19    filter_channels = in_channels # it needs to be removed from future version.20    self.in_channels = in_channels21    self.filter_channels = filter_channels22    self.kernel_size = kernel_size23    self.p_dropout = p_dropout24    self.n_flows = n_flows25    self.gin_channels = gin_channels26 27    self.log_flow = modules.Log()28    self.flows = nn.ModuleList()29    self.flows.append(modules.ElementwiseAffine(2))30    for i in range(n_flows):31      self.flows.append(modules.ConvFlow(2, filter_channels, kernel_size, n_layers=3))32      self.flows.append(modules.Flip())33 34    self.post_pre = nn.Conv1d(1, filter_channels, 1)35    self.post_proj = nn.Conv1d(filter_channels, filter_channels, 1)36    self.post_convs = modules.DDSConv(filter_channels, kernel_size, n_layers=3, p_dropout=p_dropout)37    self.post_flows = nn.ModuleList()38    self.post_flows.append(modules.ElementwiseAffine(2))39    for i in range(4):40      self.post_flows.append(modules.ConvFlow(2, filter_channels, kernel_size, n_layers=3))41      self.post_flows.append(modules.Flip())42 43    self.pre = nn.Conv1d(in_channels, filter_channels, 1)44    self.proj = nn.Conv1d(filter_channels, filter_channels, 1)45    self.convs = modules.DDSConv(filter_channels, kernel_size, n_layers=3, p_dropout=p_dropout)46    if gin_channels != 0:47      self.cond = nn.Conv1d(gin_channels, filter_channels, 1)48 49  def forward(self, x, x_mask, w=None, g=None, reverse=False, noise_scale=1.0):50    x = torch.detach(x)51    x = self.pre(x)52    if g is not None:53      g = torch.detach(g)54      x = x + self.cond(g)55    x = self.convs(x, x_mask)56    x = self.proj(x) * x_mask57 58    if not reverse:59      flows = self.flows60      assert w is not None61 62      logdet_tot_q = 0 63      h_w = self.post_pre(w)64      h_w = self.post_convs(h_w, x_mask)65      h_w = self.post_proj(h_w) * x_mask66      e_q = torch.randn(w.size(0), 2, w.size(2)).to(device=x.device, dtype=x.dtype) * x_mask67      z_q = e_q68      for flow in self.post_flows:69        z_q, logdet_q = flow(z_q, x_mask, g=(x + h_w))70        logdet_tot_q += logdet_q71      z_u, z1 = torch.split(z_q, [1, 1], 1) 72      u = torch.sigmoid(z_u) * x_mask73      z0 = (w - u) * x_mask74      logdet_tot_q += torch.sum((F.logsigmoid(z_u) + F.logsigmoid(-z_u)) * x_mask, [1,2])75      logq = torch.sum(-0.5 * (math.log(2*math.pi) + (e_q**2)) * x_mask, [1,2]) - logdet_tot_q76 77      logdet_tot = 078      z0, logdet = self.log_flow(z0, x_mask)79      logdet_tot += logdet80      z = torch.cat([z0, z1], 1)81      for flow in flows:82        z, logdet = flow(z, x_mask, g=x, reverse=reverse)83        logdet_tot = logdet_tot + logdet84      nll = torch.sum(0.5 * (math.log(2*math.pi) + (z**2)) * x_mask, [1,2]) - logdet_tot85      return nll + logq # [b]86    else:87      flows = list(reversed(self.flows))88      flows = flows[:-2] + [flows[-1]] # remove a useless vflow89      z = torch.randn(x.size(0), 2, x.size(2)).to(device=x.device, dtype=x.dtype) * noise_scale90      for flow in flows:91        z = flow(z, x_mask, g=x, reverse=reverse)92      z0, z1 = torch.split(z, [1, 1], 1)93      logw = z094      return logw95 96 97class DurationPredictor(nn.Module):98  def __init__(self, in_channels, filter_channels, kernel_size, p_dropout, gin_channels=0):99    super().__init__()100 101    self.in_channels = in_channels102    self.filter_channels = filter_channels103    self.kernel_size = kernel_size104    self.p_dropout = p_dropout105    self.gin_channels = gin_channels106 107    self.drop = nn.Dropout(p_dropout)108    self.conv_1 = nn.Conv1d(in_channels, filter_channels, kernel_size, padding=kernel_size//2)109    self.norm_1 = modules.LayerNorm(filter_channels)110    self.conv_2 = nn.Conv1d(filter_channels, filter_channels, kernel_size, padding=kernel_size//2)111    self.norm_2 = modules.LayerNorm(filter_channels)112    self.proj = nn.Conv1d(filter_channels, 1, 1)113 114    if gin_channels != 0:115      self.cond = nn.Conv1d(gin_channels, in_channels, 1)116 117  def forward(self, x, x_mask, g=None):118    x = torch.detach(x)119    if g is not None:120      g = torch.detach(g)121      x = x + self.cond(g)122    x = self.conv_1(x * x_mask)123    x = torch.relu(x)124    x = self.norm_1(x)125    x = self.drop(x)126    x = self.conv_2(x * x_mask)127    x = torch.relu(x)128    x = self.norm_2(x)129    x = self.drop(x)130    x = self.proj(x * x_mask)131    return x * x_mask132 133 134class TextEncoder(nn.Module):135  def __init__(self,136      n_vocab,137      out_channels,138      hidden_channels,139      filter_channels,140      n_heads,141      n_layers,142      kernel_size,143      p_dropout):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 154    self.emb = nn.Embedding(n_vocab, hidden_channels)155    nn.init.normal_(self.emb.weight, 0.0, hidden_channels**-0.5)156 157    self.encoder = attentions.Encoder(158      hidden_channels,159      filter_channels,160      n_heads,161      n_layers,162      kernel_size,163      p_dropout)164    self.proj= nn.Conv1d(hidden_channels, out_channels * 2, 1)165 166  def forward(self, x, x_lengths):167    x = self.emb(x) * math.sqrt(self.hidden_channels) # [b, t, h]168    x = torch.transpose(x, 1, -1) # [b, h, t]169    x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(x.dtype)170 171    x = self.encoder(x * x_mask, x_mask)172    stats = self.proj(x) * x_mask173 174    m, logs = torch.split(stats, self.out_channels, dim=1)175    return x, m, logs, x_mask176 177 178class ResidualCouplingBlock(nn.Module):179  def __init__(self,180      channels,181      hidden_channels,182      kernel_size,183      dilation_rate,184      n_layers,185      n_flows=4,186      gin_channels=0):187    super().__init__()188    self.channels = channels189    self.hidden_channels = hidden_channels190    self.kernel_size = kernel_size191    self.dilation_rate = dilation_rate192    self.n_layers = n_layers193    self.n_flows = n_flows194    self.gin_channels = gin_channels195 196    self.flows = nn.ModuleList()197    for i in range(n_flows):198      self.flows.append(modules.ResidualCouplingLayer(channels, hidden_channels, kernel_size, dilation_rate, n_layers, gin_channels=gin_channels, mean_only=True))199      self.flows.append(modules.Flip())200 201  def forward(self, x, x_mask, g=None, reverse=False):202    if not reverse:203      for flow in self.flows:204        x, _ = flow(x, x_mask, g=g, reverse=reverse)205    else:206      for flow in reversed(self.flows):207        x = flow(x, x_mask, g=g, reverse=reverse)208    return x209 210 211class PosteriorEncoder(nn.Module):212  def __init__(self,213      in_channels,214      out_channels,215      hidden_channels,216      kernel_size,217      dilation_rate,218      n_layers,219      gin_channels=0):220    super().__init__()221    self.in_channels = in_channels222    self.out_channels = out_channels223    self.hidden_channels = hidden_channels224    self.kernel_size = kernel_size225    self.dilation_rate = dilation_rate226    self.n_layers = n_layers227    self.gin_channels = gin_channels228 229    self.pre = nn.Conv1d(in_channels, hidden_channels, 1)230    self.enc = modules.WN(hidden_channels, kernel_size, dilation_rate, n_layers, gin_channels=gin_channels)231    self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)232 233  def forward(self, x, x_lengths, g=None):234    x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(x.dtype)235    x = self.pre(x) * x_mask236    x = self.enc(x, x_mask, g=g)237    stats = self.proj(x) * x_mask238    m, logs = torch.split(stats, self.out_channels, dim=1)239    z = (m + torch.randn_like(m) * torch.exp(logs)) * x_mask240    return z, m, logs, x_mask241 242 243class Generator(torch.nn.Module):244    def __init__(self, initial_channel, resblock, resblock_kernel_sizes, resblock_dilation_sizes, upsample_rates, upsample_initial_channel, upsample_kernel_sizes, gin_channels=0):245        super(Generator, self).__init__()246        self.num_kernels = len(resblock_kernel_sizes)247        self.num_upsamples = len(upsample_rates)248        self.conv_pre = Conv1d(initial_channel, upsample_initial_channel, 7, 1, padding=3)249        resblock = modules.ResBlock1 if resblock == '1' else modules.ResBlock2250 251        self.ups = nn.ModuleList()252        for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)):253            self.ups.append(weight_norm(254                ConvTranspose1d(upsample_initial_channel//(2**i), upsample_initial_channel//(2**(i+1)),255                                k, u, padding=(k-u)//2)))256 257        self.resblocks = nn.ModuleList()258        for i in range(len(self.ups)):259            ch = upsample_initial_channel//(2**(i+1))260            for j, (k, d) in enumerate(zip(resblock_kernel_sizes, resblock_dilation_sizes)):261                self.resblocks.append(resblock(ch, k, d))262 263        self.conv_post = Conv1d(ch, 1, 7, 1, padding=3, bias=False)264        self.ups.apply(init_weights)265 266        if gin_channels != 0:267            self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1)268 269    def forward(self, x, g=None):270        x = self.conv_pre(x)271        if g is not None:272          x = x + self.cond(g)273 274        for i in range(self.num_upsamples):275            x = F.leaky_relu(x, modules.LRELU_SLOPE)276            x = self.ups[i](x)277            xs = None278            for j in range(self.num_kernels):279                if xs is None:280                    xs = self.resblocks[i*self.num_kernels+j](x)281                else:282                    xs += self.resblocks[i*self.num_kernels+j](x)283            x = xs / self.num_kernels284        x = F.leaky_relu(x)285        x = self.conv_post(x)286        x = torch.tanh(x)287 288        return x289 290    def remove_weight_norm(self):291        print('Removing weight norm...')292        for l in self.ups:293            remove_weight_norm(l)294        for l in self.resblocks:295            l.remove_weight_norm()296 297 298class DiscriminatorP(torch.nn.Module):299    def __init__(self, period, kernel_size=5, stride=3, use_spectral_norm=False):300        super(DiscriminatorP, self).__init__()301        self.period = period302        self.use_spectral_norm = use_spectral_norm303        norm_f = weight_norm if use_spectral_norm == False else spectral_norm304        self.convs = nn.ModuleList([305            norm_f(Conv2d(1, 32, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),306            norm_f(Conv2d(32, 128, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),307            norm_f(Conv2d(128, 512, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),308            norm_f(Conv2d(512, 1024, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),309            norm_f(Conv2d(1024, 1024, (kernel_size, 1), 1, padding=(get_padding(kernel_size, 1), 0))),310        ])311        self.conv_post = norm_f(Conv2d(1024, 1, (3, 1), 1, padding=(1, 0)))312 313    def forward(self, x):314        fmap = []315 316        # 1d to 2d317        b, c, t = x.shape318        if t % self.period != 0: # pad first319            n_pad = self.period - (t % self.period)320            x = F.pad(x, (0, n_pad), "reflect")321            t = t + n_pad322        x = x.view(b, c, t // self.period, self.period)323 324        for l in self.convs:325            x = l(x)326            x = F.leaky_relu(x, modules.LRELU_SLOPE)327            fmap.append(x)328        x = self.conv_post(x)329        fmap.append(x)330        x = torch.flatten(x, 1, -1)331 332        return x, fmap333 334 335class DiscriminatorS(torch.nn.Module):336    def __init__(self, use_spectral_norm=False):337        super(DiscriminatorS, self).__init__()338        norm_f = weight_norm if use_spectral_norm == False else spectral_norm339        self.convs = nn.ModuleList([340            norm_f(Conv1d(1, 16, 15, 1, padding=7)),341            norm_f(Conv1d(16, 64, 41, 4, groups=4, padding=20)),342            norm_f(Conv1d(64, 256, 41, 4, groups=16, padding=20)),343            norm_f(Conv1d(256, 1024, 41, 4, groups=64, padding=20)),344            norm_f(Conv1d(1024, 1024, 41, 4, groups=256, padding=20)),345            norm_f(Conv1d(1024, 1024, 5, 1, padding=2)),346        ])347        self.conv_post = norm_f(Conv1d(1024, 1, 3, 1, padding=1))348 349    def forward(self, x):350        fmap = []351 352        for l in self.convs:353            x = l(x)354            x = F.leaky_relu(x, modules.LRELU_SLOPE)355            fmap.append(x)356        x = self.conv_post(x)357        fmap.append(x)358        x = torch.flatten(x, 1, -1)359 360        return x, fmap361 362 363class MultiPeriodDiscriminator(torch.nn.Module):364    def __init__(self, use_spectral_norm=False):365        super(MultiPeriodDiscriminator, self).__init__()366        periods = [2,3,5,7,11]367 368        discs = [DiscriminatorS(use_spectral_norm=use_spectral_norm)]369        discs = discs + [DiscriminatorP(i, use_spectral_norm=use_spectral_norm) for i in periods]370        self.discriminators = nn.ModuleList(discs)371 372    def forward(self, y, y_hat):373        y_d_rs = []374        y_d_gs = []375        fmap_rs = []376        fmap_gs = []377        for i, d in enumerate(self.discriminators):378            y_d_r, fmap_r = d(y)379            y_d_g, fmap_g = d(y_hat)380            y_d_rs.append(y_d_r)381            y_d_gs.append(y_d_g)382            fmap_rs.append(fmap_r)383            fmap_gs.append(fmap_g)384 385        return y_d_rs, y_d_gs, fmap_rs, fmap_gs386 387 388 389class SynthesizerTrn(nn.Module):390  """391  Synthesizer for Training392  """393 394  def __init__(self, 395    n_vocab,396    spec_channels,397    segment_size,398    inter_channels,399    hidden_channels,400    filter_channels,401    n_heads,402    n_layers,403    kernel_size,404    p_dropout,405    resblock, 406    resblock_kernel_sizes, 407    resblock_dilation_sizes, 408    upsample_rates, 409    upsample_initial_channel, 410    upsample_kernel_sizes,411    n_speakers=0,412    gin_channels=0,413    use_sdp=True,414    **kwargs):415 416    super().__init__()417    self.n_vocab = n_vocab418    self.spec_channels = spec_channels419    self.inter_channels = inter_channels420    self.hidden_channels = hidden_channels421    self.filter_channels = filter_channels422    self.n_heads = n_heads423    self.n_layers = n_layers424    self.kernel_size = kernel_size425    self.p_dropout = p_dropout426    self.resblock = resblock427    self.resblock_kernel_sizes = resblock_kernel_sizes428    self.resblock_dilation_sizes = resblock_dilation_sizes429    self.upsample_rates = upsample_rates430    self.upsample_initial_channel = upsample_initial_channel431    self.upsample_kernel_sizes = upsample_kernel_sizes432    self.segment_size = segment_size433    self.n_speakers = n_speakers434    self.gin_channels = gin_channels435 436    self.use_sdp = use_sdp437 438    self.enc_p = TextEncoder(n_vocab,439        inter_channels,440        hidden_channels,441        filter_channels,442        n_heads,443        n_layers,444        kernel_size,445        p_dropout)446    self.dec = Generator(inter_channels, resblock, resblock_kernel_sizes, resblock_dilation_sizes, upsample_rates, upsample_initial_channel, upsample_kernel_sizes, gin_channels=gin_channels)447    self.enc_q = PosteriorEncoder(spec_channels, inter_channels, hidden_channels, 5, 1, 16, gin_channels=gin_channels)448    self.flow = ResidualCouplingBlock(inter_channels, hidden_channels, 5, 1, 4, gin_channels=gin_channels)449 450    if use_sdp:451      self.dp = StochasticDurationPredictor(hidden_channels, 192, 3, 0.5, 4, gin_channels=gin_channels)452    else:453      self.dp = DurationPredictor(hidden_channels, 256, 3, 0.5, gin_channels=gin_channels)454 455    if n_speakers > 1:456      self.emb_g = nn.Embedding(n_speakers, gin_channels)457 458  def forward(self, x, x_lengths, y, y_lengths, sid=None):459 460    x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths)461    if self.n_speakers > 0:462      g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1]463    else:464      g = None465 466    z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g)467    z_p = self.flow(z, y_mask, g=g)468 469    with torch.no_grad():470      # negative cross-entropy471      s_p_sq_r = torch.exp(-2 * logs_p) # [b, d, t]472      neg_cent1 = torch.sum(-0.5 * math.log(2 * math.pi) - logs_p, [1], keepdim=True) # [b, 1, t_s]473      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]474      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]475      neg_cent4 = torch.sum(-0.5 * (m_p ** 2) * s_p_sq_r, [1], keepdim=True) # [b, 1, t_s]476      neg_cent = neg_cent1 + neg_cent2 + neg_cent3 + neg_cent4477 478      attn_mask = torch.unsqueeze(x_mask, 2) * torch.unsqueeze(y_mask, -1)479      attn = monotonic_align.maximum_path(neg_cent, attn_mask.squeeze(1)).unsqueeze(1).detach()480 481    w = attn.sum(2)482    if self.use_sdp:483      l_length = self.dp(x, x_mask, w, g=g)484      l_length = l_length / torch.sum(x_mask)485    else:486      logw_ = torch.log(w + 1e-6) * x_mask487      logw = self.dp(x, x_mask, g=g)488      l_length = torch.sum((logw - logw_)**2, [1,2]) / torch.sum(x_mask) # for averaging 489 490    # expand prior491    m_p = torch.matmul(attn.squeeze(1), m_p.transpose(1, 2)).transpose(1, 2)492    logs_p = torch.matmul(attn.squeeze(1), logs_p.transpose(1, 2)).transpose(1, 2)493 494    z_slice, ids_slice = commons.rand_slice_segments(z, y_lengths, self.segment_size)495    o = self.dec(z_slice, g=g)496    return o, l_length, attn, ids_slice, x_mask, y_mask, (z, z_p, m_p, logs_p, m_q, logs_q)497 498  def infer(self, x, x_lengths, sid=None, noise_scale=1, length_scale=1, noise_scale_w=1., max_len=None):499    x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths)500    if self.n_speakers > 0:501      g = self.emb_g(sid).unsqueeze(-1) # [b, h, 1]502    else:503      g = None504 505    if self.use_sdp:506      logw = self.dp(x, x_mask, g=g, reverse=True, noise_scale=noise_scale_w)507    else:508      logw = self.dp(x, x_mask, g=g)509    w = torch.exp(logw) * x_mask * length_scale510    w_ceil = torch.ceil(w)511    y_lengths = torch.clamp_min(torch.sum(w_ceil, [1, 2]), 1).long()512    y_mask = torch.unsqueeze(commons.sequence_mask(y_lengths, None), 1).to(x_mask.dtype)513    attn_mask = torch.unsqueeze(x_mask, 2) * torch.unsqueeze(y_mask, -1)514    attn = commons.generate_path(w_ceil, attn_mask)515 516    m_p = torch.matmul(attn.squeeze(1), m_p.transpose(1, 2)).transpose(1, 2) # [b, t', t], [b, t, d] -> [b, d, t']517    logs_p = torch.matmul(attn.squeeze(1), logs_p.transpose(1, 2)).transpose(1, 2) # [b, t', t], [b, t, d] -> [b, d, t']518 519    z_p = m_p + torch.randn_like(m_p) * torch.exp(logs_p) * noise_scale520    z = self.flow(z_p, y_mask, g=g, reverse=True)521    o = self.dec((z * y_mask)[:,:,:max_len], g=g)522    return o, attn, y_mask, (z, z_p, m_p, logs_p)523 524  def voice_conversion(self, y, y_lengths, sid_src, sid_tgt):525    assert self.n_speakers > 0, "n_speakers have to be larger than 0."526    g_src = self.emb_g(sid_src).unsqueeze(-1)527    g_tgt = self.emb_g(sid_tgt).unsqueeze(-1)528    z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g_src)529    z_p = self.flow(z, y_mask, g=g_src)530    z_hat = self.flow(z_p, y_mask, g=g_tgt, reverse=True)531    o_hat = self.dec(z_hat * y_mask, g=g_tgt)532    return o_hat, y_mask, (z, z_p, z_hat)533 534