CoolFace
Apppublic

ORI-Muchim/StarRailTTS

sourceHugging Facemitupdated 3y agoView on Hugging Face
0likes
models.py541 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 = 063            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        if self.n_vocab != 0:155            self.emb = nn.Embedding(n_vocab, hidden_channels)156            nn.init.normal_(self.emb.weight, 0.0, hidden_channels ** -0.5)157 158        self.encoder = attentions.Encoder(159            hidden_channels,160            filter_channels,161            n_heads,162            n_layers,163            kernel_size,164            p_dropout)165        self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)166 167    def forward(self, x, x_lengths):168        if self.n_vocab != 0:169            x = self.emb(x) * math.sqrt(self.hidden_channels)  # [b, t, h]170        x = torch.transpose(x, 1, -1)  # [b, h, t]171        x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(x.dtype)172 173        x = self.encoder(x * x_mask, x_mask)174        stats = self.proj(x) * x_mask175 176        m, logs = torch.split(stats, self.out_channels, dim=1)177        return x, m, logs, x_mask178 179 180class ResidualCouplingBlock(nn.Module):181    def __init__(self,182                 channels,183                 hidden_channels,184                 kernel_size,185                 dilation_rate,186                 n_layers,187                 n_flows=4,188                 gin_channels=0):189        super().__init__()190        self.channels = channels191        self.hidden_channels = hidden_channels192        self.kernel_size = kernel_size193        self.dilation_rate = dilation_rate194        self.n_layers = n_layers195        self.n_flows = n_flows196        self.gin_channels = gin_channels197 198        self.flows = nn.ModuleList()199        for i in range(n_flows):200            self.flows.append(201                modules.ResidualCouplingLayer(channels, hidden_channels, kernel_size, dilation_rate, n_layers,202                                              gin_channels=gin_channels, mean_only=True))203            self.flows.append(modules.Flip())204 205    def forward(self, x, x_mask, g=None, reverse=False):206        if not reverse:207            for flow in self.flows:208                x, _ = flow(x, x_mask, g=g, reverse=reverse)209        else:210            for flow in reversed(self.flows):211                x = flow(x, x_mask, g=g, reverse=reverse)212        return x213 214 215class PosteriorEncoder(nn.Module):216    def __init__(self,217                 in_channels,218                 out_channels,219                 hidden_channels,220                 kernel_size,221                 dilation_rate,222                 n_layers,223                 gin_channels=0):224        super().__init__()225        self.in_channels = in_channels226        self.out_channels = out_channels227        self.hidden_channels = hidden_channels228        self.kernel_size = kernel_size229        self.dilation_rate = dilation_rate230        self.n_layers = n_layers231        self.gin_channels = gin_channels232 233        self.pre = nn.Conv1d(in_channels, hidden_channels, 1)234        self.enc = modules.WN(hidden_channels, kernel_size, dilation_rate, n_layers, gin_channels=gin_channels)235        self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)236 237    def forward(self, x, x_lengths, g=None):238        x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(x.dtype)239        x = self.pre(x) * x_mask240        x = self.enc(x, x_mask, g=g)241        stats = self.proj(x) * x_mask242        m, logs = torch.split(stats, self.out_channels, dim=1)243        z = (m + torch.randn_like(m) * torch.exp(logs)) * x_mask244        return z, m, logs, x_mask245 246 247class Generator(torch.nn.Module):248    def __init__(self, initial_channel, resblock, resblock_kernel_sizes, resblock_dilation_sizes, upsample_rates,249                 upsample_initial_channel, upsample_kernel_sizes, gin_channels=0):250        super(Generator, self).__init__()251        self.num_kernels = len(resblock_kernel_sizes)252        self.num_upsamples = len(upsample_rates)253        self.conv_pre = Conv1d(initial_channel, upsample_initial_channel, 7, 1, padding=3)254        resblock = modules.ResBlock1 if resblock == '1' else modules.ResBlock2255 256        self.ups = nn.ModuleList()257        for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)):258            self.ups.append(weight_norm(259                ConvTranspose1d(upsample_initial_channel // (2 ** i), upsample_initial_channel // (2 ** (i + 1)),260                                k, u, padding=(k - u) // 2)))261 262        self.resblocks = nn.ModuleList()263        for i in range(len(self.ups)):264            ch = upsample_initial_channel // (2 ** (i + 1))265            for j, (k, d) in enumerate(zip(resblock_kernel_sizes, resblock_dilation_sizes)):266                self.resblocks.append(resblock(ch, k, d))267 268        self.conv_post = Conv1d(ch, 1, 7, 1, padding=3, bias=False)269        self.ups.apply(init_weights)270 271        if gin_channels != 0:272            self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1)273 274    def forward(self, x, g=None):275        x = self.conv_pre(x)276        if g is not None:277            x = x + self.cond(g)278 279        for i in range(self.num_upsamples):280            x = F.leaky_relu(x, modules.LRELU_SLOPE)281            x = self.ups[i](x)282            xs = None283            for j in range(self.num_kernels):284                if xs is None:285                    xs = self.resblocks[i * self.num_kernels + j](x)286                else:287                    xs += self.resblocks[i * self.num_kernels + j](x)288            x = xs / self.num_kernels289        x = F.leaky_relu(x)290        x = self.conv_post(x)291        x = torch.tanh(x)292 293        return x294 295    def remove_weight_norm(self):296        print('Removing weight norm...')297        for l in self.ups:298            remove_weight_norm(l)299        for l in self.resblocks:300            l.remove_weight_norm()301 302 303class DiscriminatorP(torch.nn.Module):304    def __init__(self, period, kernel_size=5, stride=3, use_spectral_norm=False):305        super(DiscriminatorP, self).__init__()306        self.period = period307        self.use_spectral_norm = use_spectral_norm308        norm_f = weight_norm if use_spectral_norm == False else spectral_norm309        self.convs = nn.ModuleList([310            norm_f(Conv2d(1, 32, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),311            norm_f(Conv2d(32, 128, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),312            norm_f(Conv2d(128, 512, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),313            norm_f(Conv2d(512, 1024, (kernel_size, 1), (stride, 1), padding=(get_padding(kernel_size, 1), 0))),314            norm_f(Conv2d(1024, 1024, (kernel_size, 1), 1, padding=(get_padding(kernel_size, 1), 0))),315        ])316        self.conv_post = norm_f(Conv2d(1024, 1, (3, 1), 1, padding=(1, 0)))317 318    def forward(self, x):319        fmap = []320 321        # 1d to 2d322        b, c, t = x.shape323        if t % self.period != 0:  # pad first324            n_pad = self.period - (t % self.period)325            x = F.pad(x, (0, n_pad), "reflect")326            t = t + n_pad327        x = x.view(b, c, t // self.period, self.period)328 329        for l in self.convs:330            x = l(x)331            x = F.leaky_relu(x, modules.LRELU_SLOPE)332            fmap.append(x)333        x = self.conv_post(x)334        fmap.append(x)335        x = torch.flatten(x, 1, -1)336 337        return x, fmap338 339 340class DiscriminatorS(torch.nn.Module):341    def __init__(self, use_spectral_norm=False):342        super(DiscriminatorS, self).__init__()343        norm_f = weight_norm if use_spectral_norm == False else spectral_norm344        self.convs = nn.ModuleList([345            norm_f(Conv1d(1, 16, 15, 1, padding=7)),346            norm_f(Conv1d(16, 64, 41, 4, groups=4, padding=20)),347            norm_f(Conv1d(64, 256, 41, 4, groups=16, padding=20)),348            norm_f(Conv1d(256, 1024, 41, 4, groups=64, padding=20)),349            norm_f(Conv1d(1024, 1024, 41, 4, groups=256, padding=20)),350            norm_f(Conv1d(1024, 1024, 5, 1, padding=2)),351        ])352        self.conv_post = norm_f(Conv1d(1024, 1, 3, 1, padding=1))353 354    def forward(self, x):355        fmap = []356 357        for l in self.convs:358            x = l(x)359            x = F.leaky_relu(x, modules.LRELU_SLOPE)360            fmap.append(x)361        x = self.conv_post(x)362        fmap.append(x)363        x = torch.flatten(x, 1, -1)364 365        return x, fmap366 367 368class MultiPeriodDiscriminator(torch.nn.Module):369    def __init__(self, use_spectral_norm=False):370        super(MultiPeriodDiscriminator, self).__init__()371        periods = [2, 3, 5, 7, 11]372 373        discs = [DiscriminatorS(use_spectral_norm=use_spectral_norm)]374        discs = discs + [DiscriminatorP(i, use_spectral_norm=use_spectral_norm) for i in periods]375        self.discriminators = nn.ModuleList(discs)376 377    def forward(self, y, y_hat):378        y_d_rs = []379        y_d_gs = []380        fmap_rs = []381        fmap_gs = []382        for i, d in enumerate(self.discriminators):383            y_d_r, fmap_r = d(y)384            y_d_g, fmap_g = d(y_hat)385            y_d_rs.append(y_d_r)386            y_d_gs.append(y_d_g)387            fmap_rs.append(fmap_r)388            fmap_gs.append(fmap_g)389 390        return y_d_rs, y_d_gs, fmap_rs, fmap_gs391 392 393class SynthesizerTrn(nn.Module):394    """395  Synthesizer for Training396  """397 398    def __init__(self,399                 n_vocab,400                 spec_channels,401                 segment_size,402                 inter_channels,403                 hidden_channels,404                 filter_channels,405                 n_heads,406                 n_layers,407                 kernel_size,408                 p_dropout,409                 resblock,410                 resblock_kernel_sizes,411                 resblock_dilation_sizes,412                 upsample_rates,413                 upsample_initial_channel,414                 upsample_kernel_sizes,415                 n_speakers=0,416                 gin_channels=0,417                 use_sdp=True,418                 **kwargs):419 420        super().__init__()421        self.n_vocab = n_vocab422        self.spec_channels = spec_channels423        self.inter_channels = inter_channels424        self.hidden_channels = hidden_channels425        self.filter_channels = filter_channels426        self.n_heads = n_heads427        self.n_layers = n_layers428        self.kernel_size = kernel_size429        self.p_dropout = p_dropout430        self.resblock = resblock431        self.resblock_kernel_sizes = resblock_kernel_sizes432        self.resblock_dilation_sizes = resblock_dilation_sizes433        self.upsample_rates = upsample_rates434        self.upsample_initial_channel = upsample_initial_channel435        self.upsample_kernel_sizes = upsample_kernel_sizes436        self.segment_size = segment_size437        self.n_speakers = n_speakers438        self.gin_channels = gin_channels439 440        self.use_sdp = use_sdp441 442        self.enc_p = TextEncoder(n_vocab,443                                 inter_channels,444                                 hidden_channels,445                                 filter_channels,446                                 n_heads,447                                 n_layers,448                                 kernel_size,449                                 p_dropout)450        self.dec = Generator(inter_channels, resblock, resblock_kernel_sizes, resblock_dilation_sizes, upsample_rates,451                             upsample_initial_channel, upsample_kernel_sizes, gin_channels=gin_channels)452        self.enc_q = PosteriorEncoder(spec_channels, inter_channels, hidden_channels, 5, 1, 16,453                                      gin_channels=gin_channels)454        self.flow = ResidualCouplingBlock(inter_channels, hidden_channels, 5, 1, 4, gin_channels=gin_channels)455 456        if use_sdp:457            self.dp = StochasticDurationPredictor(hidden_channels, 192, 3, 0.5, 4, gin_channels=gin_channels)458        else:459            self.dp = DurationPredictor(hidden_channels, 256, 3, 0.5, gin_channels=gin_channels)460 461        if n_speakers > 1:462            self.emb_g = nn.Embedding(n_speakers, gin_channels)463 464    def forward(self, x, x_lengths, y, y_lengths, sid=None):465 466        x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths)467        if self.n_speakers > 1:468            g = self.emb_g(sid).unsqueeze(-1)  # [b, h, 1]469        else:470            g = None471 472        z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g)473        z_p = self.flow(z, y_mask, g=g)474 475        with torch.no_grad():476            # negative cross-entropy477            s_p_sq_r = torch.exp(-2 * logs_p)  # [b, d, t]478            neg_cent1 = torch.sum(-0.5 * math.log(2 * math.pi) - logs_p, [1], keepdim=True)  # [b, 1, t_s]479            neg_cent2 = torch.matmul(-0.5 * (z_p ** 2).transpose(1, 2),480                                     s_p_sq_r)  # [b, t_t, d] x [b, d, t_s] = [b, t_t, t_s]481            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]482            neg_cent4 = torch.sum(-0.5 * (m_p ** 2) * s_p_sq_r, [1], keepdim=True)  # [b, 1, t_s]483            neg_cent = neg_cent1 + neg_cent2 + neg_cent3 + neg_cent4484 485            attn_mask = torch.unsqueeze(x_mask, 2) * torch.unsqueeze(y_mask, -1)486            attn = monotonic_align.maximum_path(neg_cent, attn_mask.squeeze(1)).unsqueeze(1).detach()487 488        w = attn.sum(2)489        if self.use_sdp:490            l_length = self.dp(x, x_mask, w, g=g)491            l_length = l_length / torch.sum(x_mask)492        else:493            logw_ = torch.log(w + 1e-6) * x_mask494            logw = self.dp(x, x_mask, g=g)495            l_length = torch.sum((logw - logw_) ** 2, [1, 2]) / torch.sum(x_mask)  # for averaging496 497        # expand prior498        m_p = torch.matmul(attn.squeeze(1), m_p.transpose(1, 2)).transpose(1, 2)499        logs_p = torch.matmul(attn.squeeze(1), logs_p.transpose(1, 2)).transpose(1, 2)500 501        z_slice, ids_slice = commons.rand_slice_segments(z, y_lengths, self.segment_size)502        o = self.dec(z_slice, g=g)503        return o, l_length, attn, ids_slice, x_mask, y_mask, (z, z_p, m_p, logs_p, m_q, logs_q)504 505    def infer(self, x, x_lengths, sid=None, noise_scale=1, length_scale=1, noise_scale_w=1., max_len=None):506        x, m_p, logs_p, x_mask = self.enc_p(x, x_lengths)507        if self.n_speakers > 1:508            g = self.emb_g(sid).unsqueeze(-1)  # [b, h, 1]509        else:510            g = None511 512        if self.use_sdp:513            logw = self.dp(x, x_mask, g=g, reverse=True, noise_scale=noise_scale_w)514        else:515            logw = self.dp(x, x_mask, g=g)516        w = torch.exp(logw) * x_mask * length_scale517        w_ceil = torch.ceil(w)518        y_lengths = torch.clamp_min(torch.sum(w_ceil, [1, 2]), 1).long()519        y_mask = torch.unsqueeze(commons.sequence_mask(y_lengths, None), 1).to(x_mask.dtype)520        attn_mask = torch.unsqueeze(x_mask, 2) * torch.unsqueeze(y_mask, -1)521        attn = commons.generate_path(w_ceil, attn_mask)522 523        m_p = torch.matmul(attn.squeeze(1), m_p.transpose(1, 2)).transpose(1, 2)  # [b, t', t], [b, t, d] -> [b, d, t']524        logs_p = torch.matmul(attn.squeeze(1), logs_p.transpose(1, 2)).transpose(1,525                                                                                 2)  # [b, t', t], [b, t, d] -> [b, d, t']526 527        z_p = m_p + torch.randn_like(m_p) * torch.exp(logs_p) * noise_scale528        z = self.flow(z_p, y_mask, g=g, reverse=True)529        o = self.dec((z * y_mask)[:, :, :max_len], g=g)530        return o, attn, y_mask, (z, z_p, m_p, logs_p)531 532    def voice_conversion(self, y, y_lengths, sid_src, sid_tgt):533        assert self.n_speakers > 1, "n_speakers have to be larger than 1."534        g_src = self.emb_g(sid_src).unsqueeze(-1)535        g_tgt = self.emb_g(sid_tgt).unsqueeze(-1)536        z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g_src)537        z_p = self.flow(z, y_mask, g=g_src)538        z_hat = self.flow(z_p, y_mask, g=g_tgt, reverse=True)539        o_hat = self.dec(z_hat * y_mask, g=g_tgt)540        return o_hat, y_mask, (z, z_p, z_hat)541