Download GeometryForcing/algorithms/vae/common/base_vae.py from BonanDing/worldmem-baseline-evals: direct link, hf CLI and curl.
- Browser
- Download file 1.4 kB
-
https://huggingface.co/BonanDing/worldmem-baseline-evals/resolve/main/GeometryForcing/algorithms/vae/common/base_vae.py
- Command line
-
hf download hf://BonanDing/worldmem-baseline-evals/GeometryForcing/algorithms/vae/common/base_vae.py
-
curl -L -o base_vae.py https://huggingface.co/BonanDing/worldmem-baseline-evals/resolve/main/GeometryForcing/algorithms/vae/common/base_vae.py
1.4 kB
| from typing import Tuple | |
| from abc import ABC, abstractmethod | |
| import torch | |
| import torch.nn as nn | |
| from .distribution import DiagonalGaussianDistribution | |
| class VAE(ABC, nn.Module): | |
| """ | |
| Base VAE class. | |
| """ | |
| def encode(self, x: torch.Tensor) -> DiagonalGaussianDistribution: | |
| """ | |
| Encode input tensor x into latent distribution. | |
| """ | |
| raise NotImplementedError | |
| def decode(self, z: torch.Tensor) -> torch.Tensor: | |
| """ | |
| Decode latent tensor z into original space. | |
| """ | |
| raise NotImplementedError | |
| def forward( | |
| self, sample: torch.Tensor, sample_posterior: bool = True | |
| ) -> Tuple[torch.Tensor, DiagonalGaussianDistribution]: | |
| """ | |
| Forward pass. | |
| Returns: | |
| - dec: reconstructed input, uses mode of latent distribution if sample_posterior is False, otherwise sample from it | |
| - posterior: latent distribution | |
| """ | |
| posterior = self.encode(sample) | |
| if sample_posterior: | |
| z = posterior.sample() | |
| else: | |
| z = posterior.mode() | |
| dec = self.decode(z) | |
| return dec, posterior | |
| def from_pretrained(cls, path: str, **kwargs) -> "VAE": | |
| """ | |
| Load pretrained model from path, with additional kwargs. | |
| """ | |
| raise NotImplementedError | |