MCPcopy Create free account
hub / github.com/KeepTryingTo/Pytorch-GAN / encode

Method encode

PyTorch-VAE/models/dfcvae.py:90–105  ·  view source on GitHub ↗

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)

Source from the content-addressed store, hash-verified

88
89
90 def encode(self, input: Tensor) -> List[Tensor]:
91 """
92 Encodes the input by passing through the encoder network
93 and returns the latent codes.
94 :param input: (Tensor) Input tensor to encoder [N x C x H x W]
95 :return: (Tensor) List of latent codes
96 """
97 result = self.encoder(input)
98 result = torch.flatten(result, start_dim=1)
99
100 # Split the result into mu and var components
101 # of the latent Gaussian distribution
102 mu = self.fc_mu(result)
103 log_var = self.fc_var(result)
104
105 return [mu, log_var]
106
107 def decode(self, z: Tensor) -> Tensor:
108 """

Callers 1

forwardMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected