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
