Encodes the input by passing through the encoder network and returns the latent codes. :param input: (Tensor) Input tensor to encoder [N x C x H x W] :return: (Tensor) List of latent codes
(self, input: Tensor)
| 76 | nn.Tanh()) |
| 77 | |
| 78 | def encode(self, input: Tensor) -> List[Tensor]: |
| 79 | """ |
| 80 | Encodes the input by passing through the encoder network |
| 81 | and returns the latent codes. |
| 82 | :param input: (Tensor) Input tensor to encoder [N x C x H x W] |
| 83 | :return: (Tensor) List of latent codes |
| 84 | """ |
| 85 | result = self.encoder(input) |
| 86 | result = torch.flatten(result, start_dim=1) |
| 87 | |
| 88 | # Split the result into mu and var components |
| 89 | # of the latent Gaussian distribution |
| 90 | mu = self.fc_mu(result) |
| 91 | log_var = self.fc_var(result) |
| 92 | |
| 93 | return [mu, log_var] |
| 94 | |
| 95 | def decode(self, z: Tensor) -> Tensor: |
| 96 | """ |