High-level encode with tiling support
(self, x, return_dict=True)
| 262 | return dec |
| 263 | |
| 264 | def encode(self, x, return_dict=True): |
| 265 | """High-level encode with tiling support""" |
| 266 | if x.shape[-2] >= self.tile_sample_min_height and x.shape[-1] >= self.tile_sample_min_width: |
| 267 | mu = self.tiled_encode(x) |
| 268 | # DiagonalGaussianDistribution expects [mu, logvar] concatenated along channel dim |
| 269 | # It will split by channel, so we need to provide 2*z_dim channels |
| 270 | # Use zeros for logvar since we're using deterministic encoding (mode) |
| 271 | logvar = torch.zeros_like(mu) |
| 272 | latent = torch.cat([mu, logvar], dim=1) |
| 273 | from diffusers.models.autoencoders.vae import DiagonalGaussianDistribution |
| 274 | from diffusers.models.modeling_outputs import AutoencoderKLOutput |
| 275 | posterior = DiagonalGaussianDistribution(latent) |
| 276 | if not return_dict: |
| 277 | return (posterior,) |
| 278 | return AutoencoderKLOutput(latent_dist=posterior) |
| 279 | else: |
| 280 | return self.vae.encode(x, return_dict=return_dict) |
| 281 | |
| 282 | def decode(self, z, return_dict=True): |
| 283 | """High-level decode with tiling support""" |
nothing calls this directly
no test coverage detected