| 764 | |
| 765 | |
| 766 | class MelStyleEncoderVAE(nn.Module): |
| 767 | def __init__(self, spec_channels, z_latent_dim, emb_dim): |
| 768 | super().__init__() |
| 769 | self.ref_encoder = MelStyleEncoder(spec_channels, style_vector_dim=emb_dim) |
| 770 | self.fc1 = nn.Linear(emb_dim, z_latent_dim) |
| 771 | self.fc2 = nn.Linear(emb_dim, z_latent_dim) |
| 772 | self.fc3 = nn.Linear(z_latent_dim, emb_dim) |
| 773 | self.z_latent_dim = z_latent_dim |
| 774 | |
| 775 | def reparameterize(self, mu, logvar): |
| 776 | if self.training: |
| 777 | std = torch.exp(0.5 * logvar) |
| 778 | eps = torch.randn_like(std) |
| 779 | return eps.mul(std).add_(mu) |
| 780 | else: |
| 781 | return mu |
| 782 | |
| 783 | def forward(self, inputs, mask=None): |
| 784 | enc_out = self.ref_encoder(inputs.squeeze(-1), mask).squeeze(-1) |
| 785 | mu = self.fc1(enc_out) |
| 786 | logvar = self.fc2(enc_out) |
| 787 | posterior = D.Normal(mu, torch.exp(logvar)) |
| 788 | kl_divergence = D.kl_divergence( |
| 789 | posterior, D.Normal(torch.zeros_like(mu), torch.ones_like(logvar)) |
| 790 | ) |
| 791 | loss_kl = kl_divergence.mean() |
| 792 | |
| 793 | z = posterior.rsample() |
| 794 | style_embed = self.fc3(z) |
| 795 | |
| 796 | return style_embed.unsqueeze(-1), loss_kl |
| 797 | |
| 798 | def infer(self, inputs=None, random_sample=False, manual_latent=None): |
| 799 | if manual_latent is None: |
| 800 | if random_sample: |
| 801 | dev = next(self.parameters()).device |
| 802 | posterior = D.Normal( |
| 803 | torch.zeros(1, self.z_latent_dim, device=dev), |
| 804 | torch.ones(1, self.z_latent_dim, device=dev), |
| 805 | ) |
| 806 | z = posterior.rsample() |
| 807 | else: |
| 808 | enc_out = self.ref_encoder(inputs.transpose(1, 2)) |
| 809 | mu = self.fc1(enc_out) |
| 810 | z = mu |
| 811 | else: |
| 812 | z = manual_latent |
| 813 | style_embed = self.fc3(z) |
| 814 | return style_embed.unsqueeze(-1), z |
| 815 | |
| 816 | |
| 817 | class ActNorm(nn.Module): |
nothing calls this directly
no outgoing calls
no test coverage detected