(self, inputs, mask=None)
| 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: |
nothing calls this directly
no outgoing calls
no test coverage detected