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

Method encode

PyTorch-VAE/models/iwae.py:78–93  ·  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

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 """

Callers 1

forwardMethod · 0.95

Calls

no outgoing calls

Tested by

no test coverage detected