#!/usr/bin/env python3
"""Train a one-input linear regression model on synthetic CPU data.

The target relationship is y = 3x - 2 plus small random noise. The validation
set is generated independently and is never used for optimizer updates.
"""

from __future__ import annotations

import argparse
from pathlib import Path

import torch
from torch import nn

torch.set_num_threads(1)


def linear_forward(x: torch.Tensor, weight: torch.Tensor, bias: torch.Tensor) -> torch.Tensor:
    """Return y_hat = weight * x + bias without using a PyTorch layer."""
    return weight * x + bias


def mse_gradients(
    x: torch.Tensor,
    target: torch.Tensor,
    weight: torch.Tensor,
    bias: torch.Tensor,
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]:
    """Return MSE, d(loss)/d(weight), and d(loss)/d(bias) by formula."""
    prediction = linear_forward(x, weight, bias)
    error = prediction - target
    loss = (error**2).mean()
    weight_gradient = (2 * error * x).mean()
    bias_gradient = (2 * error).mean()
    return loss, weight_gradient, bias_gradient


def make_dataset(
    count: int,
    seed: int,
    noise_standard_deviation: float = 0.15,
) -> tuple[torch.Tensor, torch.Tensor]:
    """Create an independent float32 dataset with shapes (count, 1)."""
    generator = torch.Generator(device="cpu").manual_seed(seed)
    x = torch.rand((count, 1), generator=generator, dtype=torch.float32) * 4 - 2
    noise = torch.randn((count, 1), generator=generator, dtype=torch.float32)
    target = 3 * x - 2 + noise_standard_deviation * noise
    return x, target


def validation_mse(model: nn.Module, x: torch.Tensor, target: torch.Tensor) -> float:
    """Measure MSE without changing parameters or building a backward graph."""
    model.eval()
    with torch.no_grad():
        return nn.functional.mse_loss(model(x), target).item()


def train_model(
    train_x: torch.Tensor,
    train_target: torch.Tensor,
    *,
    seed: int = 2026,
    epochs: int = 600,
    learning_rate: float = 0.05,
) -> tuple[nn.Linear, list[float]]:
    """Train nn.Linear(1, 1) using full-batch gradient descent on CPU."""
    torch.manual_seed(seed)
    model = nn.Linear(in_features=1, out_features=1)
    optimizer = torch.optim.SGD(model.parameters(), lr=learning_rate)
    history: list[float] = []

    model.train()
    for _ in range(epochs):
        optimizer.zero_grad()
        prediction = model(train_x)
        loss = nn.functional.mse_loss(prediction, train_target)
        loss.backward()
        optimizer.step()
        history.append(loss.item())

    return model, history


def save_and_verify_checkpoint(
    model: nn.Linear,
    checkpoint_path: Path,
    example_x: torch.Tensor,
) -> None:
    """Save a requested state_dict and verify that a fresh model restores it."""
    checkpoint_path.parent.mkdir(parents=True, exist_ok=True)
    torch.save({"model_state_dict": model.state_dict()}, checkpoint_path)

    restored = nn.Linear(1, 1)
    checkpoint = torch.load(checkpoint_path, map_location="cpu", weights_only=True)
    restored.load_state_dict(checkpoint["model_state_dict"])
    restored.eval()

    model.eval()
    with torch.no_grad():
        original_prediction = model(example_x)
        restored_prediction = restored(example_x)
    if not torch.equal(original_prediction, restored_prediction):
        raise AssertionError("restored model predictions do not match")


def run(epochs: int = 600, checkpoint_path: Path | None = None) -> dict[str, float]:
    """Run the complete experiment and return its measured values."""
    train_x, train_target = make_dataset(count=256, seed=10)
    validation_x, validation_target = make_dataset(count=128, seed=20)

    # A simple baseline always predicts the mean training target.
    baseline_prediction = train_target.mean().expand_as(validation_target)
    baseline_mse = nn.functional.mse_loss(
        baseline_prediction, validation_target
    ).item()

    model, history = train_model(train_x, train_target, epochs=epochs)
    model_mse = validation_mse(model, validation_x, validation_target)
    learned_weight = model.weight.item()
    learned_bias = model.bias.item()

    print(f"train shape: {tuple(train_x.shape)}, dtype: {train_x.dtype}")
    print(f"validation shape: {tuple(validation_x.shape)}")
    print(f"initial training MSE: {history[0]:.6f}")
    print(f"final training MSE: {history[-1]:.6f}")
    print(f"baseline validation MSE: {baseline_mse:.6f}")
    print(f"model validation MSE: {model_mse:.6f}")
    print(f"learned equation: y = {learned_weight:.4f}x + {learned_bias:.4f}")

    if checkpoint_path is not None:
        save_and_verify_checkpoint(model, checkpoint_path, validation_x[:8])
        print(f"checkpoint saved and verified: {checkpoint_path}")

    assert model_mse < baseline_mse * 0.05, "model did not beat the mean baseline"
    assert model_mse < 0.08, "validation error is unexpectedly high"
    assert abs(learned_weight - 3.0) < 0.15, "learned weight is not close to 3"
    assert abs(learned_bias + 2.0) < 0.15, "learned bias is not close to -2"

    return {
        "baseline_mse": baseline_mse,
        "model_mse": model_mse,
        "weight": learned_weight,
        "bias": learned_bias,
    }


def parse_args() -> argparse.Namespace:
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--epochs", type=int, default=600)
    parser.add_argument(
        "--checkpoint",
        type=Path,
        help="optional explicit path at which to save and verify a checkpoint",
    )
    return parser.parse_args()


if __name__ == "__main__":
    arguments = parse_args()
    run(epochs=arguments.epochs, checkpoint_path=arguments.checkpoint)
