| 554 | |
| 555 | |
| 556 | class 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) |