CoolFace
Apppublic

Clicko777/RVC_HFv2

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
models.py1144 linesDownload Raw Back to infer_pack
1import math, pdb, os2from time import time as ttime3import torch4from torch import nn5from torch.nn import functional as F6from lib.infer_pack import modules7from lib.infer_pack import attentions8from lib.infer_pack import commons9from lib.infer_pack.commons import init_weights, get_padding10from torch.nn import Conv1d, ConvTranspose1d, AvgPool1d, Conv2d11from torch.nn.utils import weight_norm, remove_weight_norm, spectral_norm12from lib.infer_pack.commons import init_weights13import numpy as np14from lib.infer_pack import commons15 16 17class TextEncoder256(nn.Module):18    def __init__(19        self,20        out_channels,21        hidden_channels,22        filter_channels,23        n_heads,24        n_layers,25        kernel_size,26        p_dropout,27        f0=True,28    ):29        super().__init__()30        self.out_channels = out_channels31        self.hidden_channels = hidden_channels32        self.filter_channels = filter_channels33        self.n_heads = n_heads34        self.n_layers = n_layers35        self.kernel_size = kernel_size36        self.p_dropout = p_dropout37        self.emb_phone = nn.Linear(256, hidden_channels)38        self.lrelu = nn.LeakyReLU(0.1, inplace=True)39        if f0 == True:40            self.emb_pitch = nn.Embedding(256, hidden_channels)  # pitch 25641        self.encoder = attentions.Encoder(42            hidden_channels, filter_channels, n_heads, n_layers, kernel_size, p_dropout43        )44        self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)45 46    def forward(self, phone, pitch, lengths):47        if pitch == None:48            x = self.emb_phone(phone)49        else:50            x = self.emb_phone(phone) + self.emb_pitch(pitch)51        x = x * math.sqrt(self.hidden_channels)  # [b, t, h]52        x = self.lrelu(x)53        x = torch.transpose(x, 1, -1)  # [b, h, t]54        x_mask = torch.unsqueeze(commons.sequence_mask(lengths, x.size(2)), 1).to(55            x.dtype56        )57        x = self.encoder(x * x_mask, x_mask)58        stats = self.proj(x) * x_mask59 60        m, logs = torch.split(stats, self.out_channels, dim=1)61        return m, logs, x_mask62 63 64class TextEncoder768(nn.Module):65    def __init__(66        self,67        out_channels,68        hidden_channels,69        filter_channels,70        n_heads,71        n_layers,72        kernel_size,73        p_dropout,74        f0=True,75    ):76        super().__init__()77        self.out_channels = out_channels78        self.hidden_channels = hidden_channels79        self.filter_channels = filter_channels80        self.n_heads = n_heads81        self.n_layers = n_layers82        self.kernel_size = kernel_size83        self.p_dropout = p_dropout84        self.emb_phone = nn.Linear(768, hidden_channels)85        self.lrelu = nn.LeakyReLU(0.1, inplace=True)86        if f0 == True:87            self.emb_pitch = nn.Embedding(256, hidden_channels)  # pitch 25688        self.encoder = attentions.Encoder(89            hidden_channels, filter_channels, n_heads, n_layers, kernel_size, p_dropout90        )91        self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)92 93    def forward(self, phone, pitch, lengths):94        if pitch == None:95            x = self.emb_phone(phone)96        else:97            x = self.emb_phone(phone) + self.emb_pitch(pitch)98        x = x * math.sqrt(self.hidden_channels)  # [b, t, h]99        x = self.lrelu(x)100        x = torch.transpose(x, 1, -1)  # [b, h, t]101        x_mask = torch.unsqueeze(commons.sequence_mask(lengths, x.size(2)), 1).to(102            x.dtype103        )104        x = self.encoder(x * x_mask, x_mask)105        stats = self.proj(x) * x_mask106 107        m, logs = torch.split(stats, self.out_channels, dim=1)108        return m, logs, x_mask109 110 111class ResidualCouplingBlock(nn.Module):112    def __init__(113        self,114        channels,115        hidden_channels,116        kernel_size,117        dilation_rate,118        n_layers,119        n_flows=4,120        gin_channels=0,121    ):122        super().__init__()123        self.channels = channels124        self.hidden_channels = hidden_channels125        self.kernel_size = kernel_size126        self.dilation_rate = dilation_rate127        self.n_layers = n_layers128        self.n_flows = n_flows129        self.gin_channels = gin_channels130 131        self.flows = nn.ModuleList()132        for i in range(n_flows):133            self.flows.append(134                modules.ResidualCouplingLayer(135                    channels,136                    hidden_channels,137                    kernel_size,138                    dilation_rate,139                    n_layers,140                    gin_channels=gin_channels,141                    mean_only=True,142                )143            )144            self.flows.append(modules.Flip())145 146    def forward(self, x, x_mask, g=None, reverse=False):147        if not reverse:148            for flow in self.flows:149                x, _ = flow(x, x_mask, g=g, reverse=reverse)150        else:151            for flow in reversed(self.flows):152                x = flow(x, x_mask, g=g, reverse=reverse)153        return x154 155    def remove_weight_norm(self):156        for i in range(self.n_flows):157            self.flows[i * 2].remove_weight_norm()158 159 160class PosteriorEncoder(nn.Module):161    def __init__(162        self,163        in_channels,164        out_channels,165        hidden_channels,166        kernel_size,167        dilation_rate,168        n_layers,169        gin_channels=0,170    ):171        super().__init__()172        self.in_channels = in_channels173        self.out_channels = out_channels174        self.hidden_channels = hidden_channels175        self.kernel_size = kernel_size176        self.dilation_rate = dilation_rate177        self.n_layers = n_layers178        self.gin_channels = gin_channels179 180        self.pre = nn.Conv1d(in_channels, hidden_channels, 1)181        self.enc = modules.WN(182            hidden_channels,183            kernel_size,184            dilation_rate,185            n_layers,186            gin_channels=gin_channels,187        )188        self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)189 190    def forward(self, x, x_lengths, g=None):191        x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(192            x.dtype193        )194        x = self.pre(x) * x_mask195        x = self.enc(x, x_mask, g=g)196        stats = self.proj(x) * x_mask197        m, logs = torch.split(stats, self.out_channels, dim=1)198        z = (m + torch.randn_like(m) * torch.exp(logs)) * x_mask199        return z, m, logs, x_mask200 201    def remove_weight_norm(self):202        self.enc.remove_weight_norm()203 204 205class Generator(torch.nn.Module):206    def __init__(207        self,208        initial_channel,209        resblock,210        resblock_kernel_sizes,211        resblock_dilation_sizes,212        upsample_rates,213        upsample_initial_channel,214        upsample_kernel_sizes,215        gin_channels=0,216    ):217        super(Generator, self).__init__()218        self.num_kernels = len(resblock_kernel_sizes)219        self.num_upsamples = len(upsample_rates)220        self.conv_pre = Conv1d(221            initial_channel, upsample_initial_channel, 7, 1, padding=3222        )223        resblock = modules.ResBlock1 if resblock == "1" else modules.ResBlock2224 225        self.ups = nn.ModuleList()226        for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)):227            self.ups.append(228                weight_norm(229                    ConvTranspose1d(230                        upsample_initial_channel // (2**i),231                        upsample_initial_channel // (2 ** (i + 1)),232                        k,233                        u,234                        padding=(k - u) // 2,235                    )236                )237            )238 239        self.resblocks = nn.ModuleList()240        for i in range(len(self.ups)):241            ch = upsample_initial_channel // (2 ** (i + 1))242            for j, (k, d) in enumerate(243                zip(resblock_kernel_sizes, resblock_dilation_sizes)244            ):245                self.resblocks.append(resblock(ch, k, d))246 247        self.conv_post = Conv1d(ch, 1, 7, 1, padding=3, bias=False)248        self.ups.apply(init_weights)249 250        if gin_channels != 0:251            self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1)252 253    def forward(self, x, g=None):254        x = self.conv_pre(x)255        if g is not None:256            x = x + self.cond(g)257 258        for i in range(self.num_upsamples):259            x = F.leaky_relu(x, modules.LRELU_SLOPE)260            x = self.ups[i](x)261            xs = None262            for j in range(self.num_kernels):263                if xs is None:264                    xs = self.resblocks[i * self.num_kernels + j](x)265                else:266                    xs += self.resblocks[i * self.num_kernels + j](x)267            x = xs / self.num_kernels268        x = F.leaky_relu(x)269        x = self.conv_post(x)270        x = torch.tanh(x)271 272        return x273 274    def remove_weight_norm(self):275        for l in self.ups:276            remove_weight_norm(l)277        for l in self.resblocks:278            l.remove_weight_norm()279 280 281class SineGen(torch.nn.Module):282    """Definition of sine generator283    SineGen(samp_rate, harmonic_num = 0,284            sine_amp = 0.1, noise_std = 0.003,285            voiced_threshold = 0,286            flag_for_pulse=False)287    samp_rate: sampling rate in Hz288    harmonic_num: number of harmonic overtones (default 0)289    sine_amp: amplitude of sine-wavefrom (default 0.1)290    noise_std: std of Gaussian noise (default 0.003)291    voiced_thoreshold: F0 threshold for U/V classification (default 0)292    flag_for_pulse: this SinGen is used inside PulseGen (default False)293    Note: when flag_for_pulse is True, the first time step of a voiced294        segment is always sin(np.pi) or cos(0)295    """296 297    def __init__(298        self,299        samp_rate,300        harmonic_num=0,301        sine_amp=0.1,302        noise_std=0.003,303        voiced_threshold=0,304        flag_for_pulse=False,305    ):306        super(SineGen, self).__init__()307        self.sine_amp = sine_amp308        self.noise_std = noise_std309        self.harmonic_num = harmonic_num310        self.dim = self.harmonic_num + 1311        self.sampling_rate = samp_rate312        self.voiced_threshold = voiced_threshold313 314    def _f02uv(self, f0):315        # generate uv signal316        uv = torch.ones_like(f0)317        uv = uv * (f0 > self.voiced_threshold)318        if uv.device.type == "privateuseone":  # for DirectML319            uv = uv.float()320        return uv321 322    def forward(self, f0, upp):323        """sine_tensor, uv = forward(f0)324        input F0: tensor(batchsize=1, length, dim=1)325                  f0 for unvoiced steps should be 0326        output sine_tensor: tensor(batchsize=1, length, dim)327        output uv: tensor(batchsize=1, length, 1)328        """329        with torch.no_grad():330            f0 = f0[:, None].transpose(1, 2)331            f0_buf = torch.zeros(f0.shape[0], f0.shape[1], self.dim, device=f0.device)332            # fundamental component333            f0_buf[:, :, 0] = f0[:, :, 0]334            for idx in np.arange(self.harmonic_num):335                f0_buf[:, :, idx + 1] = f0_buf[:, :, 0] * (336                    idx + 2337                )  # idx + 2: the (idx+1)-th overtone, (idx+2)-th harmonic338            rad_values = (f0_buf / self.sampling_rate) % 1  ###%1意味着n_har的乘积无法后处理优化339            rand_ini = torch.rand(340                f0_buf.shape[0], f0_buf.shape[2], device=f0_buf.device341            )342            rand_ini[:, 0] = 0343            rad_values[:, 0, :] = rad_values[:, 0, :] + rand_ini344            tmp_over_one = torch.cumsum(rad_values, 1)  # % 1  #####%1意味着后面的cumsum无法再优化345            tmp_over_one *= upp346            tmp_over_one = F.interpolate(347                tmp_over_one.transpose(2, 1),348                scale_factor=upp,349                mode="linear",350                align_corners=True,351            ).transpose(2, 1)352            rad_values = F.interpolate(353                rad_values.transpose(2, 1), scale_factor=upp, mode="nearest"354            ).transpose(355                2, 1356            )  #######357            tmp_over_one %= 1358            tmp_over_one_idx = (tmp_over_one[:, 1:, :] - tmp_over_one[:, :-1, :]) < 0359            cumsum_shift = torch.zeros_like(rad_values)360            cumsum_shift[:, 1:, :] = tmp_over_one_idx * -1.0361            sine_waves = torch.sin(362                torch.cumsum(rad_values + cumsum_shift, dim=1) * 2 * np.pi363            )364            sine_waves = sine_waves * self.sine_amp365            uv = self._f02uv(f0)366            uv = F.interpolate(367                uv.transpose(2, 1), scale_factor=upp, mode="nearest"368            ).transpose(2, 1)369            noise_amp = uv * self.noise_std + (1 - uv) * self.sine_amp / 3370            noise = noise_amp * torch.randn_like(sine_waves)371            sine_waves = sine_waves * uv + noise372        return sine_waves, uv, noise373 374 375class SourceModuleHnNSF(torch.nn.Module):376    """SourceModule for hn-nsf377    SourceModule(sampling_rate, harmonic_num=0, sine_amp=0.1,378                 add_noise_std=0.003, voiced_threshod=0)379    sampling_rate: sampling_rate in Hz380    harmonic_num: number of harmonic above F0 (default: 0)381    sine_amp: amplitude of sine source signal (default: 0.1)382    add_noise_std: std of additive Gaussian noise (default: 0.003)383        note that amplitude of noise in unvoiced is decided384        by sine_amp385    voiced_threshold: threhold to set U/V given F0 (default: 0)386    Sine_source, noise_source = SourceModuleHnNSF(F0_sampled)387    F0_sampled (batchsize, length, 1)388    Sine_source (batchsize, length, 1)389    noise_source (batchsize, length 1)390    uv (batchsize, length, 1)391    """392 393    def __init__(394        self,395        sampling_rate,396        harmonic_num=0,397        sine_amp=0.1,398        add_noise_std=0.003,399        voiced_threshod=0,400        is_half=True,401    ):402        super(SourceModuleHnNSF, self).__init__()403 404        self.sine_amp = sine_amp405        self.noise_std = add_noise_std406        self.is_half = is_half407        # to produce sine waveforms408        self.l_sin_gen = SineGen(409            sampling_rate, harmonic_num, sine_amp, add_noise_std, voiced_threshod410        )411 412        # to merge source harmonics into a single excitation413        self.l_linear = torch.nn.Linear(harmonic_num + 1, 1)414        self.l_tanh = torch.nn.Tanh()415 416    def forward(self, x, upp=None):417        sine_wavs, uv, _ = self.l_sin_gen(x, upp)418        if self.is_half:419            sine_wavs = sine_wavs.half()420        sine_merge = self.l_tanh(self.l_linear(sine_wavs))421        return sine_merge, None, None  # noise, uv422 423 424class GeneratorNSF(torch.nn.Module):425    def __init__(426        self,427        initial_channel,428        resblock,429        resblock_kernel_sizes,430        resblock_dilation_sizes,431        upsample_rates,432        upsample_initial_channel,433        upsample_kernel_sizes,434        gin_channels,435        sr,436        is_half=False,437    ):438        super(GeneratorNSF, self).__init__()439        self.num_kernels = len(resblock_kernel_sizes)440        self.num_upsamples = len(upsample_rates)441 442        self.f0_upsamp = torch.nn.Upsample(scale_factor=np.prod(upsample_rates))443        self.m_source = SourceModuleHnNSF(444            sampling_rate=sr, harmonic_num=0, is_half=is_half445        )446        self.noise_convs = nn.ModuleList()447        self.conv_pre = Conv1d(448            initial_channel, upsample_initial_channel, 7, 1, padding=3449        )450        resblock = modules.ResBlock1 if resblock == "1" else modules.ResBlock2451 452        self.ups = nn.ModuleList()453        for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)):454            c_cur = upsample_initial_channel // (2 ** (i + 1))455            self.ups.append(456                weight_norm(457                    ConvTranspose1d(458                        upsample_initial_channel // (2**i),459                        upsample_initial_channel // (2 ** (i + 1)),460                        k,461                        u,462                        padding=(k - u) // 2,463                    )464                )465            )466            if i + 1 < len(upsample_rates):467                stride_f0 = np.prod(upsample_rates[i + 1 :])468                self.noise_convs.append(469                    Conv1d(470                        1,471                        c_cur,472                        kernel_size=stride_f0 * 2,473                        stride=stride_f0,474                        padding=stride_f0 // 2,475                    )476                )477            else:478                self.noise_convs.append(Conv1d(1, c_cur, kernel_size=1))479 480        self.resblocks = nn.ModuleList()481        for i in range(len(self.ups)):482            ch = upsample_initial_channel // (2 ** (i + 1))483            for j, (k, d) in enumerate(484                zip(resblock_kernel_sizes, resblock_dilation_sizes)485            ):486                self.resblocks.append(resblock(ch, k, d))487 488        self.conv_post = Conv1d(ch, 1, 7, 1, padding=3, bias=False)489        self.ups.apply(init_weights)490 491        if gin_channels != 0:492            self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1)493 494        self.upp = np.prod(upsample_rates)495 496    def forward(self, x, f0, g=None):497        har_source, noi_source, uv = self.m_source(f0, self.upp)498        har_source = har_source.transpose(1, 2)499        x = self.conv_pre(x)500        if g is not None:501            x = x + self.cond(g)502 503        for i in range(self.num_upsamples):504            x = F.leaky_relu(x, modules.LRELU_SLOPE)505            x = self.ups[i](x)506            x_source = self.noise_convs[i](har_source)507            x = x + x_source508            xs = None509            for j in range(self.num_kernels):510                if xs is None:511                    xs = self.resblocks[i * self.num_kernels + j](x)512                else:513                    xs += self.resblocks[i * self.num_kernels + j](x)514            x = xs / self.num_kernels515        x = F.leaky_relu(x)516        x = self.conv_post(x)517        x = torch.tanh(x)518        return x519 520    def remove_weight_norm(self):521        for l in self.ups:522            remove_weight_norm(l)523        for l in self.resblocks:524            l.remove_weight_norm()525 526 527sr2sr = {528    "32k": 32000,529    "40k": 40000,530    "48k": 48000,531}532 533 534class SynthesizerTrnMs256NSFsid(nn.Module):535    def __init__(536        self,537        spec_channels,538        segment_size,539        inter_channels,540        hidden_channels,541        filter_channels,542        n_heads,543        n_layers,544        kernel_size,545        p_dropout,546        resblock,547        resblock_kernel_sizes,548        resblock_dilation_sizes,549        upsample_rates,550        upsample_initial_channel,551        upsample_kernel_sizes,552        spk_embed_dim,553        gin_channels,554        sr,555        **kwargs556    ):557        super().__init__()558        if type(sr) == type("strr"):559            sr = sr2sr[sr]560        self.spec_channels = spec_channels561        self.inter_channels = inter_channels562        self.hidden_channels = hidden_channels563        self.filter_channels = filter_channels564        self.n_heads = n_heads565        self.n_layers = n_layers566        self.kernel_size = kernel_size567        self.p_dropout = p_dropout568        self.resblock = resblock569        self.resblock_kernel_sizes = resblock_kernel_sizes570        self.resblock_dilation_sizes = resblock_dilation_sizes571        self.upsample_rates = upsample_rates572        self.upsample_initial_channel = upsample_initial_channel573        self.upsample_kernel_sizes = upsample_kernel_sizes574        self.segment_size = segment_size575        self.gin_channels = gin_channels576        # self.hop_length = hop_length#577        self.spk_embed_dim = spk_embed_dim578        self.enc_p = TextEncoder256(579            inter_channels,580            hidden_channels,581            filter_channels,582            n_heads,583            n_layers,584            kernel_size,585            p_dropout,586        )587        self.dec = GeneratorNSF(588            inter_channels,589            resblock,590            resblock_kernel_sizes,591            resblock_dilation_sizes,592            upsample_rates,593            upsample_initial_channel,594            upsample_kernel_sizes,595            gin_channels=gin_channels,596            sr=sr,597            is_half=kwargs["is_half"],598        )599        self.enc_q = PosteriorEncoder(600            spec_channels,601            inter_channels,602            hidden_channels,603            5,604            1,605            16,606            gin_channels=gin_channels,607        )608        self.flow = ResidualCouplingBlock(609            inter_channels, hidden_channels, 5, 1, 3, gin_channels=gin_channels610        )611        self.emb_g = nn.Embedding(self.spk_embed_dim, gin_channels)612        print("gin_channels:", gin_channels, "self.spk_embed_dim:", self.spk_embed_dim)613 614    def remove_weight_norm(self):615        self.dec.remove_weight_norm()616        self.flow.remove_weight_norm()617        self.enc_q.remove_weight_norm()618 619    def forward(620        self, phone, phone_lengths, pitch, pitchf, y, y_lengths, ds621    ):  # 这里ds是id,[bs,1]622        # print(1,pitch.shape)#[bs,t]623        g = self.emb_g(ds).unsqueeze(-1)  # [b, 256, 1]##1是t,广播的624        m_p, logs_p, x_mask = self.enc_p(phone, pitch, phone_lengths)625        z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g)626        z_p = self.flow(z, y_mask, g=g)627        z_slice, ids_slice = commons.rand_slice_segments(628            z, y_lengths, self.segment_size629        )630        # print(-1,pitchf.shape,ids_slice,self.segment_size,self.hop_length,self.segment_size//self.hop_length)631        pitchf = commons.slice_segments2(pitchf, ids_slice, self.segment_size)632        # print(-2,pitchf.shape,z_slice.shape)633        o = self.dec(z_slice, pitchf, g=g)634        return o, ids_slice, x_mask, y_mask, (z, z_p, m_p, logs_p, m_q, logs_q)635 636    def infer(self, phone, phone_lengths, pitch, nsff0, sid, rate=None):637        g = self.emb_g(sid).unsqueeze(-1)638        m_p, logs_p, x_mask = self.enc_p(phone, pitch, phone_lengths)639        z_p = (m_p + torch.exp(logs_p) * torch.randn_like(m_p) * 0.66666) * x_mask640        if rate:641            head = int(z_p.shape[2] * rate)642            z_p = z_p[:, :, -head:]643            x_mask = x_mask[:, :, -head:]644            nsff0 = nsff0[:, -head:]645        z = self.flow(z_p, x_mask, g=g, reverse=True)646        o = self.dec(z * x_mask, nsff0, g=g)647        return o, x_mask, (z, z_p, m_p, logs_p)648 649 650class SynthesizerTrnMs768NSFsid(nn.Module):651    def __init__(652        self,653        spec_channels,654        segment_size,655        inter_channels,656        hidden_channels,657        filter_channels,658        n_heads,659        n_layers,660        kernel_size,661        p_dropout,662        resblock,663        resblock_kernel_sizes,664        resblock_dilation_sizes,665        upsample_rates,666        upsample_initial_channel,667        upsample_kernel_sizes,668        spk_embed_dim,669        gin_channels,670        sr,671        **kwargs672    ):673        super().__init__()674        if type(sr) == type("strr"):675            sr = sr2sr[sr]676        self.spec_channels = spec_channels677        self.inter_channels = inter_channels678        self.hidden_channels = hidden_channels679        self.filter_channels = filter_channels680        self.n_heads = n_heads681        self.n_layers = n_layers682        self.kernel_size = kernel_size683        self.p_dropout = p_dropout684        self.resblock = resblock685        self.resblock_kernel_sizes = resblock_kernel_sizes686        self.resblock_dilation_sizes = resblock_dilation_sizes687        self.upsample_rates = upsample_rates688        self.upsample_initial_channel = upsample_initial_channel689        self.upsample_kernel_sizes = upsample_kernel_sizes690        self.segment_size = segment_size691        self.gin_channels = gin_channels692        # self.hop_length = hop_length#693        self.spk_embed_dim = spk_embed_dim694        self.enc_p = TextEncoder768(695            inter_channels,696            hidden_channels,697            filter_channels,698            n_heads,699            n_layers,700            kernel_size,701            p_dropout,702        )703        self.dec = GeneratorNSF(704            inter_channels,705            resblock,706            resblock_kernel_sizes,707            resblock_dilation_sizes,708            upsample_rates,709            upsample_initial_channel,710            upsample_kernel_sizes,711            gin_channels=gin_channels,712            sr=sr,713            is_half=kwargs["is_half"],714        )715        self.enc_q = PosteriorEncoder(716            spec_channels,717            inter_channels,718            hidden_channels,719            5,720            1,721            16,722            gin_channels=gin_channels,723        )724        self.flow = ResidualCouplingBlock(725            inter_channels, hidden_channels, 5, 1, 3, gin_channels=gin_channels726        )727        self.emb_g = nn.Embedding(self.spk_embed_dim, gin_channels)728        print("gin_channels:", gin_channels, "self.spk_embed_dim:", self.spk_embed_dim)729 730    def remove_weight_norm(self):731        self.dec.remove_weight_norm()732        self.flow.remove_weight_norm()733        self.enc_q.remove_weight_norm()734 735    def forward(736        self, phone, phone_lengths, pitch, pitchf, y, y_lengths, ds737    ):  # 这里ds是id,[bs,1]738        # print(1,pitch.shape)#[bs,t]739        g = self.emb_g(ds).unsqueeze(-1)  # [b, 256, 1]##1是t,广播的740        m_p, logs_p, x_mask = self.enc_p(phone, pitch, phone_lengths)741        z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g)742        z_p = self.flow(z, y_mask, g=g)743        z_slice, ids_slice = commons.rand_slice_segments(744            z, y_lengths, self.segment_size745        )746        # print(-1,pitchf.shape,ids_slice,self.segment_size,self.hop_length,self.segment_size//self.hop_length)747        pitchf = commons.slice_segments2(pitchf, ids_slice, self.segment_size)748        # print(-2,pitchf.shape,z_slice.shape)749        o = self.dec(z_slice, pitchf, g=g)750        return o, ids_slice, x_mask, y_mask, (z, z_p, m_p, logs_p, m_q, logs_q)751 752    def infer(self, phone, phone_lengths, pitch, nsff0, sid, rate=None):753        g = self.emb_g(sid).unsqueeze(-1)754        m_p, logs_p, x_mask = self.enc_p(phone, pitch, phone_lengths)755        z_p = (m_p + torch.exp(logs_p) * torch.randn_like(m_p) * 0.66666) * x_mask756        if rate:757            head = int(z_p.shape[2] * rate)758            z_p = z_p[:, :, -head:]759            x_mask = x_mask[:, :, -head:]760            nsff0 = nsff0[:, -head:]761        z = self.flow(z_p, x_mask, g=g, reverse=True)762        o = self.dec(z * x_mask, nsff0, g=g)763        return o, x_mask, (z, z_p, m_p, logs_p)764 765 766class SynthesizerTrnMs256NSFsid_nono(nn.Module):767    def __init__(768        self,769        spec_channels,770        segment_size,771        inter_channels,772        hidden_channels,773        filter_channels,774        n_heads,775        n_layers,776        kernel_size,777        p_dropout,778        resblock,779        resblock_kernel_sizes,780        resblock_dilation_sizes,781        upsample_rates,782        upsample_initial_channel,783        upsample_kernel_sizes,784        spk_embed_dim,785        gin_channels,786        sr=None,787        **kwargs788    ):789        super().__init__()790        self.spec_channels = spec_channels791        self.inter_channels = inter_channels792        self.hidden_channels = hidden_channels793        self.filter_channels = filter_channels794        self.n_heads = n_heads795        self.n_layers = n_layers796        self.kernel_size = kernel_size797        self.p_dropout = p_dropout798        self.resblock = resblock799        self.resblock_kernel_sizes = resblock_kernel_sizes800        self.resblock_dilation_sizes = resblock_dilation_sizes801        self.upsample_rates = upsample_rates802        self.upsample_initial_channel = upsample_initial_channel803        self.upsample_kernel_sizes = upsample_kernel_sizes804        self.segment_size = segment_size805        self.gin_channels = gin_channels806        # self.hop_length = hop_length#807        self.spk_embed_dim = spk_embed_dim808        self.enc_p = TextEncoder256(809            inter_channels,810            hidden_channels,811            filter_channels,812            n_heads,813            n_layers,814            kernel_size,815            p_dropout,816            f0=False,817        )818        self.dec = Generator(819            inter_channels,820            resblock,821            resblock_kernel_sizes,822            resblock_dilation_sizes,823            upsample_rates,824            upsample_initial_channel,825            upsample_kernel_sizes,826            gin_channels=gin_channels,827        )828        self.enc_q = PosteriorEncoder(829            spec_channels,830            inter_channels,831            hidden_channels,832            5,833            1,834            16,835            gin_channels=gin_channels,836        )837        self.flow = ResidualCouplingBlock(838            inter_channels, hidden_channels, 5, 1, 3, gin_channels=gin_channels839        )840        self.emb_g = nn.Embedding(self.spk_embed_dim, gin_channels)841        print("gin_channels:", gin_channels, "self.spk_embed_dim:", self.spk_embed_dim)842 843    def remove_weight_norm(self):844        self.dec.remove_weight_norm()845        self.flow.remove_weight_norm()846        self.enc_q.remove_weight_norm()847 848    def forward(self, phone, phone_lengths, y, y_lengths, ds):  # 这里ds是id,[bs,1]849        g = self.emb_g(ds).unsqueeze(-1)  # [b, 256, 1]##1是t,广播的850        m_p, logs_p, x_mask = self.enc_p(phone, None, phone_lengths)851        z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g)852        z_p = self.flow(z, y_mask, g=g)853        z_slice, ids_slice = commons.rand_slice_segments(854            z, y_lengths, self.segment_size855        )856        o = self.dec(z_slice, g=g)857        return o, ids_slice, x_mask, y_mask, (z, z_p, m_p, logs_p, m_q, logs_q)858 859    def infer(self, phone, phone_lengths, sid, rate=None):860        g = self.emb_g(sid).unsqueeze(-1)861        m_p, logs_p, x_mask = self.enc_p(phone, None, phone_lengths)862        z_p = (m_p + torch.exp(logs_p) * torch.randn_like(m_p) * 0.66666) * x_mask863        if rate:864            head = int(z_p.shape[2] * rate)865            z_p = z_p[:, :, -head:]866            x_mask = x_mask[:, :, -head:]867        z = self.flow(z_p, x_mask, g=g, reverse=True)868        o = self.dec(z * x_mask, g=g)869        return o, x_mask, (z, z_p, m_p, logs_p)870 871 872class SynthesizerTrnMs768NSFsid_nono(nn.Module):873    def __init__(874        self,875        spec_channels,876        segment_size,877        inter_channels,878        hidden_channels,879        filter_channels,880        n_heads,881        n_layers,882        kernel_size,883        p_dropout,884        resblock,885        resblock_kernel_sizes,886        resblock_dilation_sizes,887        upsample_rates,888        upsample_initial_channel,889        upsample_kernel_sizes,890        spk_embed_dim,891        gin_channels,892        sr=None,893        **kwargs894    ):895        super().__init__()896        self.spec_channels = spec_channels897        self.inter_channels = inter_channels898        self.hidden_channels = hidden_channels899        self.filter_channels = filter_channels900        self.n_heads = n_heads901        self.n_layers = n_layers902        self.kernel_size = kernel_size903        self.p_dropout = p_dropout904        self.resblock = resblock905        self.resblock_kernel_sizes = resblock_kernel_sizes906        self.resblock_dilation_sizes = resblock_dilation_sizes907        self.upsample_rates = upsample_rates908        self.upsample_initial_channel = upsample_initial_channel909        self.upsample_kernel_sizes = upsample_kernel_sizes910        self.segment_size = segment_size911        self.gin_channels = gin_channels912        # self.hop_length = hop_length#913        self.spk_embed_dim = spk_embed_dim914        self.enc_p = TextEncoder768(915            inter_channels,916            hidden_channels,917            filter_channels,918            n_heads,919            n_layers,920            kernel_size,921            p_dropout,922            f0=False,923        )924        self.dec = Generator(925            inter_channels,926            resblock,927            resblock_kernel_sizes,928            resblock_dilation_sizes,929            upsample_rates,930            upsample_initial_channel,931            upsample_kernel_sizes,932            gin_channels=gin_channels,933        )934        self.enc_q = PosteriorEncoder(935            spec_channels,936            inter_channels,937            hidden_channels,938            5,939            1,940            16,941            gin_channels=gin_channels,942        )943        self.flow = ResidualCouplingBlock(944            inter_channels, hidden_channels, 5, 1, 3, gin_channels=gin_channels945        )946        self.emb_g = nn.Embedding(self.spk_embed_dim, gin_channels)947        print("gin_channels:", gin_channels, "self.spk_embed_dim:", self.spk_embed_dim)948 949    def remove_weight_norm(self):950        self.dec.remove_weight_norm()951        self.flow.remove_weight_norm()952        self.enc_q.remove_weight_norm()953 954    def forward(self, phone, phone_lengths, y, y_lengths, ds):  # 这里ds是id,[bs,1]955        g = self.emb_g(ds).unsqueeze(-1)  # [b, 256, 1]##1是t,广播的956        m_p, logs_p, x_mask = self.enc_p(phone, None, phone_lengths)957        z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g)958        z_p = self.flow(z, y_mask, g=g)959        z_slice, ids_slice = commons.rand_slice_segments(960            z, y_lengths, self.segment_size961        )962        o = self.dec(z_slice, g=g)963        return o, ids_slice, x_mask, y_mask, (z, z_p, m_p, logs_p, m_q, logs_q)964 965    def infer(self, phone, phone_lengths, sid, rate=None):966        g = self.emb_g(sid).unsqueeze(-1)967        m_p, logs_p, x_mask = self.enc_p(phone, None, phone_lengths)968        z_p = (m_p + torch.exp(logs_p) * torch.randn_like(m_p) * 0.66666) * x_mask969        if rate:970            head = int(z_p.shape[2] * rate)971            z_p = z_p[:, :, -head:]972            x_mask = x_mask[:, :, -head:]973        z = self.flow(z_p, x_mask, g=g, reverse=True)974        o = self.dec(z * x_mask, g=g)975        return o, x_mask, (z, z_p, m_p, logs_p)976 977 978class MultiPeriodDiscriminator(torch.nn.Module):979    def __init__(self, use_spectral_norm=False):980        super(MultiPeriodDiscriminator, self).__init__()981        periods = [2, 3, 5, 7, 11, 17]982        # periods = [3, 5, 7, 11, 17, 23, 37]983 984        discs = [DiscriminatorS(use_spectral_norm=use_spectral_norm)]985        discs = discs + [986            DiscriminatorP(i, use_spectral_norm=use_spectral_norm) for i in periods987        ]988        self.discriminators = nn.ModuleList(discs)989 990    def forward(self, y, y_hat):991        y_d_rs = []  #992        y_d_gs = []993        fmap_rs = []994        fmap_gs = []995        for i, d in enumerate(self.discriminators):996            y_d_r, fmap_r = d(y)997            y_d_g, fmap_g = d(y_hat)998            # for j in range(len(fmap_r)):999            #     print(i,j,y.shape,y_hat.shape,fmap_r[j].shape,fmap_g[j].shape)1000            y_d_rs.append(y_d_r)1001            y_d_gs.append(y_d_g)1002            fmap_rs.append(fmap_r)1003            fmap_gs.append(fmap_g)1004 1005        return y_d_rs, y_d_gs, fmap_rs, fmap_gs1006 1007 1008class MultiPeriodDiscriminatorV2(torch.nn.Module):1009    def __init__(self, use_spectral_norm=False):1010        super(MultiPeriodDiscriminatorV2, self).__init__()1011        # periods = [2, 3, 5, 7, 11, 17]1012        periods = [2, 3, 5, 7, 11, 17, 23, 37]1013 1014        discs = [DiscriminatorS(use_spectral_norm=use_spectral_norm)]1015        discs = discs + [1016            DiscriminatorP(i, use_spectral_norm=use_spectral_norm) for i in periods1017        ]1018        self.discriminators = nn.ModuleList(discs)1019 1020    def forward(self, y, y_hat):1021        y_d_rs = []  #1022        y_d_gs = []1023        fmap_rs = []1024        fmap_gs = []1025        for i, d in enumerate(self.discriminators):1026            y_d_r, fmap_r = d(y)1027            y_d_g, fmap_g = d(y_hat)1028            # for j in range(len(fmap_r)):1029            #     print(i,j,y.shape,y_hat.shape,fmap_r[j].shape,fmap_g[j].shape)1030            y_d_rs.append(y_d_r)1031            y_d_gs.append(y_d_g)1032            fmap_rs.append(fmap_r)1033            fmap_gs.append(fmap_g)1034 1035        return y_d_rs, y_d_gs, fmap_rs, fmap_gs1036 1037 1038class DiscriminatorS(torch.nn.Module):1039    def __init__(self, use_spectral_norm=False):1040        super(DiscriminatorS, self).__init__()1041        norm_f = weight_norm if use_spectral_norm == False else spectral_norm1042        self.convs = nn.ModuleList(1043            [1044                norm_f(Conv1d(1, 16, 15, 1, padding=7)),1045                norm_f(Conv1d(16, 64, 41, 4, groups=4, padding=20)),1046                norm_f(Conv1d(64, 256, 41, 4, groups=16, padding=20)),1047                norm_f(Conv1d(256, 1024, 41, 4, groups=64, padding=20)),1048                norm_f(Conv1d(1024, 1024, 41, 4, groups=256, padding=20)),1049                norm_f(Conv1d(1024, 1024, 5, 1, padding=2)),1050            ]1051        )1052        self.conv_post = norm_f(Conv1d(1024, 1, 3, 1, padding=1))1053 1054    def forward(self, x):1055        fmap = []1056 1057        for l in self.convs:1058            x = l(x)1059            x = F.leaky_relu(x, modules.LRELU_SLOPE)1060            fmap.append(x)1061        x = self.conv_post(x)1062        fmap.append(x)1063        x = torch.flatten(x, 1, -1)1064 1065        return x, fmap1066 1067 1068class DiscriminatorP(torch.nn.Module):1069    def __init__(self, period, kernel_size=5, stride=3, use_spectral_norm=False):1070        super(DiscriminatorP, self).__init__()1071        self.period = period1072        self.use_spectral_norm = use_spectral_norm1073        norm_f = weight_norm if use_spectral_norm == False else spectral_norm1074        self.convs = nn.ModuleList(1075            [1076                norm_f(1077                    Conv2d(1078                        1,1079                        32,1080                        (kernel_size, 1),1081                        (stride, 1),1082                        padding=(get_padding(kernel_size, 1), 0),1083                    )1084                ),1085                norm_f(1086                    Conv2d(1087                        32,1088                        128,1089                        (kernel_size, 1),1090                        (stride, 1),1091                        padding=(get_padding(kernel_size, 1), 0),1092                    )1093                ),1094                norm_f(1095                    Conv2d(1096                        128,1097                        512,1098                        (kernel_size, 1),1099                        (stride, 1),1100                        padding=(get_padding(kernel_size, 1), 0),1101                    )1102                ),1103                norm_f(1104                    Conv2d(1105                        512,1106                        1024,1107                        (kernel_size, 1),1108                        (stride, 1),1109                        padding=(get_padding(kernel_size, 1), 0),1110                    )1111                ),1112                norm_f(1113                    Conv2d(1114                        1024,1115                        1024,1116                        (kernel_size, 1),1117                        1,1118                        padding=(get_padding(kernel_size, 1), 0),1119                    )1120                ),1121            ]1122        )1123        self.conv_post = norm_f(Conv2d(1024, 1, (3, 1), 1, padding=(1, 0)))1124 1125    def forward(self, x):1126        fmap = []1127 1128        # 1d to 2d1129        b, c, t = x.shape1130        if t % self.period != 0:  # pad first1131            n_pad = self.period - (t % self.period)1132            x = F.pad(x, (0, n_pad), "reflect")1133            t = t + n_pad1134        x = x.view(b, c, t // self.period, self.period)1135 1136        for l in self.convs:1137            x = l(x)1138            x = F.leaky_relu(x, modules.LRELU_SLOPE)1139            fmap.append(x)1140        x = self.conv_post(x)1141        fmap.append(x)1142        x = torch.flatten(x, 1, -1)1143 1144        return x, fmap