CoolFace
Apppublic

arcanus/koala2

sourceHugging Faceupdated 2y agoView on Hugging Face
0likes
models.py1143 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 weight_norm, remove_weight_norm, spectral_norm
12from lib.infer_pack.commons import init_weights
13import numpy as np
14from lib.infer_pack import commons
15
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_channels
31        self.hidden_channels = hidden_channels
32        self.filter_channels = filter_channels
33        self.n_heads = n_heads
34        self.n_layers = n_layers
35        self.kernel_size = kernel_size
36        self.p_dropout = p_dropout
37        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 256
41        self.encoder = attentions.Encoder(
42            hidden_channels, filter_channels, n_heads, n_layers, kernel_size, p_dropout
43        )
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.dtype
56        )
57        x = self.encoder(x * x_mask, x_mask)
58        stats = self.proj(x) * x_mask
59
60        m, logs = torch.split(stats, self.out_channels, dim=1)
61        return m, logs, x_mask
62
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_channels
78        self.hidden_channels = hidden_channels
79        self.filter_channels = filter_channels
80        self.n_heads = n_heads
81        self.n_layers = n_layers
82        self.kernel_size = kernel_size
83        self.p_dropout = p_dropout
84        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 256
88        self.encoder = attentions.Encoder(
89            hidden_channels, filter_channels, n_heads, n_layers, kernel_size, p_dropout
90        )
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.dtype
103        )
104        x = self.encoder(x * x_mask, x_mask)
105        stats = self.proj(x) * x_mask
106
107        m, logs = torch.split(stats, self.out_channels, dim=1)
108        return m, logs, x_mask
109
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 = channels
124        self.hidden_channels = hidden_channels
125        self.kernel_size = kernel_size
126        self.dilation_rate = dilation_rate
127        self.n_layers = n_layers
128        self.n_flows = n_flows
129        self.gin_channels = gin_channels
130
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 x
154
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_channels
173        self.out_channels = out_channels
174        self.hidden_channels = hidden_channels
175        self.kernel_size = kernel_size
176        self.dilation_rate = dilation_rate
177        self.n_layers = n_layers
178        self.gin_channels = gin_channels
179
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.dtype
193        )
194        x = self.pre(x) * x_mask
195        x = self.enc(x, x_mask, g=g)
196        stats = self.proj(x) * x_mask
197        m, logs = torch.split(stats, self.out_channels, dim=1)
198        z = (m + torch.randn_like(m) * torch.exp(logs)) * x_mask
199        return z, m, logs, x_mask
200
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=3
222        )
223        resblock = modules.ResBlock1 if resblock == "1" else modules.ResBlock2
224
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 = None
262            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_kernels
268        x = F.leaky_relu(x)
269        x = self.conv_post(x)
270        x = torch.tanh(x)
271
272        return x
273
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 generator
283    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 Hz
288    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 voiced
294        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_amp
308        self.noise_std = noise_std
309        self.harmonic_num = harmonic_num
310        self.dim = self.harmonic_num + 1
311        self.sampling_rate = samp_rate
312        self.voiced_threshold = voiced_threshold
313
314    def _f02uv(self, f0):
315        # generate uv signal
316        uv = torch.ones_like(f0)
317        uv = uv * (f0 > self.voiced_threshold)
318        return uv
319
320    def forward(self, f0, upp):
321        """sine_tensor, uv = forward(f0)
322        input F0: tensor(batchsize=1, length, dim=1)
323                  f0 for unvoiced steps should be 0
324        output sine_tensor: tensor(batchsize=1, length, dim)
325        output uv: tensor(batchsize=1, length, 1)
326        """
327        with torch.no_grad():
328            f0 = f0[:, None].transpose(1, 2)
329            f0_buf = torch.zeros(f0.shape[0], f0.shape[1], self.dim, device=f0.device)
330            # fundamental component
331            f0_buf[:, :, 0] = f0[:, :, 0]
332            for idx in np.arange(self.harmonic_num):
333                f0_buf[:, :, idx + 1] = f0_buf[:, :, 0] * (
334                    idx + 2
335                )  # idx + 2: the (idx+1)-th overtone, (idx+2)-th harmonic
336            rad_values = (f0_buf / self.sampling_rate) % 1  ###%1 means that the product of n_har cannot be post-processed and optimized
337            rand_ini = torch.rand(
338                f0_buf.shape[0], f0_buf.shape[2], device=f0_buf.device
339            )
340            rand_ini[:, 0] = 0
341            rad_values[:, 0, :] = rad_values[:, 0, :] + rand_ini
342            tmp_over_one = torch.cumsum(rad_values, 1)  # % 1  #####%1 means that the following cumsum can no longer be optimized
343            tmp_over_one *= upp
344            tmp_over_one = F.interpolate(
345                tmp_over_one.transpose(2, 1),
346                scale_factor=upp,
347                mode="linear",
348                align_corners=True,
349            ).transpose(2, 1)
350            rad_values = F.interpolate(
351                rad_values.transpose(2, 1), scale_factor=upp, mode="nearest"
352            ).transpose(
353                2, 1
354            )  #######
355            tmp_over_one %= 1
356            tmp_over_one_idx = (tmp_over_one[:, 1:, :] - tmp_over_one[:, :-1, :]) < 0
357            cumsum_shift = torch.zeros_like(rad_values)
358            cumsum_shift[:, 1:, :] = tmp_over_one_idx * -1.0
359            sine_waves = torch.sin(
360                torch.cumsum(rad_values + cumsum_shift, dim=1) * 2 * np.pi
361            )
362            sine_waves = sine_waves * self.sine_amp
363            uv = self._f02uv(f0)
364            uv = F.interpolate(
365                uv.transpose(2, 1), scale_factor=upp, mode="nearest"
366            ).transpose(2, 1)
367            noise_amp = uv * self.noise_std + (1 - uv) * self.sine_amp / 3
368            noise = noise_amp * torch.randn_like(sine_waves)
369            sine_waves = sine_waves * uv + noise
370        return sine_waves, uv, noise
371
372
373class SourceModuleHnNSF(torch.nn.Module):
374    """SourceModule for hn-nsf
375    SourceModule(sampling_rate, harmonic_num=0, sine_amp=0.1,
376                 add_noise_std=0.003, voiced_threshod=0)
377    sampling_rate: sampling_rate in Hz
378    harmonic_num: number of harmonic above F0 (default: 0)
379    sine_amp: amplitude of sine source signal (default: 0.1)
380    add_noise_std: std of additive Gaussian noise (default: 0.003)
381        note that amplitude of noise in unvoiced is decided
382        by sine_amp
383    voiced_threshold: threhold to set U/V given F0 (default: 0)
384    Sine_source, noise_source = SourceModuleHnNSF(F0_sampled)
385    F0_sampled (batchsize, length, 1)
386    Sine_source (batchsize, length, 1)
387    noise_source (batchsize, length 1)
388    uv (batchsize, length, 1)
389    """
390
391    def __init__(
392        self,
393        sampling_rate,
394        harmonic_num=0,
395        sine_amp=0.1,
396        add_noise_std=0.003,
397        voiced_threshod=0,
398        is_half=True,
399    ):
400        super(SourceModuleHnNSF, self).__init__()
401
402        self.sine_amp = sine_amp
403        self.noise_std = add_noise_std
404        self.is_half = is_half
405        # to produce sine waveforms
406        self.l_sin_gen = SineGen(
407            sampling_rate, harmonic_num, sine_amp, add_noise_std, voiced_threshod
408        )
409
410        # to merge source harmonics into a single excitation
411        self.l_linear = torch.nn.Linear(harmonic_num + 1, 1)
412        self.l_tanh = torch.nn.Tanh()
413
414    def forward(self, x, upp=None):
415        sine_wavs, uv, _ = self.l_sin_gen(x, upp)
416        if self.is_half:
417            sine_wavs = sine_wavs.half()
418        sine_merge = self.l_tanh(self.l_linear(sine_wavs))
419        return sine_merge, None, None  # noise, uv
420
421
422class GeneratorNSF(torch.nn.Module):
423    def __init__(
424        self,
425        initial_channel,
426        resblock,
427        resblock_kernel_sizes,
428        resblock_dilation_sizes,
429        upsample_rates,
430        upsample_initial_channel,
431        upsample_kernel_sizes,
432        gin_channels,
433        sr,
434        is_half=False,
435    ):
436        super(GeneratorNSF, self).__init__()
437        self.num_kernels = len(resblock_kernel_sizes)
438        self.num_upsamples = len(upsample_rates)
439
440        self.f0_upsamp = torch.nn.Upsample(scale_factor=np.prod(upsample_rates))
441        self.m_source = SourceModuleHnNSF(
442            sampling_rate=sr, harmonic_num=0, is_half=is_half
443        )
444        self.noise_convs = nn.ModuleList()
445        self.conv_pre = Conv1d(
446            initial_channel, upsample_initial_channel, 7, 1, padding=3
447        )
448        resblock = modules.ResBlock1 if resblock == "1" else modules.ResBlock2
449
450        self.ups = nn.ModuleList()
451        for i, (u, k) in enumerate(zip(upsample_rates, upsample_kernel_sizes)):
452            c_cur = upsample_initial_channel // (2 ** (i + 1))
453            self.ups.append(
454                weight_norm(
455                    ConvTranspose1d(
456                        upsample_initial_channel // (2**i),
457                        upsample_initial_channel // (2 ** (i + 1)),
458                        k,
459                        u,
460                        padding=(k - u) // 2,
461                    )
462                )
463            )
464            if i + 1 < len(upsample_rates):
465                stride_f0 = np.prod(upsample_rates[i + 1 :])
466                self.noise_convs.append(
467                    Conv1d(
468                        1,
469                        c_cur,
470                        kernel_size=stride_f0 * 2,
471                        stride=stride_f0,
472                        padding=stride_f0 // 2,
473                    )
474                )
475            else:
476                self.noise_convs.append(Conv1d(1, c_cur, kernel_size=1))
477
478        self.resblocks = nn.ModuleList()
479        for i in range(len(self.ups)):
480            ch = upsample_initial_channel // (2 ** (i + 1))
481            for j, (k, d) in enumerate(
482                zip(resblock_kernel_sizes, resblock_dilation_sizes)
483            ):
484                self.resblocks.append(resblock(ch, k, d))
485
486        self.conv_post = Conv1d(ch, 1, 7, 1, padding=3, bias=False)
487        self.ups.apply(init_weights)
488
489        if gin_channels != 0:
490            self.cond = nn.Conv1d(gin_channels, upsample_initial_channel, 1)
491
492        self.upp = np.prod(upsample_rates)
493
494    def forward(self, x, f0, g=None):
495        har_source, noi_source, uv = self.m_source(f0, self.upp)
496        har_source = har_source.transpose(1, 2)
497        x = self.conv_pre(x)
498        if g is not None:
499            x = x + self.cond(g)
500
501        for i in range(self.num_upsamples):
502            x = F.leaky_relu(x, modules.LRELU_SLOPE)
503            x = self.ups[i](x)
504            x_source = self.noise_convs[i](har_source)
505            x = x + x_source
506            xs = None
507            for j in range(self.num_kernels):
508                if xs is None:
509                    xs = self.resblocks[i * self.num_kernels + j](x)
510                else:
511                    xs += self.resblocks[i * self.num_kernels + j](x)
512            x = xs / self.num_kernels
513        x = F.leaky_relu(x)
514        x = self.conv_post(x)
515        x = torch.tanh(x)
516        return x
517
518    def remove_weight_norm(self):
519        for l in self.ups:
520            remove_weight_norm(l)
521        for l in self.resblocks:
522            l.remove_weight_norm()
523
524
525sr2sr = {
526    "32k": 32000,
527    "40k": 40000,
528    "48k": 48000,
529}
530
531
532class SynthesizerTrnMs256NSFsid(nn.Module):
533    def __init__(
534        self,
535        spec_channels,
536        segment_size,
537        inter_channels,
538        hidden_channels,
539        filter_channels,
540        n_heads,
541        n_layers,
542        kernel_size,
543        p_dropout,
544        resblock,
545        resblock_kernel_sizes,
546        resblock_dilation_sizes,
547        upsample_rates,
548        upsample_initial_channel,
549        upsample_kernel_sizes,
550        spk_embed_dim,
551        gin_channels,
552        sr,
553        **kwargs
554    ):
555        super().__init__()
556        if type(sr) == type("strr"):
557            sr = sr2sr[sr]
558        self.spec_channels = spec_channels
559        self.inter_channels = inter_channels
560        self.hidden_channels = hidden_channels
561        self.filter_channels = filter_channels
562        self.n_heads = n_heads
563        self.n_layers = n_layers
564        self.kernel_size = kernel_size
565        self.p_dropout = p_dropout
566        self.resblock = resblock
567        self.resblock_kernel_sizes = resblock_kernel_sizes
568        self.resblock_dilation_sizes = resblock_dilation_sizes
569        self.upsample_rates = upsample_rates
570        self.upsample_initial_channel = upsample_initial_channel
571        self.upsample_kernel_sizes = upsample_kernel_sizes
572        self.segment_size = segment_size
573        self.gin_channels = gin_channels
574        # self.hop_length = hop_length#
575        self.spk_embed_dim = spk_embed_dim
576        self.enc_p = TextEncoder256(
577            inter_channels,
578            hidden_channels,
579            filter_channels,
580            n_heads,
581            n_layers,
582            kernel_size,
583            p_dropout,
584        )
585        self.dec = GeneratorNSF(
586            inter_channels,
587            resblock,
588            resblock_kernel_sizes,
589            resblock_dilation_sizes,
590            upsample_rates,
591            upsample_initial_channel,
592            upsample_kernel_sizes,
593            gin_channels=gin_channels,
594            sr=sr,
595            is_half=kwargs["is_half"],
596        )
597        self.enc_q = PosteriorEncoder(
598            spec_channels,
599            inter_channels,
600            hidden_channels,
601            5,
602            1,
603            16,
604            gin_channels=gin_channels,
605        )
606        self.flow = ResidualCouplingBlock(
607            inter_channels, hidden_channels, 5, 1, 3, gin_channels=gin_channels
608        )
609        self.emb_g = nn.Embedding(self.spk_embed_dim, gin_channels)
610        print("gin_channels:", gin_channels, "self.spk_embed_dim:", self.spk_embed_dim)
611
612    def remove_weight_norm(self):
613        self.dec.remove_weight_norm()
614        self.flow.remove_weight_norm()
615        self.enc_q.remove_weight_norm()
616
617    def forward(
618        self, phone, phone_lengths, pitch, pitchf, y, y_lengths, ds
619    ):  # Here ds is id, [bs,1]
620        # print(1,pitch.shape)#[bs,t]
621        g = self.emb_g(ds).unsqueeze(-1)  # [b, 256, 1]##1 is t, broadcast
622        m_p, logs_p, x_mask = self.enc_p(phone, pitch, phone_lengths)
623        z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g)
624        z_p = self.flow(z, y_mask, g=g)
625        z_slice, ids_slice = commons.rand_slice_segments(
626            z, y_lengths, self.segment_size
627        )
628        # print(-1,pitchf.shape,ids_slice,self.segment_size,self.hop_length,self.segment_size//self.hop_length)
629        pitchf = commons.slice_segments2(pitchf, ids_slice, self.segment_size)
630        # print(-2,pitchf.shape,z_slice.shape)
631        o = self.dec(z_slice, pitchf, g=g)
632        return o, ids_slice, x_mask, y_mask, (z, z_p, m_p, logs_p, m_q, logs_q)
633
634    def infer(self, phone, phone_lengths, pitch, nsff0, sid, rate=None):
635        g = self.emb_g(sid).unsqueeze(-1)
636        m_p, logs_p, x_mask = self.enc_p(phone, pitch, phone_lengths)
637        z_p = (m_p + torch.exp(logs_p) * torch.randn_like(m_p) * 0.66666) * x_mask
638        if rate:
639            head = int(z_p.shape[2] * rate)
640            z_p = z_p[:, :, -head:]
641            x_mask = x_mask[:, :, -head:]
642            nsff0 = nsff0[:, -head:]
643        z = self.flow(z_p, x_mask, g=g, reverse=True)
644        o = self.dec(z * x_mask, nsff0, g=g)
645        return o, x_mask, (z, z_p, m_p, logs_p)
646
647
648class SynthesizerTrnMs768NSFsid(nn.Module):
649    def __init__(
650        self,
651        spec_channels,
652        segment_size,
653        inter_channels,
654        hidden_channels,
655        filter_channels,
656        n_heads,
657        n_layers,
658        kernel_size,
659        p_dropout,
660        resblock,
661        resblock_kernel_sizes,
662        resblock_dilation_sizes,
663        upsample_rates,
664        upsample_initial_channel,
665        upsample_kernel_sizes,
666        spk_embed_dim,
667        gin_channels,
668        sr,
669        **kwargs
670    ):
671        super().__init__()
672        if type(sr) == type("strr"):
673            sr = sr2sr[sr]
674        self.spec_channels = spec_channels
675        self.inter_channels = inter_channels
676        self.hidden_channels = hidden_channels
677        self.filter_channels = filter_channels
678        self.n_heads = n_heads
679        self.n_layers = n_layers
680        self.kernel_size = kernel_size
681        self.p_dropout = p_dropout
682        self.resblock = resblock
683        self.resblock_kernel_sizes = resblock_kernel_sizes
684        self.resblock_dilation_sizes = resblock_dilation_sizes
685        self.upsample_rates = upsample_rates
686        self.upsample_initial_channel = upsample_initial_channel
687        self.upsample_kernel_sizes = upsample_kernel_sizes
688        self.segment_size = segment_size
689        self.gin_channels = gin_channels
690        # self.hop_length = hop_length#
691        self.spk_embed_dim = spk_embed_dim
692        self.enc_p = TextEncoder768(
693            inter_channels,
694            hidden_channels,
695            filter_channels,
696            n_heads,
697            n_layers,
698            kernel_size,
699            p_dropout,
700        )
701        self.dec = GeneratorNSF(
702            inter_channels,
703            resblock,
704            resblock_kernel_sizes,
705            resblock_dilation_sizes,
706            upsample_rates,
707            upsample_initial_channel,
708            upsample_kernel_sizes,
709            gin_channels=gin_channels,
710            sr=sr,
711            is_half=kwargs["is_half"],
712        )
713        self.enc_q = PosteriorEncoder(
714            spec_channels,
715            inter_channels,
716            hidden_channels,
717            5,
718            1,
719            16,
720            gin_channels=gin_channels,
721        )
722        self.flow = ResidualCouplingBlock(
723            inter_channels, hidden_channels, 5, 1, 3, gin_channels=gin_channels
724        )
725        self.emb_g = nn.Embedding(self.spk_embed_dim, gin_channels)
726        print("gin_channels:", gin_channels, "self.spk_embed_dim:", self.spk_embed_dim)
727
728    def remove_weight_norm(self):
729        self.dec.remove_weight_norm()
730        self.flow.remove_weight_norm()
731        self.enc_q.remove_weight_norm()
732
733    def forward(
734        self, phone, phone_lengths, pitch, pitchf, y, y_lengths, ds
735    ):  # Here ds is id,[bs,1]
736        # print(1,pitch.shape)#[bs,t]
737        g = self.emb_g(ds).unsqueeze(-1)  # [b, 256, 1]##1 is t, broadcast
738        m_p, logs_p, x_mask = self.enc_p(phone, pitch, phone_lengths)
739        z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g)
740        z_p = self.flow(z, y_mask, g=g)
741        z_slice, ids_slice = commons.rand_slice_segments(
742            z, y_lengths, self.segment_size
743        )
744        # print(-1,pitchf.shape,ids_slice,self.segment_size,self.hop_length,self.segment_size//self.hop_length)
745        pitchf = commons.slice_segments2(pitchf, ids_slice, self.segment_size)
746        # print(-2,pitchf.shape,z_slice.shape)
747        o = self.dec(z_slice, pitchf, g=g)
748        return o, ids_slice, x_mask, y_mask, (z, z_p, m_p, logs_p, m_q, logs_q)
749
750    def infer(self, phone, phone_lengths, pitch, nsff0, sid, rate=None):
751        g = self.emb_g(sid).unsqueeze(-1)
752        m_p, logs_p, x_mask = self.enc_p(phone, pitch, phone_lengths)
753        z_p = (m_p + torch.exp(logs_p) * torch.randn_like(m_p) * 0.66666) * x_mask
754        if rate:
755            head = int(z_p.shape[2] * rate)
756            z_p = z_p[:, :, -head:]
757            x_mask = x_mask[:, :, -head:]
758            nsff0 = nsff0[:, -head:]
759        z = self.flow(z_p, x_mask, g=g, reverse=True)
760        o = self.dec(z * x_mask, nsff0, g=g)
761        return o, x_mask, (z, z_p, m_p, logs_p)
762
763
764class SynthesizerTrnMs256NSFsid_nono(nn.Module):
765    def __init__(
766        self,
767        spec_channels,
768        segment_size,
769        inter_channels,
770        hidden_channels,
771        filter_channels,
772        n_heads,
773        n_layers,
774        kernel_size,
775        p_dropout,
776        resblock,
777        resblock_kernel_sizes,
778        resblock_dilation_sizes,
779        upsample_rates,
780        upsample_initial_channel,
781        upsample_kernel_sizes,
782        spk_embed_dim,
783        gin_channels,
784        sr=None,
785        **kwargs
786    ):
787        super().__init__()
788        self.spec_channels = spec_channels
789        self.inter_channels = inter_channels
790        self.hidden_channels = hidden_channels
791        self.filter_channels = filter_channels
792        self.n_heads = n_heads
793        self.n_layers = n_layers
794        self.kernel_size = kernel_size
795        self.p_dropout = p_dropout
796        self.resblock = resblock
797        self.resblock_kernel_sizes = resblock_kernel_sizes
798        self.resblock_dilation_sizes = resblock_dilation_sizes
799        self.upsample_rates = upsample_rates
800        self.upsample_initial_channel = upsample_initial_channel
801        self.upsample_kernel_sizes = upsample_kernel_sizes
802        self.segment_size = segment_size
803        self.gin_channels = gin_channels
804        # self.hop_length = hop_length#
805        self.spk_embed_dim = spk_embed_dim
806        self.enc_p = TextEncoder256(
807            inter_channels,
808            hidden_channels,
809            filter_channels,
810            n_heads,
811            n_layers,
812            kernel_size,
813            p_dropout,
814            f0=False,
815        )
816        self.dec = Generator(
817            inter_channels,
818            resblock,
819            resblock_kernel_sizes,
820            resblock_dilation_sizes,
821            upsample_rates,
822            upsample_initial_channel,
823            upsample_kernel_sizes,
824            gin_channels=gin_channels,
825        )
826        self.enc_q = PosteriorEncoder(
827            spec_channels,
828            inter_channels,
829            hidden_channels,
830            5,
831            1,
832            16,
833            gin_channels=gin_channels,
834        )
835        self.flow = ResidualCouplingBlock(
836            inter_channels, hidden_channels, 5, 1, 3, gin_channels=gin_channels
837        )
838        self.emb_g = nn.Embedding(self.spk_embed_dim, gin_channels)
839        print("gin_channels:", gin_channels, "self.spk_embed_dim:", self.spk_embed_dim)
840
841    def remove_weight_norm(self):
842        self.dec.remove_weight_norm()
843        self.flow.remove_weight_norm()
844        self.enc_q.remove_weight_norm()
845
846    def forward(self, phone, phone_lengths, y, y_lengths, ds):  # Here ds is id,[bs,1]
847        g = self.emb_g(ds).unsqueeze(-1)  # [b, 256, 1]##1 is t, broadcast
848        m_p, logs_p, x_mask = self.enc_p(phone, None, phone_lengths)
849        z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g)
850        z_p = self.flow(z, y_mask, g=g)
851        z_slice, ids_slice = commons.rand_slice_segments(
852            z, y_lengths, self.segment_size
853        )
854        o = self.dec(z_slice, g=g)
855        return o, ids_slice, x_mask, y_mask, (z, z_p, m_p, logs_p, m_q, logs_q)
856
857    def infer(self, phone, phone_lengths, sid, rate=None):
858        g = self.emb_g(sid).unsqueeze(-1)
859        m_p, logs_p, x_mask = self.enc_p(phone, None, phone_lengths)
860        z_p = (m_p + torch.exp(logs_p) * torch.randn_like(m_p) * 0.66666) * x_mask
861        if rate:
862            head = int(z_p.shape[2] * rate)
863            z_p = z_p[:, :, -head:]
864            x_mask = x_mask[:, :, -head:]
865        z = self.flow(z_p, x_mask, g=g, reverse=True)
866        o = self.dec(z * x_mask, g=g)
867        return o, x_mask, (z, z_p, m_p, logs_p)
868
869
870class SynthesizerTrnMs768NSFsid_nono(nn.Module):
871    def __init__(
872        self,
873        spec_channels,
874        segment_size,
875        inter_channels,
876        hidden_channels,
877        filter_channels,
878        n_heads,
879        n_layers,
880        kernel_size,
881        p_dropout,
882        resblock,
883        resblock_kernel_sizes,
884        resblock_dilation_sizes,
885        upsample_rates,
886        upsample_initial_channel,
887        upsample_kernel_sizes,
888        spk_embed_dim,
889        gin_channels,
890        sr=None,
891        **kwargs
892    ):
893        super().__init__()
894        self.spec_channels = spec_channels
895        self.inter_channels = inter_channels
896        self.hidden_channels = hidden_channels
897        self.filter_channels = filter_channels
898        self.n_heads = n_heads
899        self.n_layers = n_layers
900        self.kernel_size = kernel_size
901        self.p_dropout = p_dropout
902        self.resblock = resblock
903        self.resblock_kernel_sizes = resblock_kernel_sizes
904        self.resblock_dilation_sizes = resblock_dilation_sizes
905        self.upsample_rates = upsample_rates
906        self.upsample_initial_channel = upsample_initial_channel
907        self.upsample_kernel_sizes = upsample_kernel_sizes
908        self.segment_size = segment_size
909        self.gin_channels = gin_channels
910        # self.hop_length = hop_length#
911        self.spk_embed_dim = spk_embed_dim
912        self.enc_p = TextEncoder768(
913            inter_channels,
914            hidden_channels,
915            filter_channels,
916            n_heads,
917            n_layers,
918            kernel_size,
919            p_dropout,
920            f0=False,
921        )
922        self.dec = Generator(
923            inter_channels,
924            resblock,
925            resblock_kernel_sizes,
926            resblock_dilation_sizes,
927            upsample_rates,
928            upsample_initial_channel,
929            upsample_kernel_sizes,
930            gin_channels=gin_channels,
931        )
932        self.enc_q = PosteriorEncoder(
933            spec_channels,
934            inter_channels,
935            hidden_channels,
936            5,
937            1,
938            16,
939            gin_channels=gin_channels,
940        )
941        self.flow = ResidualCouplingBlock(
942            inter_channels, hidden_channels, 5, 1, 3, gin_channels=gin_channels
943        )
944        self.emb_g = nn.Embedding(self.spk_embed_dim, gin_channels)
945        print("gin_channels:", gin_channels, "self.spk_embed_dim:", self.spk_embed_dim)
946
947    def remove_weight_norm(self):
948        self.dec.remove_weight_norm()
949        self.flow.remove_weight_norm()
950        self.enc_q.remove_weight_norm()
951
952    def forward(self, phone, phone_lengths, y, y_lengths, ds):  # Here ds is id,[bs,1]
953        g = self.emb_g(ds).unsqueeze(-1)  # [b, 256, 1]##1 is t, broadcast
954        m_p, logs_p, x_mask = self.enc_p(phone, None, phone_lengths)
955        z, m_q, logs_q, y_mask = self.enc_q(y, y_lengths, g=g)
956        z_p = self.flow(z, y_mask, g=g)
957        z_slice, ids_slice = commons.rand_slice_segments(
958            z, y_lengths, self.segment_size
959        )
960        o = self.dec(z_slice, g=g)
961        return o, ids_slice, x_mask, y_mask, (z, z_p, m_p, logs_p, m_q, logs_q)
962
963    def infer(self, phone, phone_lengths, sid, rate=None):
964        g = self.emb_g(sid).unsqueeze(-1)
965        m_p, logs_p, x_mask = self.enc_p(phone, None, phone_lengths)
966        z_p = (m_p + torch.exp(logs_p) * torch.randn_like(m_p) * 0.66666) * x_mask
967        if rate:
968            head = int(z_p.shape[2] * rate)
969            z_p = z_p[:, :, -head:]
970            x_mask = x_mask[:, :, -head:]
971        z = self.flow(z_p, x_mask, g=g, reverse=True)
972        o = self.dec(z * x_mask, g=g)
973        return o, x_mask, (z, z_p, m_p, logs_p)
974
975
976class MultiPeriodDiscriminator(torch.nn.Module):
977    def __init__(self, use_spectral_norm=False):
978        super(MultiPeriodDiscriminator, self).__init__()
979        periods = [2, 3, 5, 7, 11, 17]
980        # periods = [3, 5, 7, 11, 17, 23, 37]
981
982        discs = [DiscriminatorS(use_spectral_norm=use_spectral_norm)]
983        discs = discs + [
984            DiscriminatorP(i, use_spectral_norm=use_spectral_norm) for i in periods
985        ]
986        self.discriminators = nn.ModuleList(discs)
987
988    def forward(self, y, y_hat):
989        y_d_rs = []  #
990        y_d_gs = []
991        fmap_rs = []
992        fmap_gs = []
993        for i, d in enumerate(self.discriminators):
994            y_d_r, fmap_r = d(y)
995            y_d_g, fmap_g = d(y_hat)
996            # for j in range(len(fmap_r)):
997            #     print(i,j,y.shape,y_hat.shape,fmap_r[j].shape,fmap_g[j].shape)
998            y_d_rs.append(y_d_r)
999            y_d_gs.append(y_d_g)
1000            fmap_rs.append(fmap_r)
1001            fmap_gs.append(fmap_g)
1002
1003        return y_d_rs, y_d_gs, fmap_rs, fmap_gs
1004
1005
1006class MultiPeriodDiscriminatorV2(torch.nn.Module):
1007    def __init__(self, use_spectral_norm=False):
1008        super(MultiPeriodDiscriminatorV2, self).__init__()
1009        # periods = [2, 3, 5, 7, 11, 17]
1010        periods = [2, 3, 5, 7, 11, 17, 23, 37]
1011
1012        discs = [DiscriminatorS(use_spectral_norm=use_spectral_norm)]
1013        discs = discs + [
1014            DiscriminatorP(i, use_spectral_norm=use_spectral_norm) for i in periods
1015        ]
1016        self.discriminators = nn.ModuleList(discs)
1017
1018    def forward(self, y, y_hat):
1019        y_d_rs = []  #
1020        y_d_gs = []
1021        fmap_rs = []
1022        fmap_gs = []
1023        for i, d in enumerate(self.discriminators):
1024            y_d_r, fmap_r = d(y)
1025            y_d_g, fmap_g = d(y_hat)
1026            # for j in range(len(fmap_r)):
1027            #     print(i,j,y.shape,y_hat.shape,fmap_r[j].shape,fmap_g[j].shape)
1028            y_d_rs.append(y_d_r)
1029            y_d_gs.append(y_d_g)
1030            fmap_rs.append(fmap_r)
1031            fmap_gs.append(fmap_g)
1032
1033        return y_d_rs, y_d_gs, fmap_rs, fmap_gs
1034
1035
1036class DiscriminatorS(torch.nn.Module):
1037    def __init__(self, use_spectral_norm=False):
1038        super(DiscriminatorS, self).__init__()
1039        norm_f = weight_norm if use_spectral_norm == False else spectral_norm
1040        self.convs = nn.ModuleList(
1041            [
1042                norm_f(Conv1d(1, 16, 15, 1, padding=7)),
1043                norm_f(Conv1d(16, 64, 41, 4, groups=4, padding=20)),
1044                norm_f(Conv1d(64, 256, 41, 4, groups=16, padding=20)),
1045                norm_f(Conv1d(256, 1024, 41, 4, groups=64, padding=20)),
1046                norm_f(Conv1d(1024, 1024, 41, 4, groups=256, padding=20)),
1047                norm_f(Conv1d(1024, 1024, 5, 1, padding=2)),
1048            ]
1049        )
1050        self.conv_post = norm_f(Conv1d(1024, 1, 3, 1, padding=1))
1051
1052    def forward(self, x):
1053        fmap = []
1054
1055        for l in self.convs:
1056            x = l(x)
1057            x = F.leaky_relu(x, modules.LRELU_SLOPE)
1058            fmap.append(x)
1059        x = self.conv_post(x)
1060        fmap.append(x)
1061        x = torch.flatten(x, 1, -1)
1062
1063        return x, fmap
1064
1065
1066class DiscriminatorP(torch.nn.Module):
1067    def __init__(self, period, kernel_size=5, stride=3, use_spectral_norm=False):
1068        super(DiscriminatorP, self).__init__()
1069        self.period = period
1070        self.use_spectral_norm = use_spectral_norm
1071        norm_f = weight_norm if use_spectral_norm == False else spectral_norm
1072        self.convs = nn.ModuleList(
1073            [
1074                norm_f(
1075                    Conv2d(
1076                        1,
1077                        32,
1078                        (kernel_size, 1),
1079                        (stride, 1),
1080                        padding=(get_padding(kernel_size, 1), 0),
1081                    )
1082                ),
1083                norm_f(
1084                    Conv2d(
1085                        32,
1086                        128,
1087                        (kernel_size, 1),
1088                        (stride, 1),
1089                        padding=(get_padding(kernel_size, 1), 0),
1090                    )
1091                ),
1092                norm_f(
1093                    Conv2d(
1094                        128,
1095                        512,
1096                        (kernel_size, 1),
1097                        (stride, 1),
1098                        padding=(get_padding(kernel_size, 1), 0),
1099                    )
1100                ),
1101                norm_f(
1102                    Conv2d(
1103                        512,
1104                        1024,
1105                        (kernel_size, 1),
1106                        (stride, 1),
1107                        padding=(get_padding(kernel_size, 1), 0),
1108                    )
1109                ),
1110                norm_f(
1111                    Conv2d(
1112                        1024,
1113                        1024,
1114                        (kernel_size, 1),
1115                        1,
1116                        padding=(get_padding(kernel_size, 1), 0),
1117                    )
1118                ),
1119            ]
1120        )
1121        self.conv_post = norm_f(Conv2d(1024, 1, (3, 1), 1, padding=(1, 0)))
1122
1123    def forward(self, x):
1124        fmap = []
1125
1126        # 1d to 2d
1127        b, c, t = x.shape
1128        if t % self.period != 0:  # pad first
1129            n_pad = self.period - (t % self.period)
1130            x = F.pad(x, (0, n_pad), "reflect")
1131            t = t + n_pad
1132        x = x.view(b, c, t // self.period, self.period)
1133
1134        for l in self.convs:
1135            x = l(x)
1136            x = F.leaky_relu(x, modules.LRELU_SLOPE)
1137            fmap.append(x)
1138        x = self.conv_post(x)
1139        fmap.append(x)
1140        x = torch.flatten(x, 1, -1)
1141
1142        return x, fmap
1143