import torch import torch.nn as nn from einops import rearrange def vae_sample(mean, scale): stdev = nn.functional.softplus(scale) + 1e-4 var = stdev * stdev logvar = torch.log(var) latents = torch.randn_like(mean) * stdev + mean kl = (mean * mean + var - logvar - 1).sum(1).mean() return latents, kl class VAEBottleneck(nn.Module): def __init__(self, latent_dim, dim, is_discrete=False, **kwargs): super().__init__() self.project_in = nn.Linear(latent_dim, dim * 2) self.project_out = ( nn.Linear(dim, latent_dim) if latent_dim != dim else nn.Identity() ) self.is_discrete = is_discrete def forward(self, x, transpose=True, **kwargs): if transpose: x = rearrange(x, "b d n -> b n d") x = self.project_in(x) mean, scale = x.chunk(2, dim=-1) x, kl = vae_sample(mean, scale) x = self.project_out(x) if transpose: x = rearrange(x, "b n d -> b d n") return { "z": x, "kl": kl, "mean": mean, "scale": scale, } def encode(self, x, return_info=False, **kwargs): if return_info: out = self.forward(x) latents = out["z"] # latents = rearrange(latents, "b n d -> b d n") bottleneck_info = {"kl": float(out["kl"].item())} return latents, bottleneck_info return self.forward(x) def decode(self, x, **kwargs): return x