Kit-Lemonfoot/vtuber_rvc_models
53
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 