#!/usr/bin/env python3
"""Train a small variational autoencoder on generated two-dimensional data."""

from __future__ import annotations

import argparse

import torch
from torch import nn
from torch.nn import functional as F


torch.set_num_threads(1)


CLUSTER_CENTRES = torch.tensor(
    [[-2.0, -2.0], [-2.0, 2.0], [2.0, -2.0], [2.0, 2.0]],
    dtype=torch.float32,
)


def make_dataset(
    count: int,
    seed: int,
    noise_standard_deviation: float = 0.30,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Return generated points of shape ``(count, 2)`` and cluster labels."""
    if count < 1:
        raise ValueError("count must be positive")
    generator = torch.Generator(device="cpu").manual_seed(seed)
    labels = torch.randint(0, len(CLUSTER_CENTRES), (count,), generator=generator)
    noise = torch.randn((count, 2), generator=generator) * noise_standard_deviation
    points = CLUSTER_CENTRES[labels] + noise
    return points, labels


def reparameterize(
    mean: torch.Tensor,
    log_variance: torch.Tensor,
    *,
    generator: torch.Generator | None = None,
    epsilon: torch.Tensor | None = None,
) -> torch.Tensor:
    """Sample z = mean + standard_deviation * epsilon with differentiable scaling."""
    standard_deviation = torch.exp(0.5 * log_variance)
    if epsilon is None:
        epsilon = torch.randn(
            standard_deviation.shape,
            dtype=standard_deviation.dtype,
            device=standard_deviation.device,
            generator=generator,
        )
    return mean + standard_deviation * epsilon


class TinyVAE(nn.Module):
    """A two-dimensional encoder and decoder with a Gaussian latent variable."""

    def __init__(self, hidden_size: int = 32, latent_size: int = 2) -> None:
        super().__init__()
        self.encoder = nn.Sequential(nn.Linear(2, hidden_size), nn.Tanh())
        self.mean_head = nn.Linear(hidden_size, latent_size)
        self.log_variance_head = nn.Linear(hidden_size, latent_size)
        self.decoder = nn.Sequential(
            nn.Linear(latent_size, hidden_size),
            nn.Tanh(),
            nn.Linear(hidden_size, 2),
        )

    def encode(self, x: torch.Tensor) -> tuple[torch.Tensor, torch.Tensor]:
        hidden = self.encoder(x)
        return self.mean_head(hidden), self.log_variance_head(hidden)

    def decode(self, z: torch.Tensor) -> torch.Tensor:
        return self.decoder(z)

    def forward(
        self,
        x: torch.Tensor,
        *,
        generator: torch.Generator | None = None,
    ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
        mean, log_variance = self.encode(x)
        z = reparameterize(mean, log_variance, generator=generator)
        return self.decode(z), mean, log_variance


def vae_loss(
    reconstruction: torch.Tensor,
    target: torch.Tensor,
    mean: torch.Tensor,
    log_variance: torch.Tensor,
    beta: float = 0.05,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
    """Return total, reconstruction, and KL loss averaged over examples."""
    reconstruction_loss = ((reconstruction - target) ** 2).sum(dim=1).mean()
    kl_loss = -0.5 * (
        1.0 + log_variance - mean.square() - log_variance.exp()
    ).sum(dim=1).mean()
    total = reconstruction_loss + beta * kl_loss
    return total, reconstruction_loss, kl_loss


def evaluate(
    model: TinyVAE,
    points: torch.Tensor,
) -> tuple[float, float]:
    """Use the latent mean for stable held-out reconstruction measurement."""
    model.eval()
    with torch.no_grad():
        mean, log_variance = model.encode(points)
        reconstruction = model.decode(mean)
        reconstruction_mse = F.mse_loss(reconstruction, points).item()
        kl = -0.5 * (
            1.0 + log_variance - mean.square() - log_variance.exp()
        ).sum(dim=1).mean().item()
    return reconstruction_mse, kl


def train_model(steps: int = 500, beta: float = 0.05) -> TinyVAE:
    """Train with deterministic data, batches, initialization, and latent noise."""
    if steps < 1:
        raise ValueError("steps must be positive")
    torch.manual_seed(2026)
    points, _ = make_dataset(1024, seed=50)
    batch_generator = torch.Generator(device="cpu").manual_seed(51)
    noise_generator = torch.Generator(device="cpu").manual_seed(52)
    model = TinyVAE()
    optimizer = torch.optim.Adam(model.parameters(), lr=0.01)

    model.train()
    for _ in range(steps):
        indices = torch.randint(0, len(points), (128,), generator=batch_generator)
        batch = points[indices]
        optimizer.zero_grad()
        reconstruction, mean, log_variance = model(batch, generator=noise_generator)
        loss, _, _ = vae_loss(reconstruction, batch, mean, log_variance, beta=beta)
        loss.backward()
        optimizer.step()
    return model


def sample_from_prior(model: TinyVAE, count: int, seed: int) -> torch.Tensor:
    """Decode standard-normal latent samples without tracking gradients."""
    generator = torch.Generator(device="cpu").manual_seed(seed)
    z = torch.randn((count, model.mean_head.out_features), generator=generator)
    model.eval()
    with torch.no_grad():
        return model.decode(z)


def run(steps: int = 500, smoke_test: bool = False) -> dict[str, float]:
    """Train, evaluate against a mean baseline, and generate fresh samples."""
    beta = 0.05
    effective_steps = min(steps, 3) if smoke_test else steps
    train_points, _ = make_dataset(1024, seed=50)
    validation_points, _ = make_dataset(256, seed=60)
    model = train_model(effective_steps, beta=beta)
    validation_mse, validation_kl = evaluate(model, validation_points)
    baseline = train_points.mean(dim=0).expand_as(validation_points)
    baseline_mse = F.mse_loss(baseline, validation_points).item()
    samples = sample_from_prior(model, count=8, seed=70)

    print("train and validation shapes:", tuple(train_points.shape), tuple(validation_points.shape))
    print(f"mean-predictor baseline MSE: {baseline_mse:.4f}")
    print(f"held-out deterministic reconstruction MSE: {validation_mse:.4f}")
    print(f"held-out KL per example: {validation_kl:.4f}")
    print("sample shape:", tuple(samples.shape))
    print("first three generated points:")
    for point in samples[:3]:
        print(f"  ({point[0].item():+.3f}, {point[1].item():+.3f})")
    if smoke_test:
        print("smoke test: learning thresholds skipped")
    else:
        assert validation_mse < baseline_mse * 0.20
        assert validation_kl > 0.01
        assert torch.isfinite(samples).all()

    return {
        "baseline_mse": baseline_mse,
        "validation_mse": validation_mse,
        "validation_kl": validation_kl,
    }


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--steps", type=int, default=500)
    parser.add_argument(
        "--smoke-test",
        action="store_true",
        help="run at most three optimizer steps and skip learning thresholds",
    )
    return parser.parse_args()


if __name__ == "__main__":
    arguments = parse_args()
    run(steps=arguments.steps, smoke_test=arguments.smoke_test)
