CoolFace
Apppublic

Kit-Lemonfoot/vtuber_rvc_models

sourceHugging Facecreativeml-openrail-mupdated 2y agoView on Hugging Face
53likes
models.py1126 linesDownload Raw Back to infer_pack
1import math, pdb, os
2from time import time as ttime
3import torch
4from torch import nn
5from torch.nn import functional as F
6from lib.infer_pack import modules
7from lib.infer_pack import attentions
8from lib.infer_pack import commons
9from lib.infer_pack.commons import init_weights, get_padding
10from torch.nn import Conv1d, ConvTranspose1d, AvgPool1d, Conv2d
11from torch.nn.utils import remove_weight_norm
12from torch.nn.utils.parametrizations import spectral_norm, weight_norm
13from lib.infer_pack.commons import init_weights
14import numpy as np
15from lib.infer_pack import commons
16
17
18class TextEncoder256(nn.Module):
19    def __init__(
20        self,
21        out_channels,
22        hidden_channels,
23        filter_channels,
24        n_heads,
25        n_layers,
26        kernel_size,
27        p_dropout,
28        f0=True,
29    ):
30        super().__init__()
31        self.out_channels = out_channels
32        self.hidden_channels = hidden_channels
33        self.filter_channels = filter_channels
34        self.n_heads = n_heads
35        self.n_layers = n_layers
36        self.kernel_size = kernel_size
37        self.p_dropout = p_dropout
38        self.emb_phone = nn.Linear(256, hidden_channels)
39        self.lrelu = nn.LeakyReLU(0.1, inplace=True)
40        if f0 == True:
41            self.emb_pitch = nn.Embedding(256, hidden_channels)  # pitch 256
42        self.encoder = attentions.Encoder(
43            hidden_channels, filter_channels, n_heads, n_layers, kernel_size, p_dropout
44        )
45        self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)
46
47    def forward(self, phone, pitch, lengths):
48        if pitch == None:
49            x = self.emb_phone(phone)
50        else:
51            x = self.emb_phone(phone) + self.emb_pitch(pitch)
52        x = x * math.sqrt(self.hidden_channels)  # [b, t, h]
53        x = self.lrelu(x)
54        x = torch.transpose(x, 1, -1)  # [b, h, t]
55        x_mask = torch.unsqueeze(commons.sequence_mask(lengths, x.size(2)), 1).to(
56            x.dtype
57        )
58        x = self.encoder(x * x_mask, x_mask)
59        stats = self.proj(x) * x_mask
60
61        m, logs = torch.split(stats, self.out_channels, dim=1)
62        return m, logs, x_mask
63
64
65class TextEncoder768(nn.Module):
66    def __init__(
67        self,
68        out_channels,
69        hidden_channels,
70        filter_channels,
71        n_heads,
72        n_layers,
73        kernel_size,
74        p_dropout,
75        f0=True,
76    ):
77        super().__init__()
78        self.out_channels = out_channels
79        self.hidden_channels = hidden_channels
80        self.filter_channels = filter_channels
81        self.n_heads = n_heads
82        self.n_layers = n_layers
83        self.kernel_size = kernel_size
84        self.p_dropout = p_dropout
85        self.emb_phone = nn.Linear(768, hidden_channels)
86        self.lrelu = nn.LeakyReLU(0.1, inplace=True)
87        if f0 == True:
88            self.emb_pitch = nn.Embedding(256, hidden_channels)  # pitch 256
89        self.encoder = attentions.Encoder(
90            hidden_channels, filter_channels, n_heads, n_layers, kernel_size, p_dropout
91        )
92        self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)
93
94    def forward(self, phone, pitch, lengths):
95        if pitch == None:
96            x = self.emb_phone(phone)
97        else:
98            x = self.emb_phone(phone) + self.emb_pitch(pitch)
99        x = x * math.sqrt(self.hidden_channels)  # [b, t, h]
100        x = self.lrelu(x)
101        x = torch.transpose(x, 1, -1)  # [b, h, t]
102        x_mask = torch.unsqueeze(commons.sequence_mask(lengths, x.size(2)), 1).to(
103            x.dtype
104        )
105        x = self.encoder(x * x_mask, x_mask)
106        stats = self.proj(x) * x_mask
107
108        m, logs = torch.split(stats, self.out_channels, dim=1)
109        return m, logs, x_mask
110
111
112class ResidualCouplingBlock(nn.Module):
113    def __init__(
114        self,
115        channels,
116        hidden_channels,
117        kernel_size,
118        dilation_rate,
119        n_layers,
120        n_flows=4,
121        gin_channels=0,
122    ):
123        super().__init__()
124        self.channels = channels
125        self.hidden_channels = hidden_channels
126        self.kernel_size = kernel_size
127        self.dilation_rate = dilation_rate
128        self.n_layers = n_layers
129        self.n_flows = n_flows
130        self.gin_channels = gin_channels
131
132        self.flows = nn.ModuleList()
133        for i in range(n_flows):
134            self.flows.append(
135                modules.ResidualCouplingLayer(
136                    channels,
137                    hidden_channels,
138                    kernel_size,
139                    dilation_rate,
140                    n_layers,
141                    gin_channels=gin_channels,
142                    mean_only=True,
143                )
144            )
145            self.flows.append(modules.Flip())
146
147    def forward(self, x, x_mask, g=None, reverse=False):
148        if not reverse:
149            for flow in self.flows:
150                x, _ = flow(x, x_mask, g=g, reverse=reverse)
151        else:
152            for flow in reversed(self.flows):
153                x = flow(x, x_mask, g=g, reverse=reverse)
154        return x
155
156    def remove_weight_norm(self):
157        for i in range(self.n_flows):
158            self.flows[i * 2].remove_weight_norm()
159
160
161class PosteriorEncoder(nn.Module):
162    def __init__(
163        self,
164        in_channels,
165        out_channels,
166        hidden_channels,
167        kernel_size,
168        dilation_rate,
169        n_layers,
170        gin_channels=0,
171    ):
172        super().__init__()
173        self.in_channels = in_channels
174        self.out_channels = out_channels
175        self.hidden_channels = hidden_channels
176        self.kernel_size = kernel_size
177        self.dilation_rate = dilation_rate
178        self.n_layers = n_layers
179        self.gin_channels = gin_channels
180
181        self.pre = nn.Conv1d(in_channels, hidden_channels, 1)
182        self.enc = modules.WN(
183            hidden_channels,
184            kernel_size,
185            dilation_rate,
186            n_layers,
187            gin_channels=gin_channels,
188        )
189        self.proj = nn.Conv1d(hidden_channels, out_channels * 2, 1)
190
191    def forward(self, x, x_lengths, g=None):
192        x_mask = torch.unsqueeze(commons.sequence_mask(x_lengths, x.size(2)), 1).to(
193            x.dtype
194        )
195        x = self.pre(x) * x_mask
196        x = self.enc(x, x_mask, g=g)
197        stats = self.proj(x) * x_mask
198        m, logs = torch.split(stats, self.out_channels, dim=1)
199        z = (m + torch.randn_like(m) * torch.exp(logs)) * x_mask
200        return z, m, logs, x_mask
201
202    def remove_weight_norm(self):
203        self.enc.remove_weight_norm()
204
205
206class Generator(torch.nn.Module):
207    def __init__(
208        self,
209        initial_channel,
210        resblock,
211        resblock_kernel_sizes,
212        resblock_dilation_sizes,
213        upsample_rates,
214        upsample_initial_channel,
215        upsample_kernel_sizes,
216        gin_channels=0,
217    ):
218        super(Generator, self).__init__()
219        self.num_kernels = len(resblock_kernel_sizes)
220        self.num_upsamples = len(upsample_rates)
221        self.conv_pre = Conv1d(
222            initial_channel, upsample_initial_channel, 7, 1, padding=3
223        )
224        resblock = modules.ResBlock1 if resblock == "1" else modules.ResBlock2
225
226        self.ups = nn.ModuleList()
227        for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)):
228            self.ups.append(
229                weight_norm(
230                    ConvTranspose1d(
231                        upsample_initial_channel // (2**i),
232                        upsample_initial_channel // (2 ** (i + 1)),
233                        k,
234                        u,
235                        padding=(k - u) // 2,
236                    )
237                )
238            )
239
240        self.resblocks = nn.ModuleList()
241        for i in range(len(self.ups)):
242            ch = upsample_initial_channel // (2 ** (i + 1))
243            for j, (k, d) in enumerate(
244                zip(resblock_kernel_sizes, resblock_dilation_sizes)
245            ):
246                self.resblocks.append(resblock(ch, k, d))
247
248        self.conv_post = Conv1d(ch, 1, 7, 1, padding=3, bias=False)
249        self.ups.apply(init_weights)
250
251        if gin_channels != 0:
252            self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1)
253
254    def forward(self, x, g=None):
255        x = self.conv_pre(x)
256        if g is not None:
257            x = x + self.cond(g)
258
259        for i in range(self.num_upsamples):
260            x = F.leaky_relu(x, modules.LRELU_SLOPE)
261            x = self.ups[i](x)
262            xs = None
263            for j in range(self.num_kernels):
264                if xs is None:
265                    xs = self.resblocks[i * self.num_kernels + j](x)
266                else:
267                    xs += self.resblocks[i * self.num_kernels + j](x)
268            x = xs / self.num_kernels
269        x = F.leaky_relu(x)
270        x = self.conv_post(x)
271        x = torch.tanh(x)
272
273        return x
274
275    def remove_weight_norm(self):
276        for l in self.ups:
277            remove_weight_norm(l)
278        for l in self.resblocks:
279            l.remove_weight_norm()
280
281
282class SineGen(torch.nn.Module):
283    """Definition of sine generator
284    SineGen(samp_rate, harmonic_num = 0,
285            sine_amp = 0.1, noise_std = 0.003,
286            voiced_threshold = 0,
287            flag_for_pulse=False)
288    samp_rate: sampling rate in Hz
289    harmonic_num: number of harmonic overtones (default 0)
290    sine_amp: amplitude of sine-wavefrom (default 0.1)
291    noise_std: std of Gaussian noise (default 0.003)
292    voiced_thoreshold: F0 threshold for U/V classification (default 0)
293    flag_for_pulse: this SinGen is used inside PulseGen (default False)
294    Note: when flag_for_pulse is True, the first time step of a voiced
295        segment is always sin(np.pi) or cos(0)
296    """
297
298    def __init__(
299        self,
300        samp_rate,
301        harmonic_num=0,
302        sine_amp=0.1,
303        noise_std=0.003,
304        voiced_threshold=0,
305        flag_for_pulse=False,
306    ):
307        super(SineGen, self).__init__()
308        self.sine_amp = sine_amp
309        self.noise_std = noise_std
310        self.harmonic_num = harmonic_num
311        self.dim = self.harmonic_num + 1
312        self.sampling_rate = samp_rate
313        self.voiced_threshold = voiced_threshold
314
315    def _f02uv(self, f0):
316        # generate uv signal
317        uv = torch.ones_like(f0)
318        uv = uv * (f0 > self.voiced_threshold)
319        return uv
320
321    def forward(self, f0, upp):
322        """sine_tensor, uv = forward(f0)
323        input F0: tensor(batchsize=1, length, dim=1)
324                  f0 for unvoiced steps should be 0
325        output sine_tensor: tensor(batchsize=1, length, dim)
326        output uv: tensor(batchsize=1, length, 1)
327        """
328        with torch.no_grad():
329            f0 = f0[:, None].transpose(1, 2)
330            f0_buf = torch.zeros(f0.shape[0], f0.shape[1], self.dim, device=f0.device)
331            # fundamental component
332            f0_buf[:, :, 0] = f0[:, :, 0]
333            for idx in np.arange(self.harmonic_num):
334                f0_buf[:, :, idx + 1] = f0_buf[:, :, 0] * (
335                    idx + 2
336                )  # idx + 2: the (idx+1)-th overtone, (idx+2)-th harmonic
337            rad_values = (f0_buf / self.sampling_rate) % 1  ###%1意味着n_har的乘积无法后处理优化
338            rand_ini = torch.rand(
339                f0_buf.shape[0], f0_buf.shape[2], device=f0_buf.device
340            )
341            rand_ini[:, 0] = 0
342            rad_values[:, 0, :] = rad_values[:, 0, :] + rand_ini
343            tmp_over_one = torch.cumsum(rad_values, 1)  # % 1  #####%1意味着后面的cumsum无法再优化
344            tmp_over_one *= upp
345            tmp_over_one = F.interpolate(
346                tmp_over_one.transpose(2, 1),
347                scale_factor=upp,
348                mode="linear",
349                align_corners=True,
350            ).transpose(2, 1)
351            rad_values = F.interpolate(
352                rad_values.transpose(2, 1), scale_factor=upp, mode="nearest"
353            ).transpose(
354                2, 1
355            )  #######
356            tmp_over_one %= 1
357            tmp_over_one_idx = (tmp_over_one[:, 1:, :] - tmp_over_one[:, :-1, :]) < 0
358            cumsum_shift = torch.zeros_like(rad_values)
359            cumsum_shift[:, 1:, :] = tmp_over_one_idx * -1.0
360            sine_waves = torch.sin(
361                torch.cumsum(rad_values + cumsum_shift, dim=1) * 2 * np.pi
362            )
363            sine_waves = sine_waves * self.sine_amp
364            uv = self._f02uv(f0)
365            uv = F.interpolate(
366                uv.transpose(2, 1), scale_factor=upp, mode="nearest"
367            ).transpose(2, 1)
368            noise_amp = uv * self.noise_std + (1 - uv) * self.sine_amp / 3
369            noise = noise_amp * torch.randn_like(sine_waves)
370            sine_waves = sine_waves * uv + noise
371        return sine_waves, uv, noise
372
373
374class SourceModuleHnNSF(torch.nn.Module):
375    """SourceModule for hn-nsf
376    SourceModule(sampling_rate, harmonic_num=0, sine_amp=0.1,
377                 add_noise_std=0.003, voiced_threshod=0)
378    sampling_rate: sampling_rate in Hz
379    harmonic_num: number of harmonic above F0 (default: 0)
380    sine_amp: amplitude of sine source signal (default: 0.1)
381    add_noise_std: std of additive Gaussian noise (default: 0.003)
382        note that amplitude of noise in unvoiced is decided
383        by sine_amp
384    voiced_threshold: threhold to set U/V given F0 (default: 0)
385    Sine_source, noise_source = SourceModuleHnNSF(F0_sampled)
386    F0_sampled (batchsize, length, 1)
387    Sine_source (batchsize, length, 1)
388    noise_source (batchsize, length 1)
389    uv (batchsize, length, 1)
390    """
391
392    def __init__(
393        self,
394        sampling_rate,
395        harmonic_num=0,
396        sine_amp=0.1,
397        add_noise_std=0.003,
398        voiced_threshod=0,
399        is_half=True,
400    ):
401        super(SourceModuleHnNSF, self).__init__()
402
403        self.sine_amp = sine_amp
404        self.noise_std = add_noise_std
405        self.is_half = is_half
406        # to produce sine waveforms
407        self.l_sin_gen = SineGen(
408            sampling_rate, harmonic_num, sine_amp, add_noise_std, voiced_threshod
409        )
410
411        # to merge source harmonics into a single excitation
412        self.l_linear = torch.nn.Linear(harmonic_num + 1, 1)
413        self.l_tanh = torch.nn.Tanh()
414
415    def forward(self, x, upp=None):
416        sine_wavs, uv, _ = self.l_sin_gen(x, upp)
417        if self.is_half:
418            sine_wavs = sine_wavs.half()
419        sine_merge = self.l_tanh(self.l_linear(sine_wavs))
420        return sine_merge, None, None  # noise, uv
421
422
423class GeneratorNSF(torch.nn.Module):
424    def __init__(
425        self,
426        initial_channel,
427        resblock,
428        resblock_kernel_sizes,
429        resblock_dilation_sizes,
430        upsample_rates,
431        upsample_initial_channel,
432        upsample_kernel_sizes,
433        gin_channels,
434        sr,
435        is_half=False,
436    ):
437        super(GeneratorNSF, self).__init__()
438        self.num_kernels = len(resblock_kernel_sizes)
439        self.num_upsamples = len(upsample_rates)
440
441        self.f0_upsamp = torch.nn.Upsample(scale_factor=np.prod(upsample_rates))
442        self.m_source = SourceModuleHnNSF(
443            sampling_rate=sr, harmonic_num=0, is_half=is_half
444        )
445        self.noise_convs = nn.ModuleList()
446        self.conv_pre = Conv1d(
447            initial_channel, upsample_initial_channel, 7, 1, padding=3
448        )
449        resblock = modules.ResBlock1 if resblock == "1" else modules.ResBlock2
450
451        self.ups = nn.ModuleList()
452        for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)):
453            c_cur = upsample_initial_channel // (2 ** (i + 1))
454            self.ups.append(
455                weight_norm(
456                    ConvTranspose1d(
457                        upsample_initial_channel // (2**i),
458                        upsample_initial_channel // (2 ** (i + 1)),
459                        k,
460                        u,
461                        padding=(k - u) // 2,
462                    )
463                )
464            )
465            if i + 1 < len(upsample_rates):
466                stride_f0 = np.prod(upsample_rates[i + 1 :])
467                self.noise_convs.append(
468                    Conv1d(
469                        1,
470                        c_cur,
471                        kernel_size=stride_f0 * 2,
472                        stride=stride_f0,
473                        padding=stride_f0 // 2,
474                    )
475                )
476            else:
477                self.noise_convs.append(Conv1d(1, c_cur, kernel_size=1))
478
479        self.resblocks = nn.ModuleList()
480        for i in range(len(self.ups)):
481            ch = upsample_initial_channel // (2 ** (i + 1))
482            for j, (k, d) in enumerate(
483                zip(resblock_kernel_sizes, resblock_dilation_sizes)
484            ):
485                self.resblocks.append(resblock(ch, k, d))
486
487        self.conv_post = Conv1d(ch, 1, 7, 1, padding=3, bias=False)
488        self.ups.apply(init_weights)
489
490        if gin_channels != 0:
491            self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1)
492
493        self.upp = np.prod(upsample_rates)
494
495    def forward(self, x, f0, g=None):
496        har_source, noi_source, uv = self.m_source(f0, self.upp)
497        har_source = har_source.transpose(1, 2)
498        x = self.conv_pre(x)
499        if g is not None:
500            x = x + self.cond(g)
501
502        for i in range(self.num_upsamples):
503            x = F.leaky_relu(x, modules.LRELU_SLOPE)
504            x = self.ups[i](x)
505            x_source = self.noise_convs[i](har_source)
506            x = x + x_source
507            xs = None
508            for j in range(self.num_kernels):
509                if xs is None:
510                    xs = self.resblocks[i * self.num_kernels + j](x)
511                else:
512                    xs += self.resblocks[i * self.num_kernels + j](x)
513            x = xs / self.num_kernels
514        x = F.leaky_relu(x)
515        x = self.conv_post(x)
516        x = torch.tanh(x)
517        return x
518
519    def remove_weight_norm(self):
520        for l in self.ups:
521            remove_weight_norm(l)
522        for l in self.resblocks:
523            l.remove_weight_norm()
524
525
526sr2sr = {
527    "32k": 32000,
528    "40k": 40000,
529    "48k": 48000,
530}
531
532
533class SynthesizerTrnMs256NSFsid(nn.Module):
534    def __init__(
535        self,
536        spec_channels,
537        segment_size,
538        inter_channels,
539        hidden_channels,
540        filter_channels,
541        n_heads,
542        n_layers,
543        kernel_size,
544        p_dropout,
545        resblock,
546        resblock_kernel_sizes,
547        resblock_dilation_sizes,
548        upsample_rates,
549        upsample_initial_channel,
550        upsample_kernel_sizes,
551        spk_embed_dim,
552        gin_channels,
553        sr,
554        **kwargs
555    ):
556        super().__init__()
557        if type(sr) == type("strr"):
558            sr = sr2sr[sr]
559        self.spec_channels = spec_channels
560        self.inter_channels = inter_channels
561        self.hidden_channels = hidden_channels
562        self.filter_channels = filter_channels
563        self.n_heads = n_heads
564        self.n_layers = n_layers
565        self.kernel_size = kernel_size
566        self.p_dropout = p_dropout
567        self.resblock = resblock
568        self.resblock_kernel_sizes = resblock_kernel_sizes
569        self.resblock_dilation_sizes = resblock_dilation_sizes
570        self.upsample_rates = upsample_rates
571        self.upsample_initial_channel = upsample_initial_channel
572        self.upsample_kernel_sizes = upsample_kernel_sizes
573        self.segment_size = segment_size
574        self.gin_channels = gin_channels
575        # self.hop_length = hop_length#
576        self.spk_embed_dim = spk_embed_dim
577        self.enc_p = TextEncoder256(
578            inter_channels,
579            hidden_channels,
580            filter_channels,
581            n_heads,
582            n_layers,
583            kernel_size,
584            p_dropout,
585        )
586        self.dec = GeneratorNSF(
587            inter_channels,
588            resblock,
589            resblock_kernel_sizes,
590            resblock_dilation_sizes,
591            upsample_rates,
592            upsample_initial_channel,
593            upsample_kernel_sizes,
594            gin_channels=gin_channels,
595            sr=sr,
596            is_half=kwargs["is_half"],
597        )
598        self.enc_q = PosteriorEncoder(
599            spec_channels,
600            inter_channels,
601            hidden_channels,
602            5,
603            1,
604            16,
605            gin_channels=gin_channels,
606        )
607        self.flow = ResidualCouplingBlock(
608            inter_channels, hidden_channels, 5, 1, 3, gin_channels=gin_channels
609        )
610        self.emb_g = nn.Embedding(self.spk_embed_dim, gin_channels)
611        print("gin_channels:", gin_channels, "self.spk_embed_dim:", self.spk_embed_dim)
612
613    def remove_weight_norm(self):
614        self.dec.remove_weight_norm()
615        self.flow.remove_weight_norm()
616        self.enc_q.remove_weight_norm()
617
618    def forward(
619        self, phone, phone_lengths, pitch, pitchf, y, y_lengths, ds
620    ):  # 这里ds是id,[bs,1]
621        # print(1,pitch.shape)#[bs,t]
622        g = self.emb_g(ds).unsqueeze(-1)  # [b, 256, 1]##1是t,广播的
623        m_p, logs_p, x_mask = self.enc_p(phone, pitch, phone_lengths)
624        z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g)
625        z_p = self.flow(z, y_mask, g=g)
626        z_slice, ids_slice = commons.rand_slice_segments(
627            z, y_lengths, self.segment_size
628        )
629        # print(-1,pitchf.shape,ids_slice,self.segment_size,self.hop_length,self.segment_size//self.hop_length)
630        pitchf = commons.slice_segments2(pitchf, ids_slice, self.segment_size)
631        # print(-2,pitchf.shape,z_slice.shape)
632        o = self.dec(z_slice, pitchf, g=g)
633        return o, ids_slice, x_mask, y_mask, (z, z_p, m_p, logs_p, m_q, logs_q)
634
635    def infer(self, phone, phone_lengths, pitch, nsff0, sid, max_len=None):
636        g = self.emb_g(sid).unsqueeze(-1)
637        m_p, logs_p, x_mask = self.enc_p(phone, pitch, phone_lengths)
638        z_p = (m_p + torch.exp(logs_p) * torch.randn_like(m_p) * 0.66666) * x_mask
639        z = self.flow(z_p, x_mask, g=g, reverse=True)
640        o = self.dec((z * x_mask)[:, :, :max_len], nsff0, g=g)
641        return o, x_mask, (z, z_p, m_p, logs_p)
642
643
644class SynthesizerTrnMs768NSFsid(nn.Module):
645    def __init__(
646        self,
647        spec_channels,
648        segment_size,
649        inter_channels,
650        hidden_channels,
651        filter_channels,
652        n_heads,
653        n_layers,
654        kernel_size,
655        p_dropout,
656        resblock,
657        resblock_kernel_sizes,
658        resblock_dilation_sizes,
659        upsample_rates,
660        upsample_initial_channel,
661        upsample_kernel_sizes,
662        spk_embed_dim,
663        gin_channels,
664        sr,
665        **kwargs
666    ):
667        super().__init__()
668        if type(sr) == type("strr"):
669            sr = sr2sr[sr]
670        self.spec_channels = spec_channels
671        self.inter_channels = inter_channels
672        self.hidden_channels = hidden_channels
673        self.filter_channels = filter_channels
674        self.n_heads = n_heads
675        self.n_layers = n_layers
676        self.kernel_size = kernel_size
677        self.p_dropout = p_dropout
678        self.resblock = resblock
679        self.resblock_kernel_sizes = resblock_kernel_sizes
680        self.resblock_dilation_sizes = resblock_dilation_sizes
681        self.upsample_rates = upsample_rates
682        self.upsample_initial_channel = upsample_initial_channel
683        self.upsample_kernel_sizes = upsample_kernel_sizes
684        self.segment_size = segment_size
685        self.gin_channels = gin_channels
686        # self.hop_length = hop_length#
687        self.spk_embed_dim = spk_embed_dim
688        self.enc_p = TextEncoder768(
689            inter_channels,
690            hidden_channels,
691            filter_channels,
692            n_heads,
693            n_layers,
694            kernel_size,
695            p_dropout,
696        )
697        self.dec = GeneratorNSF(
698            inter_channels,
699            resblock,
700            resblock_kernel_sizes,
701            resblock_dilation_sizes,
702            upsample_rates,
703            upsample_initial_channel,
704            upsample_kernel_sizes,
705            gin_channels=gin_channels,
706            sr=sr,
707            is_half=kwargs["is_half"],
708        )
709        self.enc_q = PosteriorEncoder(
710            spec_channels,
711            inter_channels,
712            hidden_channels,
713            5,
714            1,
715            16,
716            gin_channels=gin_channels,
717        )
718        self.flow = ResidualCouplingBlock(
719            inter_channels, hidden_channels, 5, 1, 3, gin_channels=gin_channels
720        )
721        self.emb_g = nn.Embedding(self.spk_embed_dim, gin_channels)
722        print("gin_channels:", gin_channels, "self.spk_embed_dim:", self.spk_embed_dim)
723
724    def remove_weight_norm(self):
725        self.dec.remove_weight_norm()
726        self.flow.remove_weight_norm()
727        self.enc_q.remove_weight_norm()
728
729    def forward(
730        self, phone, phone_lengths, pitch, pitchf, y, y_lengths, ds
731    ):  # 这里ds是id,[bs,1]
732        # print(1,pitch.shape)#[bs,t]
733        g = self.emb_g(ds).unsqueeze(-1)  # [b, 256, 1]##1是t,广播的
734        m_p, logs_p, x_mask = self.enc_p(phone, pitch, phone_lengths)
735        z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g)
736        z_p = self.flow(z, y_mask, g=g)
737        z_slice, ids_slice = commons.rand_slice_segments(
738            z, y_lengths, self.segment_size
739        )
740        # print(-1,pitchf.shape,ids_slice,self.segment_size,self.hop_length,self.segment_size//self.hop_length)
741        pitchf = commons.slice_segments2(pitchf, ids_slice, self.segment_size)
742        # print(-2,pitchf.shape,z_slice.shape)
743        o = self.dec(z_slice, pitchf, g=g)
744        return o, ids_slice, x_mask, y_mask, (z, z_p, m_p, logs_p, m_q, logs_q)
745
746    def infer(self, phone, phone_lengths, pitch, nsff0, sid, max_len=None):
747        g = self.emb_g(sid).unsqueeze(-1)
748        m_p, logs_p, x_mask = self.enc_p(phone, pitch, phone_lengths)
749        z_p = (m_p + torch.exp(logs_p) * torch.randn_like(m_p) * 0.66666) * x_mask
750        z = self.flow(z_p, x_mask, g=g, reverse=True)
751        o = self.dec((z * x_mask)[:, :, :max_len], nsff0, g=g)
752        return o, x_mask, (z, z_p, m_p, logs_p)
753
754
755class SynthesizerTrnMs256NSFsid_nono(nn.Module):
756    def __init__(
757        self,
758        spec_channels,
759        segment_size,
760        inter_channels,
761        hidden_channels,
762        filter_channels,
763        n_heads,
764        n_layers,
765        kernel_size,
766        p_dropout,
767        resblock,
768        resblock_kernel_sizes,
769        resblock_dilation_sizes,
770        upsample_rates,
771        upsample_initial_channel,
772        upsample_kernel_sizes,
773        spk_embed_dim,
774        gin_channels,
775        sr=None,
776        **kwargs
777    ):
778        super().__init__()
779        self.spec_channels = spec_channels
780        self.inter_channels = inter_channels
781        self.hidden_channels = hidden_channels
782        self.filter_channels = filter_channels
783        self.n_heads = n_heads
784        self.n_layers = n_layers
785        self.kernel_size = kernel_size
786        self.p_dropout = p_dropout
787        self.resblock = resblock
788        self.resblock_kernel_sizes = resblock_kernel_sizes
789        self.resblock_dilation_sizes = resblock_dilation_sizes
790        self.upsample_rates = upsample_rates
791        self.upsample_initial_channel = upsample_initial_channel
792        self.upsample_kernel_sizes = upsample_kernel_sizes
793        self.segment_size = segment_size
794        self.gin_channels = gin_channels
795        # self.hop_length = hop_length#
796        self.spk_embed_dim = spk_embed_dim
797        self.enc_p = TextEncoder256(
798            inter_channels,
799            hidden_channels,
800            filter_channels,
801            n_heads,
802            n_layers,
803            kernel_size,
804            p_dropout,
805            f0=False,
806        )
807        self.dec = Generator(
808            inter_channels,
809            resblock,
810            resblock_kernel_sizes,
811            resblock_dilation_sizes,
812            upsample_rates,
813            upsample_initial_channel,
814            upsample_kernel_sizes,
815            gin_channels=gin_channels,
816        )
817        self.enc_q = PosteriorEncoder(
818            spec_channels,
819            inter_channels,
820            hidden_channels,
821            5,
822            1,
823            16,
824            gin_channels=gin_channels,
825        )
826        self.flow = ResidualCouplingBlock(
827            inter_channels, hidden_channels, 5, 1, 3, gin_channels=gin_channels
828        )
829        self.emb_g = nn.Embedding(self.spk_embed_dim, gin_channels)
830        print("gin_channels:", gin_channels, "self.spk_embed_dim:", self.spk_embed_dim)
831
832    def remove_weight_norm(self):
833        self.dec.remove_weight_norm()
834        self.flow.remove_weight_norm()
835        self.enc_q.remove_weight_norm()
836
837    def forward(self, phone, phone_lengths, y, y_lengths, ds):  # 这里ds是id,[bs,1]
838        g = self.emb_g(ds).unsqueeze(-1)  # [b, 256, 1]##1是t,广播的
839        m_p, logs_p, x_mask = self.enc_p(phone, None, phone_lengths)
840        z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g)
841        z_p = self.flow(z, y_mask, g=g)
842        z_slice, ids_slice = commons.rand_slice_segments(
843            z, y_lengths, self.segment_size
844        )
845        o = self.dec(z_slice, g=g)
846        return o, ids_slice, x_mask, y_mask, (z, z_p, m_p, logs_p, m_q, logs_q)
847
848    def infer(self, phone, phone_lengths, sid, max_len=None):
849        g = self.emb_g(sid).unsqueeze(-1)
850        m_p, logs_p, x_mask = self.enc_p(phone, None, phone_lengths)
851        z_p = (m_p + torch.exp(logs_p) * torch.randn_like(m_p) * 0.66666) * x_mask
852        z = self.flow(z_p, x_mask, g=g, reverse=True)
853        o = self.dec((z * x_mask)[:, :, :max_len], g=g)
854        return o, x_mask, (z, z_p, m_p, logs_p)
855
856
857class SynthesizerTrnMs768NSFsid_nono(nn.Module):
858    def __init__(
859        self,
860        spec_channels,
861        segment_size,
862        inter_channels,
863        hidden_channels,
864        filter_channels,
865        n_heads,
866        n_layers,
867        kernel_size,
868        p_dropout,
869        resblock,
870        resblock_kernel_sizes,
871        resblock_dilation_sizes,
872        upsample_rates,
873        upsample_initial_channel,
874        upsample_kernel_sizes,
875        spk_embed_dim,
876        gin_channels,
877        sr=None,
878        **kwargs
879    ):
880        super().__init__()
881        self.spec_channels = spec_channels
882        self.inter_channels = inter_channels
883        self.hidden_channels = hidden_channels
884        self.filter_channels = filter_channels
885        self.n_heads = n_heads
886        self.n_layers = n_layers
887        self.kernel_size = kernel_size
888        self.p_dropout = p_dropout
889        self.resblock = resblock
890        self.resblock_kernel_sizes = resblock_kernel_sizes
891        self.resblock_dilation_sizes = resblock_dilation_sizes
892        self.upsample_rates = upsample_rates
893        self.upsample_initial_channel = upsample_initial_channel
894        self.upsample_kernel_sizes = upsample_kernel_sizes
895        self.segment_size = segment_size
896        self.gin_channels = gin_channels
897        # self.hop_length = hop_length#
898        self.spk_embed_dim = spk_embed_dim
899        self.enc_p = TextEncoder768(
900            inter_channels,
901            hidden_channels,
902            filter_channels,
903            n_heads,
904            n_layers,
905            kernel_size,
906            p_dropout,
907            f0=False,
908        )
909        self.dec = Generator(
910            inter_channels,
911            resblock,
912            resblock_kernel_sizes,
913            resblock_dilation_sizes,
914            upsample_rates,
915            upsample_initial_channel,
916            upsample_kernel_sizes,
917            gin_channels=gin_channels,
918        )
919        self.enc_q = PosteriorEncoder(
920            spec_channels,
921            inter_channels,
922            hidden_channels,
923            5,
924            1,
925            16,
926            gin_channels=gin_channels,
927        )
928        self.flow = ResidualCouplingBlock(
929            inter_channels, hidden_channels, 5, 1, 3, gin_channels=gin_channels
930        )
931        self.emb_g = nn.Embedding(self.spk_embed_dim, gin_channels)
932        print("gin_channels:", gin_channels, "self.spk_embed_dim:", self.spk_embed_dim)
933
934    def remove_weight_norm(self):
935        self.dec.remove_weight_norm()
936        self.flow.remove_weight_norm()
937        self.enc_q.remove_weight_norm()
938
939    def forward(self, phone, phone_lengths, y, y_lengths, ds):  # 这里ds是id,[bs,1]
940        g = self.emb_g(ds).unsqueeze(-1)  # [b, 256, 1]##1是t,广播的
941        m_p, logs_p, x_mask = self.enc_p(phone, None, phone_lengths)
942        z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g)
943        z_p = self.flow(z, y_mask, g=g)
944        z_slice, ids_slice = commons.rand_slice_segments(
945            z, y_lengths, self.segment_size
946        )
947        o = self.dec(z_slice, g=g)
948        return o, ids_slice, x_mask, y_mask, (z, z_p, m_p, logs_p, m_q, logs_q)
949
950    def infer(self, phone, phone_lengths, sid, max_len=None):
951        g = self.emb_g(sid).unsqueeze(-1)
952        m_p, logs_p, x_mask = self.enc_p(phone, None, phone_lengths)
953        z_p = (m_p + torch.exp(logs_p) * torch.randn_like(m_p) * 0.66666) * x_mask
954        z = self.flow(z_p, x_mask, g=g, reverse=True)
955        o = self.dec((z * x_mask)[:, :, :max_len], g=g)
956        return o, x_mask, (z, z_p, m_p, logs_p)
957
958
959class MultiPeriodDiscriminator(torch.nn.Module):
960    def __init__(self, use_spectral_norm=False):
961        super(MultiPeriodDiscriminator, self).__init__()
962        periods = [2, 3, 5, 7, 11, 17]
963        # periods = [3, 5, 7, 11, 17, 23, 37]
964
965        discs = [DiscriminatorS(use_spectral_norm=use_spectral_norm)]
966        discs = discs + [
967            DiscriminatorP(i, use_spectral_norm=use_spectral_norm) for i in periods
968        ]
969        self.discriminators = nn.ModuleList(discs)
970
971    def forward(self, y, y_hat):
972        y_d_rs = []  #
973        y_d_gs = []
974        fmap_rs = []
975        fmap_gs = []
976        for i, d in enumerate(self.discriminators):
977            y_d_r, fmap_r = d(y)
978            y_d_g, fmap_g = d(y_hat)
979            # for j in range(len(fmap_r)):
980            #     print(i,j,y.shape,y_hat.shape,fmap_r[j].shape,fmap_g[j].shape)
981            y_d_rs.append(y_d_r)
982            y_d_gs.append(y_d_g)
983            fmap_rs.append(fmap_r)
984            fmap_gs.append(fmap_g)
985
986        return y_d_rs, y_d_gs, fmap_rs, fmap_gs
987
988
989class MultiPeriodDiscriminatorV2(torch.nn.Module):
990    def __init__(self, use_spectral_norm=False):
991        super(MultiPeriodDiscriminatorV2, self).__init__()
992        # periods = [2, 3, 5, 7, 11, 17]
993        periods = [2, 3, 5, 7, 11, 17, 23, 37]
994
995        discs = [DiscriminatorS(use_spectral_norm=use_spectral_norm)]
996        discs = discs + [
997            DiscriminatorP(i, use_spectral_norm=use_spectral_norm) for i in periods
998        ]
999        self.discriminators = nn.ModuleList(discs)
1000
1001    def forward(self, y, y_hat):
1002        y_d_rs = []  #
1003        y_d_gs = []
1004        fmap_rs = []
1005        fmap_gs = []
1006        for i, d in enumerate(self.discriminators):
1007            y_d_r, fmap_r = d(y)
1008            y_d_g, fmap_g = d(y_hat)
1009            # for j in range(len(fmap_r)):
1010            #     print(i,j,y.shape,y_hat.shape,fmap_r[j].shape,fmap_g[j].shape)
1011            y_d_rs.append(y_d_r)
1012            y_d_gs.append(y_d_g)
1013            fmap_rs.append(fmap_r)
1014            fmap_gs.append(fmap_g)
1015
1016        return y_d_rs, y_d_gs, fmap_rs, fmap_gs
1017
1018
1019class DiscriminatorS(torch.nn.Module):
1020    def __init__(self, use_spectral_norm=False):
1021        super(DiscriminatorS, self).__init__()
1022        norm_f = weight_norm if use_spectral_norm == False else spectral_norm
1023        self.convs = nn.ModuleList(
1024            [
1025                norm_f(Conv1d(1, 16, 15, 1, padding=7)),
1026                norm_f(Conv1d(16, 64, 41, 4, groups=4, padding=20)),
1027                norm_f(Conv1d(64, 256, 41, 4, groups=16, padding=20)),
1028                norm_f(Conv1d(256, 1024, 41, 4, groups=64, padding=20)),
1029                norm_f(Conv1d(1024, 1024, 41, 4, groups=256, padding=20)),
1030                norm_f(Conv1d(1024, 1024, 5, 1, padding=2)),
1031            ]
1032        )
1033        self.conv_post = norm_f(Conv1d(1024, 1, 3, 1, padding=1))
1034
1035    def forward(self, x):
1036        fmap = []
1037
1038        for l in self.convs:
1039            x = l(x)
1040            x = F.leaky_relu(x, modules.LRELU_SLOPE)
1041            fmap.append(x)
1042        x = self.conv_post(x)
1043        fmap.append(x)
1044        x = torch.flatten(x, 1, -1)
1045
1046        return x, fmap
1047
1048
1049class DiscriminatorP(torch.nn.Module):
1050    def __init__(self, period, kernel_size=5, stride=3, use_spectral_norm=False):
1051        super(DiscriminatorP, self).__init__()
1052        self.period = period
1053        self.use_spectral_norm = use_spectral_norm
1054        norm_f = weight_norm if use_spectral_norm == False else spectral_norm
1055        self.convs = nn.ModuleList(
1056            [
1057                norm_f(
1058                    Conv2d(
1059                        1,
1060                        32,
1061                        (kernel_size, 1),
1062                        (stride, 1),
1063                        padding=(get_padding(kernel_size, 1), 0),
1064                    )
1065                ),
1066                norm_f(
1067                    Conv2d(
1068                        32,
1069                        128,
1070                        (kernel_size, 1),
1071                        (stride, 1),
1072                        padding=(get_padding(kernel_size, 1), 0),
1073                    )
1074                ),
1075                norm_f(
1076                    Conv2d(
1077                        128,
1078                        512,
1079                        (kernel_size, 1),
1080                        (stride, 1),
1081                        padding=(get_padding(kernel_size, 1), 0),
1082                    )
1083                ),
1084                norm_f(
1085                    Conv2d(
1086                        512,
1087                        1024,
1088                        (kernel_size, 1),
1089                        (stride, 1),
1090                        padding=(get_padding(kernel_size, 1), 0),
1091                    )
1092                ),
1093                norm_f(
1094                    Conv2d(
1095                        1024,
1096                        1024,
1097                        (kernel_size, 1),
1098                        1,
1099                        padding=(get_padding(kernel_size, 1), 0),
1100                    )
1101                ),
1102            ]
1103        )
1104        self.conv_post = norm_f(Conv2d(1024, 1, (3, 1), 1, padding=(1, 0)))
1105
1106    def forward(self, x):
1107        fmap = []
1108
1109        # 1d to 2d
1110        b, c, t = x.shape
1111        if t % self.period != 0:  # pad first
1112            n_pad = self.period - (t % self.period)
1113            x = F.pad(x, (0, n_pad), "reflect")
1114            t = t + n_pad
1115        x = x.view(b, c, t // self.period, self.period)
1116
1117        for l in self.convs:
1118            x = l(x)
1119            x = F.leaky_relu(x, modules.LRELU_SLOPE)
1120            fmap.append(x)
1121        x = self.conv_post(x)
1122        fmap.append(x)
1123        x = torch.flatten(x, 1, -1)
1124
1125        return x, fmap
1126