#!/usr/bin/env python3
"""Train small supervised PyTorch models on deterministic generated CPU data.

Examples:

    python examples/pytorch/build_supervised.py --task logistic
    python examples/pytorch/build_supervised.py --task mlp --optimizer adam
    python examples/pytorch/build_supervised.py --task cnn --checkpoint /tmp/cnn.pt

Each task creates independent train, validation, and test sets. The test set is
evaluated once after validation has selected the best in-memory model state.
Nothing is downloaded, and no file is written unless ``--checkpoint`` is used.
"""

from __future__ import annotations

import argparse
import copy
from dataclasses import dataclass
from pathlib import Path

import torch
from torch import nn


torch.set_num_threads(1)


@dataclass(frozen=True)
class DatasetSplit:
    """Features and integer class labels for one independent data split."""

    features: torch.Tensor
    targets: torch.Tensor


@dataclass(frozen=True)
class SupervisedData:
    """Three splits; only ``train`` may influence parameter updates."""

    train: DatasetSplit
    validation: DatasetSplit
    test: DatasetSplit
    class_names: tuple[str, ...]


def _make_linear_split(count: int, seed: int) -> DatasetSplit:
    generator = torch.Generator(device="cpu").manual_seed(seed)
    features = torch.randn((count, 2), generator=generator)
    score = 1.4 * features[:, 0] - 0.9 * features[:, 1] + 0.2
    targets = score.gt(0).long()
    return DatasetSplit(features.float(), targets)


def _make_xor_split(count: int, seed: int) -> DatasetSplit:
    generator = torch.Generator(device="cpu").manual_seed(seed)
    corners = torch.randint(0, 2, (count, 2), generator=generator) * 2 - 1
    features = corners.float() + 0.28 * torch.randn(
        (count, 2), generator=generator
    )
    targets = corners[:, 0].eq(corners[:, 1]).long()
    return DatasetSplit(features, targets)


def _make_image_split(count: int, seed: int, image_size: int = 8) -> DatasetSplit:
    """Make noisy vertical, horizontal, and diagonal line images."""
    generator = torch.Generator(device="cpu").manual_seed(seed)
    targets = torch.arange(count, dtype=torch.long) % 3
    order = torch.randperm(count, generator=generator)
    targets = targets[order]
    images = 0.12 * torch.randn(
        (count, 1, image_size, image_size), generator=generator
    )
    for row, label in enumerate(targets.tolist()):
        if label == 0:
            column = int(torch.randint(1, image_size - 1, (), generator=generator))
            images[row, 0, :, column - 1 : column + 1] += 1.0
        elif label == 1:
            line = int(torch.randint(1, image_size - 1, (), generator=generator))
            images[row, 0, line - 1 : line + 1, :] += 1.0
        else:
            direction = int(torch.randint(0, 2, (), generator=generator))
            for index in range(image_size):
                column = index if direction == 0 else image_size - 1 - index
                images[row, 0, index, column] += 1.25
    return DatasetSplit(images.float(), targets)


def make_datasets(
    task: str,
    *,
    train_count: int = 900,
    validation_count: int = 300,
    test_count: int = 300,
) -> SupervisedData:
    """Create independent generated splits for ``logistic``, ``mlp``, or ``cnn``."""
    if min(train_count, validation_count, test_count) < 6:
        raise ValueError("each split needs at least 6 examples")
    makers = {
        "logistic": _make_linear_split,
        "mlp": _make_xor_split,
        "cnn": _make_image_split,
    }
    if task not in makers:
        raise ValueError("task must be 'logistic', 'mlp', or 'cnn'")
    maker = makers[task]
    names = ("negative", "positive") if task != "cnn" else (
        "vertical",
        "horizontal",
        "diagonal",
    )
    return SupervisedData(
        train=maker(train_count, 101),
        validation=maker(validation_count, 202),
        test=maker(test_count, 303),
        class_names=names,
    )


def make_model(task: str) -> nn.Module:
    """Return the smallest suitable model for one generated task."""
    if task == "logistic":
        return nn.Linear(2, 2)
    if task == "mlp":
        return nn.Sequential(
            nn.Linear(2, 16),
            nn.ReLU(),
            nn.Linear(16, 16),
            nn.ReLU(),
            nn.Linear(16, 2),
        )
    if task == "cnn":
        return nn.Sequential(
            nn.Conv2d(1, 8, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.MaxPool2d(kernel_size=2),
            nn.Conv2d(8, 12, kernel_size=3, padding=1),
            nn.ReLU(),
            nn.AdaptiveAvgPool2d((1, 1)),
            nn.Flatten(),
            nn.Linear(12, 3),
        )
    raise ValueError("task must be 'logistic', 'mlp', or 'cnn'")


def majority_baseline(train_targets: torch.Tensor, targets: torch.Tensor) -> float:
    """Return accuracy from always predicting the most common training class."""
    majority_class = int(torch.bincount(train_targets).argmax().item())
    return targets.eq(majority_class).float().mean().item()


def classification_metrics(
    logits: torch.Tensor, targets: torch.Tensor, class_count: int
) -> dict[str, float | torch.Tensor]:
    """Calculate accuracy, macro precision/recall/F1, and a confusion matrix."""
    if logits.shape != (targets.shape[0], class_count):
        raise ValueError("expected logits shaped (examples, classes)")
    predictions = logits.argmax(dim=1)
    confusion = torch.zeros((class_count, class_count), dtype=torch.long)
    for actual, predicted in zip(targets.tolist(), predictions.tolist()):
        confusion[actual, predicted] += 1
    true_positive = confusion.diag().float()
    predicted_positive = confusion.sum(dim=0).float()
    actual_positive = confusion.sum(dim=1).float()
    precision = true_positive / predicted_positive.clamp_min(1)
    recall = true_positive / actual_positive.clamp_min(1)
    f1 = 2 * precision * recall / (precision + recall).clamp_min(1e-12)
    return {
        "accuracy": predictions.eq(targets).float().mean().item(),
        "macro_precision": precision.mean().item(),
        "macro_recall": recall.mean().item(),
        "macro_f1": f1.mean().item(),
        "confusion_matrix": confusion,
    }


def _evaluate(
    model: nn.Module, split: DatasetSplit, class_count: int
) -> dict[str, float | torch.Tensor]:
    model.eval()
    with torch.no_grad():
        logits = model(split.features)
        loss = nn.functional.cross_entropy(logits, split.targets).item()
    result = classification_metrics(logits, split.targets, class_count)
    result["loss"] = loss
    return result


def _minibatches(
    count: int, batch_size: int, generator: torch.Generator
) -> list[torch.Tensor]:
    order = torch.randperm(count, generator=generator)
    return list(order.split(batch_size))


def _save_and_verify(
    model: nn.Module,
    task: str,
    checkpoint_path: Path,
    example_features: torch.Tensor,
) -> None:
    checkpoint_path.parent.mkdir(parents=True, exist_ok=True)
    torch.save({"task": task, "model_state_dict": model.state_dict()}, checkpoint_path)
    payload = torch.load(checkpoint_path, map_location="cpu", weights_only=True)
    restored = make_model(str(payload["task"]))
    restored.load_state_dict(payload["model_state_dict"])
    model.eval()
    restored.eval()
    with torch.no_grad():
        original_logits = model(example_features)
        restored_logits = restored(example_features)
    if not torch.equal(original_logits, restored_logits):
        raise AssertionError("checkpoint reload changed model logits")


def run(
    *,
    task: str,
    epochs: int | None = None,
    batch_size: int = 64,
    optimizer_name: str = "adam",
    learning_rate: float | None = None,
    weight_decay: float = 1e-4,
    checkpoint_path: Path | None = None,
    train_count: int = 900,
    validation_count: int = 300,
    test_count: int = 300,
) -> dict[str, float | torch.Tensor]:
    """Train one task, select by validation loss, then evaluate test data once."""
    if batch_size < 1:
        raise ValueError("batch_size must be at least 1")
    default_epochs = {"logistic": 35, "mlp": 55, "cnn": 30}
    if task not in default_epochs:
        raise ValueError("task must be 'logistic', 'mlp', or 'cnn'")
    epochs = default_epochs[task] if epochs is None else epochs
    if epochs < 1:
        raise ValueError("epochs must be at least 1")
    if optimizer_name not in {"sgd", "adam"}:
        raise ValueError("optimizer_name must be 'sgd' or 'adam'")

    torch.manual_seed(404)
    data = make_datasets(
        task,
        train_count=train_count,
        validation_count=validation_count,
        test_count=test_count,
    )
    model = make_model(task)
    rate = learning_rate
    if rate is None:
        rate = 0.08 if optimizer_name == "sgd" else 0.01
    if optimizer_name == "sgd":
        optimizer = torch.optim.SGD(
            model.parameters(), lr=rate, momentum=0.9, weight_decay=weight_decay
        )
    else:
        optimizer = torch.optim.Adam(
            model.parameters(), lr=rate, weight_decay=weight_decay
        )

    class_count = len(data.class_names)
    shuffle_generator = torch.Generator(device="cpu").manual_seed(505)
    best_validation_loss = float("inf")
    best_state: dict[str, torch.Tensor] | None = None
    first_training_loss: float | None = None
    final_training_loss = float("nan")

    for _ in range(epochs):
        model.train()
        total_loss = 0.0
        seen = 0
        for indices in _minibatches(
            data.train.targets.shape[0], batch_size, shuffle_generator
        ):
            optimizer.zero_grad(set_to_none=True)
            logits = model(data.train.features[indices])
            loss = nn.functional.cross_entropy(logits, data.train.targets[indices])
            loss.backward()
            optimizer.step()
            total_loss += loss.item() * indices.numel()
            seen += indices.numel()
        final_training_loss = total_loss / seen
        if first_training_loss is None:
            first_training_loss = final_training_loss
        validation = _evaluate(model, data.validation, class_count)
        validation_loss = float(validation["loss"])
        if validation_loss < best_validation_loss:
            best_validation_loss = validation_loss
            best_state = copy.deepcopy(model.state_dict())

    if best_state is None or first_training_loss is None:
        raise AssertionError("training did not produce a model state")
    model.load_state_dict(best_state)
    validation = _evaluate(model, data.validation, class_count)
    test = _evaluate(model, data.test, class_count)
    baseline = majority_baseline(data.train.targets, data.test.targets)

    confusion = test["confusion_matrix"]
    if not isinstance(confusion, torch.Tensor):
        raise AssertionError("confusion matrix must be a tensor")
    print(f"task: {task}")
    print(f"train feature shape: {tuple(data.train.features.shape)}")
    print(f"validation feature shape: {tuple(data.validation.features.shape)}")
    print(f"test feature shape: {tuple(data.test.features.shape)}")
    print(f"first/final epoch training loss: {first_training_loss:.4f} / {final_training_loss:.4f}")
    print(f"best validation loss: {best_validation_loss:.4f}")
    print(f"majority test baseline accuracy: {baseline:.3f}")
    print(f"test accuracy: {float(test['accuracy']):.3f}")
    print(f"test macro F1: {float(test['macro_f1']):.3f}")
    print(f"test confusion matrix:\n{confusion}")

    if checkpoint_path is not None:
        _save_and_verify(model, task, checkpoint_path, data.test.features[:8])
        print(f"checkpoint saved and verified: {checkpoint_path}")

    return {
        "initial_training_loss": first_training_loss,
        "final_training_loss": final_training_loss,
        "best_validation_loss": best_validation_loss,
        "validation_accuracy": float(validation["accuracy"]),
        "test_loss": float(test["loss"]),
        "test_accuracy": float(test["accuracy"]),
        "test_macro_precision": float(test["macro_precision"]),
        "test_macro_recall": float(test["macro_recall"]),
        "test_macro_f1": float(test["macro_f1"]),
        "baseline_accuracy": baseline,
        "confusion_matrix": confusion,
    }


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--task", choices=("logistic", "mlp", "cnn"), required=True)
    parser.add_argument("--epochs", type=int)
    parser.add_argument("--batch-size", type=int, default=64)
    parser.add_argument("--optimizer", choices=("sgd", "adam"), default="adam")
    parser.add_argument("--learning-rate", type=float)
    parser.add_argument("--weight-decay", type=float, default=1e-4)
    parser.add_argument("--checkpoint", type=Path)
    return parser.parse_args()


def main() -> None:
    arguments = parse_args()
    run(
        task=arguments.task,
        epochs=arguments.epochs,
        batch_size=arguments.batch_size,
        optimizer_name=arguments.optimizer,
        learning_rate=arguments.learning_rate,
        weight_decay=arguments.weight_decay,
        checkpoint_path=arguments.checkpoint,
    )


if __name__ == "__main__":
    main()
