(self, inputs=None, random_sample=False, manual_latent=None)
| 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