Args: x: input tensor (B, C, H, W) in range [-1, 1] scaled with self.scale and shifted with self.shift return_posterior: return the posterior distribution
(self, x: torch.Tensor, return_posterior=False)
| 493 | |
| 494 | @torch.no_grad() |
| 495 | def encode(self, x: torch.Tensor, return_posterior=False): |
| 496 | """ |
| 497 | Args: |
| 498 | x: input tensor (B, C, H, W) in range [-1, 1] scaled with |
| 499 | self.scale and shifted with self.shift |
| 500 | return_posterior: return the posterior distribution |
| 501 | """ |
| 502 | h = self.encoder(x) |
| 503 | moments = self.quant_conv(h) |
| 504 | posterior = DiagonalGaussianDistribution(moments) |
| 505 | if return_posterior: |
| 506 | return posterior |
| 507 | latent = posterior.mode() |
| 508 | return (latent + self.shift) * self.scale |
| 509 | |
| 510 | @torch.no_grad() |
| 511 | def decode(self, z: torch.Tensor): |
no test coverage detected