"""Project a trained MNIST network's hidden activations onto row(W) or null(W).

Self-contained. Run: python classifier_subspaces.py [--quick] [--out DIRECTORY]
Train in float32, then cast the frozen network and inputs to float64 once for all
three evaluations. No learned values change between intervention conditions.
"""
from __future__ import annotations

import argparse
import hashlib
import json
import time
import urllib.request
from pathlib import Path

import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt
import numpy as np
import torch
from torch import nn

HERE = Path(__file__).resolve().parent
MNIST_URL = "https://storage.googleapis.com/tensorflow/tf-keras-datasets/mnist.npz"


def load_mnist(cache):
    cache.mkdir(parents=True, exist_ok=True)
    path = cache / "mnist.npz"
    if not path.exists():
        temporary = path.with_suffix(".download")
        urllib.request.urlretrieve(MNIST_URL, temporary)
        temporary.replace(path)
    with np.load(path, allow_pickle=False) as data:
        arrays = [data[key].copy() for key in ("x_train", "y_train", "x_test", "y_test")]
    return arrays, hashlib.sha256(path.read_bytes()).hexdigest()


class FixedProjection(nn.Module):
    def __init__(self, matrix):
        super().__init__()
        self.register_buffer("matrix", matrix.clone())

    def forward(self, hidden):
        return hidden @ self.matrix.T


def train(x, y, args):
    torch.manual_seed(args.seed)
    model = nn.Sequential(nn.Linear(784, 128), nn.ReLU(), nn.Linear(128, 10))
    optimizer = torch.optim.Adam(model.parameters(), lr=args.lr)
    generator = torch.Generator().manual_seed(args.seed + 1)
    for _ in range(args.epochs):
        for indices in torch.randperm(len(x), generator=generator).split(args.batch_size):
            optimizer.zero_grad(set_to_none=True)
            nn.functional.cross_entropy(model(x[indices]), y[indices]).backward()
            optimizer.step()
    return model.eval()


@torch.inference_mode()
def logits(model, x):
    return torch.cat([model(batch) for batch in x.split(1024)])


def summarize(z, y, baseline):
    probabilities = z.softmax(1)
    predicted = z.argmax(1)
    return {
        "accuracy": float((predicted == y).double().mean()),
        "correct": int((predicted == y).sum()), "total": len(y),
        "changed_predictions": int((predicted != baseline.argmax(1)).sum()),
        "distinct_predicted_digits": predicted.unique().tolist(),
        "max_logit_difference_from_original": float((z - baseline).abs().max()),
        "max_probability_difference_from_original": float((probabilities - baseline.softmax(1)).abs().max()),
        "prediction_sha256": hashlib.sha256(predicted.numpy().tobytes()).hexdigest(),
    }


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--quick", action="store_true")
    parser.add_argument("--seed", type=int, default=42)
    parser.add_argument("--epochs", type=int)
    parser.add_argument("--batch-size", type=int, default=512)
    parser.add_argument("--lr", type=float, default=0.002)
    parser.add_argument("--threads", type=int, default=2)
    parser.add_argument("--cache", type=Path, default=Path.home() / ".cache/mirror-span")
    parser.add_argument("--out", type=Path)
    args = parser.parse_args()
    args.epochs = args.epochs or (2 if args.quick else 8)
    out = args.out or HERE / ("out-classifier-quick" if args.quick else "out-classifier")
    out.mkdir(parents=True, exist_ok=True)
    torch.set_num_threads(args.threads)
    torch.use_deterministic_algorithms(True)
    started = time.perf_counter()
    (train_images, train_labels, test_images, test_labels), data_hash = load_mnist(args.cache)
    if args.quick:
        train_images, train_labels = train_images[:2000], train_labels[:2000]
        test_images, test_labels = test_images[:1000], test_labels[:1000]
    train_x = torch.from_numpy(train_images.reshape(-1, 784).astype(np.float32) / 255)
    test_x = torch.from_numpy(test_images.reshape(-1, 784).astype(np.float32) / 255)
    train_y = torch.from_numpy(train_labels.astype(np.int64))
    test_y = torch.from_numpy(test_labels.astype(np.int64))
    print(f"Training on {len(train_x):,} original images; evaluating {len(test_x):,}.", flush=True)
    model = train(train_x, train_y, args)
    float32_predictions = logits(model, test_x).argmax(1)
    model = model.double().requires_grad_(False)
    test_x = test_x.double()
    snapshot = {key: value.clone() for key, value in model.state_dict().items()}
    W, bias = model[2].weight, model[2].bias
    _, singular, vh = torch.linalg.svd(W, full_matrices=False)
    tolerance = float(max(W.shape) * torch.finfo(W.dtype).eps * singular.max())
    rank = int((singular > tolerance).sum())
    q = vh[:rank].T
    row = q @ q.T
    null = torch.eye(128, dtype=torch.float64) - row
    matrices = {"original": torch.eye(128, dtype=torch.float64), "row": row, "null": null}
    scores = {}
    baseline = logits(model, test_x)
    for name, matrix in matrices.items():
        wrapped = nn.Sequential(model[:2], FixedProjection(matrix), model[2]).eval()
        scores[name] = logits(wrapped, test_x)
    conditions = {name: summarize(z, test_y, baseline) for name, z in scores.items()}
    geometry = {
        "head_rank": rank, "row_space_dimension": rank, "null_space_dimension": 128 - rank,
        "svd_tolerance": tolerance, "singular_values": singular.tolist(),
        "row_idempotence_max_abs": float((row @ row - row).abs().max()),
        "null_idempotence_max_abs": float((null @ null - null).abs().max()),
        "row_null_product_max_abs": float((row @ null).abs().max()),
        "W_row_minus_W_max_abs": float((W @ row - W).abs().max()),
        "W_null_max_abs": float((W @ null).abs().max()),
        "null_logits_minus_bias_max_abs": float((scores["null"] - bias).abs().max()),
        "learned_tensors_identical": all(torch.equal(snapshot[key], value) for key, value in model.state_dict().items()),
        "changed_predictions_after_float64_cast": int((float32_predictions != baseline.argmax(1)).sum()),
    }
    assert geometry["learned_tensors_identical"]
    assert torch.allclose(scores["row"], baseline, atol=1e-10, rtol=0)
    assert torch.allclose(scores["null"], bias.expand_as(baseline), atol=1e-10, rtol=0)
    assert torch.equal(scores["row"].argmax(1), baseline.argmax(1))
    assert torch.equal(scores["null"].argmax(1), torch.full_like(test_y, int(bias.argmax())))
    for key in ("row_idempotence_max_abs", "null_idempotence_max_abs", "row_null_product_max_abs"):
        assert geometry[key] < 1e-10
    selection = "First three original test images of each digit, in test-index order within each digit."
    samples = []
    for digit in range(10):
        for index in np.flatnonzero(test_labels == digit)[:3]:
            index = int(index)
            samples.append({"testIndex": index, "digit": digit,
                            "pixels": test_images[index].reshape(-1).tolist(),
                            "conditions": {name: {"logits": z[index].tolist(),
                                                   "probabilities": z[index].softmax(0).tolist(),
                                                   "prediction": int(z[index].argmax())}
                                           for name, z in scores.items()}})
    (out / "widget-samples.json").write_text(json.dumps({"selection": selection, "rank": rank,
            "nullity": 128 - rank, "samples": samples}, separators=(",", ":")) + "\n")
    torch.save(model.state_dict(), out / "model_seed42.pt" if args.seed == 42 else out / f"model_seed{args.seed}.pt")
    labels = ("Original", "Keep row space", "Keep null space")
    colors = ("#3465a4", "#008b80", "#cf5565")
    fig, axes = plt.subplots(1, 2, figsize=(9, 3.8), layout="constrained")
    bars = axes[0].bar(labels, [100 * conditions[name]["accuracy"] for name in matrices], color=colors)
    axes[0].bar_label(bars, fmt="%.2f%%"); axes[0].set_ylim(0, 110); axes[0].set_ylabel("Test accuracy (%)")
    bars = axes[1].bar(labels, [128, rank, 128 - rank], color=colors)
    axes[1].bar_label(bars); axes[1].set_ylim(0, 145); axes[1].set_ylabel("Independent retained directions")
    fig.suptitle("Same trained grayscale MNIST network, same 10,000 test images" if not args.quick else "Quick smoke test")
    fig.savefig(out / "accuracy_dimensions.png", dpi=170, bbox_inches="tight", facecolor="white")
    plt.close(fig)
    payload = {
        "config": {"seed": args.seed, "epochs": args.epochs, "learning_rate": args.lr,
                   "batch_size": args.batch_size, "threads": args.threads, "quick": args.quick,
                   "train_images": len(train_x), "test_images": len(test_x), "architecture": "784 -> 128 ReLU -> 10",
                   "optimizer": "Adam", "training_dtype": "float32", "evaluation_dtype": "float64",
                   "dataset_sha256": data_hash, "torch": torch.__version__, "numpy": np.__version__},
        "definitions": {"intervention_location": "After hidden ReLU, before final Linear; no activation after projection.",
                        "row": "Orthogonal projection onto row space of the trained output weights.",
                        "null": "Orthogonal projection onto null space of the trained output weights.",
                        "sample_selection": selection},
        "conditions": conditions, "checks": geometry,
        "bias": bias.tolist(), "constant_null_prediction": int(bias.argmax()),
        "elapsed_seconds": time.perf_counter() - started,
    }
    (out / "results.json").write_text(json.dumps(payload, indent=2) + "\n")
    print(json.dumps({"conditions": conditions, "checks": geometry}, indent=2), flush=True)
    print(f"Wrote {out.resolve()} in {payload['elapsed_seconds']:.1f}s", flush=True)


if __name__ == "__main__":
    main()
