MCPcopy Create free account
hub / github.com/MYZY-AI/Muyan-TTS / ReferenceEncoder

Class ReferenceEncoder

sovits/models.py:628–683  ·  view source on GitHub ↗

inputs --- [N, Ty/r, n_mels*r] mels outputs --- [N, ref_enc_gru_size]

Source from the content-addressed store, hash-verified

626
627
628class ReferenceEncoder(nn.Module):
629 """
630 inputs --- [N, Ty/r, n_mels*r] mels
631 outputs --- [N, ref_enc_gru_size]
632 """
633
634 def __init__(self, spec_channels, gin_channels=0):
635 super().__init__()
636 self.spec_channels = spec_channels
637 ref_enc_filters = [32, 32, 64, 64, 128, 128]
638 K = len(ref_enc_filters)
639 filters = [1] + ref_enc_filters
640 convs = [
641 weight_norm(
642 nn.Conv2d(
643 in_channels=filters[i],
644 out_channels=filters[i + 1],
645 kernel_size=(3, 3),
646 stride=(2, 2),
647 padding=(1, 1),
648 )
649 )
650 for i in range(K)
651 ]
652 self.convs = nn.ModuleList(convs)
653
654 out_channels = self.calculate_channels(spec_channels, 3, 2, 1, K)
655 self.gru = nn.GRU(
656 input_size=ref_enc_filters[-1] * out_channels,
657 hidden_size=256 // 2,
658 batch_first=True,
659 )
660 self.proj = nn.Linear(128, gin_channels)
661
662 def forward(self, inputs):
663 N = inputs.size(0)
664 out = inputs.view(N, 1, -1, self.spec_channels) # [N, 1, Ty, n_freqs]
665 for conv in self.convs:
666 out = conv(out)
667 # out = wn(out)
668 out = F.relu(out) # [N, 128, Ty//2^K, n_mels//2^K]
669
670 out = out.transpose(1, 2) # [N, Ty//2^K, 128, n_mels//2^K]
671 T = out.size(1)
672 N = out.size(0)
673 out = out.contiguous().view(N, T, -1) # [N, Ty//2^K, 128*n_mels//2^K]
674
675 self.gru.flatten_parameters()
676 memory, out = self.gru(out) # out --- [1, N, 128]
677
678 return self.proj(out.squeeze(0)).unsqueeze(-1)
679
680 def calculate_channels(self, L, kernel_size, stride, pad, n_convs):
681 for i in range(n_convs):
682 L = (L - kernel_size + 2 * pad) // stride + 1
683 return L
684
685

Callers

nothing calls this directly

Calls

no outgoing calls

Tested by

no test coverage detected