MCPcopy Create free account
hub / github.com/TextGeneratorio/text-generator.io / Decoder

Class Decoder

questions/inference_server/istftnet.py:556–619  ·  view source on GitHub ↗

Source from the content-addressed store, hash-verified

554
555
556class Decoder(nn.Module):
557 def __init__(
558 self,
559 dim_in=512,
560 F0_channel=512,
561 style_dim=64,
562 dim_out=80,
563 resblock_kernel_sizes=[3, 7, 11],
564 upsample_rates=[10, 6],
565 upsample_initial_channel=512,
566 resblock_dilation_sizes=[[1, 3, 5], [1, 3, 5], [1, 3, 5]],
567 upsample_kernel_sizes=[20, 12],
568 gen_istft_n_fft=20,
569 gen_istft_hop_size=5,
570 ):
571 super().__init__()
572
573 self.decode = nn.ModuleList()
574
575 self.encode = AdainResBlk1d(dim_in + 2, 1024, style_dim)
576
577 self.decode.append(AdainResBlk1d(1024 + 2 + 64, 1024, style_dim))
578 self.decode.append(AdainResBlk1d(1024 + 2 + 64, 1024, style_dim))
579 self.decode.append(AdainResBlk1d(1024 + 2 + 64, 1024, style_dim))
580 self.decode.append(AdainResBlk1d(1024 + 2 + 64, 512, style_dim, upsample=True))
581
582 self.F0_conv = weight_norm(nn.Conv1d(1, 1, kernel_size=3, stride=2, groups=1, padding=1))
583
584 self.N_conv = weight_norm(nn.Conv1d(1, 1, kernel_size=3, stride=2, groups=1, padding=1))
585
586 self.asr_res = nn.Sequential(
587 weight_norm(nn.Conv1d(512, 64, kernel_size=1)),
588 )
589
590 self.generator = Generator(
591 style_dim,
592 resblock_kernel_sizes,
593 upsample_rates,
594 upsample_initial_channel,
595 resblock_dilation_sizes,
596 upsample_kernel_sizes,
597 gen_istft_n_fft,
598 gen_istft_hop_size,
599 )
600
601 def forward(self, asr, F0_curve, N, s):
602 F0 = self.F0_conv(F0_curve.unsqueeze(1))
603 N = self.N_conv(N.unsqueeze(1))
604
605 x = torch.cat([asr, F0, N], axis=1)
606 x = self.encode(x, s)
607
608 asr_res = self.asr_res(asr)
609
610 res = True
611 for block in self.decode:
612 if res:
613 x = torch.cat([x, asr_res, F0, N], axis=1)

Callers 1

build_modelFunction · 0.85

Calls

no outgoing calls

Tested by

no test coverage detected