"""Colored MNIST variant: suppress or replace colour-associated hidden components.

Run: python colored_mnist_intervention.py [--quick] [--out out-colored]
Install the adjacent requirements.txt first. This script is self-contained.

Ten-class variant, not the IRM binary setup. Each grayscale digit is rendered in one of ten
palette colours (3 x 784 = 2352 inputs). In training, digit d is drawn in colour d with
probability 0.9, otherwise a uniform random colour. Evaluation domains: matched (same rule),
independent (uniform colour), misleading (colour of digit d+1 mod 10 with probability 0.9).

Main path: image -> unchanged trained extractor (2352 -> 128 ReLU) -> explicit fixed layer
A = I - lambda U U^T on the hidden vector -> unchanged trained head (128 -> 10). U spans the
top singular directions of within-image hidden differences across all ten colours, estimated
on a calibration split. Rank and strength are selected on validation data only. Controls:
same-rank random subspaces, and projections onto null(W) / row(W) of the head, whose effects
on the logits are exact identities.

Replacement extension: transplant the selected components from the same grayscale image
rendered in each of the ten palette colours, without retraining or reselecting the subspace.
All-palette evaluation uses no label to choose donors. Aligned/shifted donor summaries use
known digit labels as diagnostic conditions, NOT as deployable prediction methods.
"""

from __future__ import annotations

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

import numpy as np
import torch
from torch import nn

import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt

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


def load_mnist(cache: Path):
    cache.mkdir(parents=True, exist_ok=True)
    path = cache / "mnist.npz"
    if not path.exists():
        print(f"Downloading {MNIST_URL}", flush=True)
        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()


def save_figure(fig, out, name):
    fig.savefig(out / name, dpi=170, bbox_inches="tight", facecolor="white")
    plt.close(fig)

PALETTE = torch.tensor([
    [1.0, 0.0, 0.0], [0.0, 1.0, 0.0], [0.0, 0.0, 1.0], [1.0, 1.0, 0.0], [1.0, 0.0, 1.0],
    [0.0, 1.0, 1.0], [1.0, 0.5, 0.0], [0.5, 0.0, 1.0], [0.5, 1.0, 0.0], [1.0, 1.0, 1.0],
])
PALETTE_NAMES = ("red", "green", "blue", "yellow", "magenta", "cyan", "orange", "purple", "lime", "white")
HIDDEN = 128
RANKS = (1, 2, 3, 4, 8, 16, 32)   # rank 3 added after seeing the 3-direction spectrum
STRENGTHS = (0.5, 1.0)
RANDOM_DRAWS = 3
TRAIN_DOMAIN = ("aligned", 0.9)
DOMAINS = {"matched": ("aligned", 0.9), "independent": ("aligned", 0.0), "misleading": ("shifted", 0.9)}
COLOUR_SEEDS = {"train": 1000, "calibration": 1001, "validation": 1002, "test": 1003}


def sample_colours(labels: np.ndarray, mapping: str, p: float, seed: int) -> np.ndarray:
    rng = np.random.default_rng(seed)
    base = labels if mapping == "aligned" else (labels + 1) % 10
    keep = rng.random(len(labels)) < p
    return np.where(keep, base, rng.integers(0, 10, len(labels))).astype(np.int64)


def render(x: torch.Tensor, colours: np.ndarray) -> torch.Tensor:
    c = PALETTE[torch.from_numpy(np.asarray(colours))]
    return (c[:, :, None] * x[:, None, :]).reshape(len(x), 3 * 784)


def build_model(seed: int) -> nn.Sequential:
    torch.manual_seed(seed)
    return nn.Sequential(nn.Linear(3 * 784, HIDDEN), nn.ReLU(), nn.Linear(HIDDEN, 10))


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


class FixedLinear(nn.Module):
    """An explicit, non-trainable linear layer on the hidden representation."""

    def __init__(self, matrix: torch.Tensor):
        super().__init__()
        self.register_buffer("matrix", matrix.clone())

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


@torch.inference_mode()
def hidden(model, x):
    return torch.cat([model[:2](batch) for batch in x.split(2048)])


@torch.inference_mode()
def head(model, h):
    return model[2](h)


def score(logits, y):
    pred = logits.argmax(1)
    correct = pred == y
    prob = logits.softmax(1)
    return {
        "accuracy": float(correct.double().mean()),
        "correct": int(correct.sum()),
        "total": int(len(y)),
        "mean_true_class_probability": float(prob[torch.arange(len(y)), y].double().mean()),
        "per_digit_accuracy": {str(d): float(correct[y == d].double().mean()) for d in range(10)},
    }, correct, pred


def transitions(before, after):
    return {"correct_to_wrong": int((before & ~after).sum()), "wrong_to_correct": int((~before & after).sum()),
            "correct_both": int((before & after).sum()), "wrong_both": int((~before & ~after).sum())}


def suppression(U: torch.Tensor, strength: float) -> torch.Tensor:
    return torch.eye(HIDDEN) - strength * (U @ U.T)


def colour_subspace(model, x_cal):
    stacks = torch.stack([hidden(model, render(x_cal, np.full(len(x_cal), j))) for j in range(10)], dim=1)
    differences = (stacks - stacks.mean(dim=1, keepdim=True)).reshape(-1, HIDDEN)
    _, singular, vt = torch.linalg.svd(differences, full_matrices=False)
    return vt.T.contiguous(), singular


def random_subspace(rank: int, seed: int) -> torch.Tensor:
    q, _ = torch.linalg.qr(torch.randn(HIDDEN, rank, generator=torch.Generator().manual_seed(seed)))
    return q


def evaluate(model, matrix, rendered: dict, labels, baseline_correct=None):
    out = {}
    for domain, x in rendered.items():
        h = hidden(model, x)
        logits = head(model, h if matrix is None else h @ matrix.T)
        result, correct, _ = score(logits, labels)
        if baseline_correct is not None:
            result["transitions_vs_baseline"] = transitions(baseline_correct[domain], correct)
        out[domain] = result
    return out


def replace_components(h, donor_h, U):
    """Keep recipient's perpendicular component; copy donor's selected coordinates."""
    return h + ((donor_h - h) @ U) @ U.T


@torch.inference_mode()
def evaluate_replacement(model, U, test_x, test_y, rendered, seed, out, export=False):
    donor_hidden = [hidden(model, render(test_x, np.full(len(test_x), c))) for c in range(10)]
    donor_logits = torch.stack([head(model, h) for h in donor_hidden], dim=1)
    random_bases = [random_subspace(U.shape[1], 500 + 10 * seed + draw) for draw in range(RANDOM_DRAWS)]
    indices = [int(torch.where(test_y == d)[0][0]) for d in range(10)]
    samples = [{"testIndex": i, "digit": int(test_y[i]),
                "pixels": (test_x[i] * 255).round().to(torch.uint8).tolist(), "domains": {}} for i in indices]
    checks = {"self_patch_logits_max_abs_diff": 0.0, "selected_coordinates_max_abs_diff": 0.0,
              "complementary_logit_effects_max_abs_diff": 0.0,
              "perpendicular_component_max_abs_diff": 0.0, "full_patch_vs_donor_logits_max_abs_diff": 0.0}

    def summary(logits, baseline, donor_colours):
        # logits is N x K x 10; each source image has K donor conditions.
        k = logits.shape[1]
        y = test_y[:, None].expand(-1, k).reshape(-1)
        result, correct, pred = score(logits.reshape(-1, 10), y)
        base = baseline[:, None, :].expand(-1, k, -1).reshape(-1, 10)
        base_pred = base.argmax(1)
        donor_colours = donor_colours.reshape(-1)
        rows = torch.arange(len(y))
        reference = donor_logits[torch.arange(len(test_y))[:, None], donor_colours.reshape(len(test_y), k)].reshape(-1, 10)
        result.update({"unique_images": len(test_y), "donors_per_image": k,
                       "changed_predictions": int((pred != base_pred).sum()),
                       "transitions_vs_baseline": transitions(base_pred == y, correct),
                       "mean_donor_class_probability": float(logits.reshape(-1, 10).softmax(1)[rows, donor_colours].double().mean()),
                       "baseline_mean_donor_class_probability": float(base.softmax(1)[rows, donor_colours].double().mean()),
                       "prediction_agreement_with_recoloring": float((pred == reference.argmax(1)).double().mean()),
                       "mean_abs_logit_difference_from_recoloring": float((logits.reshape(-1, 10) - reference).abs().double().mean())})
        return result

    policies = {"aligned_donor": test_y, "shifted_donor": (test_y + 1) % 10}
    all_colours = torch.arange(10).expand(len(test_y), -1)
    rows = torch.arange(len(test_y))
    results = {}
    for domain, x in rendered.items():
        h = hidden(model, x)
        baseline = head(model, h)
        source_colours = (x.reshape(-1, 3, 784).amax(2)[:, None, :] - PALETTE[None, :, :]).abs().sum(2).argmin(1)
        patched_logits, complement_logits = [], []
        for colour, donor_h in enumerate(donor_hidden):
            patched = replace_components(h, donor_h, U)
            delta = patched - h
            checks["selected_coordinates_max_abs_diff"] = max(checks["selected_coordinates_max_abs_diff"],
                float((patched @ U - donor_h @ U).abs().max()))
            checks["perpendicular_component_max_abs_diff"] = max(checks["perpendicular_component_max_abs_diff"],
                float((delta - (delta @ U) @ U.T).abs().max()))
            logits = head(model, patched)
            complement = head(model, donor_h - ((donor_h - h) @ U) @ U.T)
            checks["complementary_logit_effects_max_abs_diff"] = max(checks["complementary_logit_effects_max_abs_diff"],
                float((logits + complement - baseline - donor_logits[:, colour]).abs().max()))
            same = source_colours == colour
            if same.any():
                checks["self_patch_logits_max_abs_diff"] = max(checks["self_patch_logits_max_abs_diff"],
                    float((logits[same] - baseline[same]).abs().max()))
                assert torch.equal(logits[same].argmax(1), baseline[same].argmax(1))
            # Replacing all 128 coordinates must reproduce the donor forward pass.
            full = replace_components(h, donor_h, torch.eye(HIDDEN))
            checks["full_patch_vs_donor_logits_max_abs_diff"] = max(checks["full_patch_vs_donor_logits_max_abs_diff"],
                float((head(model, full) - donor_logits[:, colour]).abs().max()))
            patched_logits.append(logits)
            complement_logits.append(complement)
        patched_logits = torch.stack(patched_logits, dim=1)
        complement_logits = torch.stack(complement_logits, dim=1)
        full_removed = head(model, h @ suppression(U, 1.0).T)
        records = {"all_palette": summary(patched_logits, baseline, all_colours)}
        for name, colours in policies.items():
            records[name] = summary(patched_logits[rows, colours][:, None, :], baseline, colours[:, None])
        for name, logits in (("complement_patch", complement_logits), ("recolored_reference", donor_logits)):
            records[name] = {"all_palette": summary(logits, baseline, all_colours)}
            for policy, colours in policies.items():
                records[name][policy] = summary(logits[rows, colours][:, None, :], baseline, colours[:, None])
        records["random_controls"] = []
        for basis in random_bases:
            random_logits = torch.stack([head(model, replace_components(h, donor, basis)) for donor in donor_hidden], dim=1)
            random_record = {"all_palette": summary(random_logits, baseline, all_colours)}
            for name, colours in policies.items():
                random_record[name] = summary(random_logits[rows, colours][:, None, :], baseline, colours[:, None])
            records["random_controls"].append(random_record)
        results[domain] = records

        def pack(logits):
            return {"logits": logits.tolist(), "probabilities": logits.softmax(-1).tolist(),
                    "prediction": int(logits.argmax())}

        for sample, i in zip(samples, indices):
            sample["domains"][domain] = {
                "sourceColourIndex": int(source_colours[i]), "original": pack(baseline[i]),
                "suppressed": pack(full_removed[i]),
                "sourceCoordinates": (h[i] @ U).tolist(),
                "donors": [{"colourIndex": c, "coordinates": (donor_hidden[c][i] @ U).tolist(),
                            "patched": pack(patched_logits[i, c]), "complement": pack(complement_logits[i, c]),
                            "recolored": pack(donor_logits[i, c])}
                           for c in range(10)],
            }
    assert all(value < 2e-4 for value in checks.values()), checks
    if export:
        (out / "replacement-samples.json").write_text(json.dumps(
            {"seed": seed, "rank": U.shape[1], "palette": PALETTE.tolist(), "paletteNames": PALETTE_NAMES,
             "samples": samples}, separators=(",", ":")) + "\n")
    return {"rank": U.shape[1], "test": results, "checks": checks}


def run_seed(seed, args, data, out):
    train_x, train_y, cal_x, val_x, val_y, test_x, test_y = data
    model = build_model(seed)
    train_rendered = render(train_x, sample_colours(train_y.numpy(), *TRAIN_DOMAIN, COLOUR_SEEDS["train"]))
    train(model, train_rendered, train_y, seed, args.epochs, args.batch_size, args.lr)
    for p in model.parameters():
        p.requires_grad_(False)
    snapshot = {k: v.detach().clone() for k, v in model.state_dict().items()}

    val_rendered = {d: render(val_x, sample_colours(val_y.numpy(), *spec, COLOUR_SEEDS["validation"] + i))
                    for i, (d, spec) in enumerate(DOMAINS.items())}
    test_rendered = {d: render(test_x, sample_colours(test_y.numpy(), *spec, COLOUR_SEEDS["test"] + i))
                     for i, (d, spec) in enumerate(DOMAINS.items())}

    # Baseline (no inserted layer) and the identity check at lambda = 0.
    baseline_val = evaluate(model, None, val_rendered, val_y)
    baseline_test = evaluate(model, None, test_rendered, test_y)
    baseline_correct = {}
    for d, x in test_rendered.items():
        _, correct, _ = score(head(model, hidden(model, x)), test_y)
        baseline_correct[d] = correct

    U_all, singular = colour_subspace(model, cal_x[: args.calibration])
    spectrum = (singular ** 2 / (singular ** 2).sum()).tolist()

    identity_check = evaluate(model, suppression(U_all[:, :4], 0.0), test_rendered, test_y)
    assert all(identity_check[d]["correct"] == baseline_test[d]["correct"] for d in DOMAINS)

    # Validation grid, prespecified objective: mean accuracy over the three domains.
    grid = []
    for rank in RANKS:
        U = U_all[:, :rank]
        assert torch.allclose(U.T @ U, torch.eye(rank), atol=1e-5)
        for strength in STRENGTHS:
            result = evaluate(model, suppression(U, strength), val_rendered, val_y)
            grid.append({"rank": rank, "strength": strength, "kind": "colour",
                         "validation": {d: r["accuracy"] for d, r in result.items()},
                         "objective": float(np.mean([r["accuracy"] for r in result.values()]))})
        for draw in range(RANDOM_DRAWS):
          for strength in STRENGTHS:
            result = evaluate(model, suppression(random_subspace(rank, 500 + 10 * seed + draw), strength),
                              val_rendered, val_y)
            grid.append({"rank": rank, "strength": strength, "kind": f"random_{draw}",
                         "validation": {d: r["accuracy"] for d, r in result.items()},
                         "objective": float(np.mean([r["accuracy"] for r in result.values()]))})
    colour_entries = [g for g in grid if g["kind"] == "colour"]
    selected = max(colour_entries, key=lambda g: g["objective"])
    rank, strength = selected["rank"], selected["strength"]
    U = U_all[:, :rank]
    A = suppression(U, strength)
    P = U @ U.T
    assert torch.allclose(P @ P, P, atol=1e-5)

    # Test evaluation through an explicit inserted layer, checked against direct computation.
    wrapped = nn.Sequential(model[0], model[1], FixedLinear(A), model[2]).eval()
    with torch.inference_mode():
        wrapped_logits = torch.cat([wrapped(b) for b in test_rendered["misleading"].split(2048)])
        direct_logits = head(model, hidden(model, test_rendered["misleading"]) @ A.T)
    wrapped_max_diff = float((wrapped_logits - direct_logits).abs().max())
    assert wrapped_max_diff < 1e-4

    selected_test = evaluate(model, A, test_rendered, test_y, baseline_correct)
    random_tests = [evaluate(model, suppression(random_subspace(rank, 500 + 10 * seed + draw), strength),
                             test_rendered, test_y, baseline_correct) for draw in range(RANDOM_DRAWS)]
    full_test = evaluate(model, suppression(U_all[:, :rank], 1.0), test_rendered, test_y, baseline_correct) \
        if strength != 1.0 else selected_test

    # Destructive control from the head: W+ W projects onto row(W); I - W+ W onto null(W).
    W = model[2].weight.detach()
    b = model[2].bias.detach()
    P_row = torch.linalg.pinv(W) @ W
    h_test = hidden(model, test_rendered["matched"])
    null_logits = head(model, h_test @ (torch.eye(HIDDEN) - P_row).T)
    row_logits = head(model, h_test @ P_row.T)
    null_max_diff = float((null_logits - b).abs().max())
    row_max_diff = float((row_logits - head(model, h_test)).abs().max())
    null_test = evaluate(model, torch.eye(HIDDEN) - P_row, test_rendered, test_y, baseline_correct)
    row_test = evaluate(model, P_row, test_rendered, test_y, baseline_correct)

    replacement = evaluate_replacement(model, U, test_x, test_y, test_rendered, seed, out,
                                       export=seed == args.seeds[0])
    assert all(torch.equal(snapshot[k], v) for k, v in model.state_dict().items())
    print(f"seed {seed}: baseline " + " ".join(f"{d} {r['accuracy']:.2%}" for d, r in baseline_test.items())
          + f" | selected rank {rank} λ {strength}: "
          + " ".join(f"{d} {r['accuracy']:.2%}" for d, r in selected_test.items())
          + f" | random r={rank}: " + " ".join(
              f"{d} {np.mean([t[d]['accuracy'] for t in random_tests]):.2%}" for d in DOMAINS)
          + f" | full projection r={rank} λ=1: " + " ".join(f"{d} {r['accuracy']:.2%}" for d, r in full_test.items())
          + f" | null(W): matched {null_test['matched']['accuracy']:.2%}", flush=True)

    if seed == args.seeds[0]:
        torch.save(model.state_dict(), out / f"model_seed{seed}.pt")
        torch.save({"U": U_all, "singular_values": singular, "selected_rank": rank,
                    "selected_strength": strength}, out / f"colour_subspace_seed{seed}.pt")
        export_widget_samples(model, A, test_x, test_y, test_rendered, out)
        plot_examples(test_x, test_y, test_rendered, out)

    return {
        "seed": seed,
        "baseline_validation": {d: r["accuracy"] for d, r in baseline_val.items()},
        "baseline_test": baseline_test,
        "spectrum_variance_fraction": spectrum[:32],
        "validation_grid": grid,
        "selected": {"rank": rank, "strength": strength, "validation_objective": selected["objective"]},
        "selected_test": selected_test,
        "selected_test_full_strength": full_test,
        "random_control_test": random_tests,
        "null_space_control_test": null_test,
        "row_space_identity_test": row_test,
        "replacement": replacement,
        "checks": {
            "learned_tensors_identical": True,
            "identity_at_zero_strength_matches_baseline": True,
            "U_orthonormal": True, "UUT_idempotent": True,
            "wrapped_vs_direct_logits_max_abs_diff": wrapped_max_diff,
            "null_projection_logits_minus_bias_max_abs": null_max_diff,
            "row_projection_logits_max_abs_diff": row_max_diff,
            "head_rank": int(torch.linalg.matrix_rank(W)),
            "row_space_dim": int(torch.linalg.matrix_rank(P_row)),
            "null_space_dim": HIDDEN - int(torch.linalg.matrix_rank(P_row)),
        },
    }


def export_widget_samples(model, A, test_x, test_y, test_rendered, out):
    y = test_y.numpy()
    indices = [int(np.flatnonzero(y == d)[0]) for d in range(10)]
    samples = []
    for i in indices:
        entry = {"testIndex": i, "digit": int(y[i]), "pixels": (test_x[i] * 255).round().to(torch.uint8).tolist(),
                 "domains": {}}
        for d, x in test_rendered.items():
            colour = x[i].reshape(3, 784).amax(dim=1)
            colour_index = int((PALETTE - colour).abs().sum(1).argmin())
            h = hidden(model, x[i:i + 1])
            entry["domains"][d] = {"colourIndex": colour_index, "colourName": PALETTE_NAMES[colour_index],
                                   "rgb": PALETTE[colour_index].tolist(),
                                   "predictedBefore": int(head(model, h).argmax()),
                                   "predictedAfter": int(head(model, h @ A.T).argmax())}
        samples.append(entry)
    (out / "widget-samples.json").write_text(json.dumps(
        {"palette": PALETTE.tolist(), "paletteNames": PALETTE_NAMES, "samples": samples},
        separators=(",", ":")) + "\n")


def plot_examples(test_x, test_y, test_rendered, out):
    y = test_y.numpy()
    indices = [int(np.flatnonzero(y == d)[0]) for d in range(10)]
    fig, axes = plt.subplots(3, 10, figsize=(11, 3.6), layout="constrained")
    for row, (domain, x) in enumerate(test_rendered.items()):
        for col, i in enumerate(indices):
            ax = axes[row, col]
            ax.imshow(x[i].reshape(3, 28, 28).permute(1, 2, 0).clamp(0, 1))
            ax.set_xticks([]); ax.set_yticks([])
            if col == 0:
                ax.set_ylabel(domain, fontsize=9)
            if row == 0:
                ax.set_title(f"digit {y[i]}", fontsize=9)
    fig.suptitle("The same test images under the three colour rules", fontsize=10)
    save_figure(fig, out, "colored_examples.png")


def plot_summary(per_seed, out):
    first = per_seed[0]
    fig, ax = plt.subplots(figsize=(6, 3.4), layout="constrained")
    ax.bar(range(1, 33), first["spectrum_variance_fraction"], color="#3465a4")
    ax.set_xlabel("Singular direction of within-image colour differences")
    ax.set_ylabel("Fraction of variance"); ax.set_title(f"Colour-difference spectrum (seed {first['seed']})")
    save_figure(fig, out, "colored_spectrum.png")

    fig, ax = plt.subplots(figsize=(7.5, 4), layout="constrained")
    colours = {"matched": "#3465a4", "independent": "#4e9a06", "misleading": "#d97732"}
    for domain, colour in colours.items():
        for strength, style in ((1.0, "-"), (0.5, "--")):
            ys = [100 * g["validation"][domain] for g in first["validation_grid"]
                  if g["kind"] == "colour" and g["strength"] == strength]
            ax.plot(RANKS, ys, style, marker="o", color=colour, label=f"{domain}, λ={strength}")
        rand = [100 * np.mean([g["validation"][domain] for g in first["validation_grid"]
                               if g["kind"].startswith("random") and g["rank"] == r and g["strength"] == 1.0])
                for r in RANKS]
        ax.plot(RANKS, rand, ":", marker="x", color=colour, alpha=0.7, label=f"{domain}, random r, λ=1")
        ax.axhline(100 * first["baseline_validation"][domain], color=colour, linewidth=0.8, alpha=0.5)
    ax.set_xscale("log", base=2); ax.set_xticks(RANKS, [str(r) for r in RANKS])
    ax.set_xlabel("Suppressed rank r"); ax.set_ylabel("Validation accuracy (%)")
    ax.set_title(f"Validation grid (seed {first['seed']}); thin lines = baseline")
    ax.legend(fontsize=7, ncol=3, loc="lower left"); ax.set_ylim(0, 105)
    save_figure(fig, out, "colored_grid.png")

    fig, ax = plt.subplots(figsize=(10, 4.4), layout="constrained")
    conditions = (("baseline_test", "Frozen network, no insertion"),
                  ("selected_test", "Colour directions suppressed (selected r, λ)"),
                  ("selected_test_full_strength", "Colour directions projected out (same r, λ=1)"),
                  ("random_control_test", "Random directions, same r and λ"),
                  ("null_space_control_test", "Row space of the head removed"))
    bar_colours = ("#3465a4", "#d97732", "#c4a000", "#888a85", "#a40000")
    positions = np.arange(len(DOMAINS))
    for k, ((key, label), colour) in enumerate(zip(conditions, bar_colours)):
        means, sds = [], []
        for d in DOMAINS:
            vals = []
            for s in per_seed:
                if key == "random_control_test":
                    vals.append(100 * np.mean([t[d]["accuracy"] for t in s[key]]))
                else:
                    vals.append(100 * s[key][d]["accuracy"])
            means.append(np.mean(vals)); sds.append(np.std(vals))
        bars = ax.bar(positions + (k - 2) * 0.17, means, width=0.16, yerr=sds if len(per_seed) > 1 else None,
                      capsize=3, label=label, color=colour)
        ax.bar_label(bars, fmt="%.1f", fontsize=7, padding=2)
    ax.set_xticks(positions, [f"{d}\n({DOMAINS[d][0]}, p={DOMAINS[d][1]})" for d in DOMAINS])
    ax.set_ylabel("Test accuracy (%)"); ax.set_ylim(0, 112)
    ax.set_title(f"Mean over {len(per_seed)} training seed(s); error bars = sd")
    ax.legend(fontsize=8, loc="upper center", bbox_to_anchor=(0.5, -0.22), ncol=3)
    save_figure(fig, out, "colored_results.png")


def main():
    parser = argparse.ArgumentParser(description=__doc__)
    parser.add_argument("--quick", action="store_true")
    parser.add_argument("--seeds", type=int, nargs="+", default=None)
    parser.add_argument("--epochs", type=int, default=None)
    parser.add_argument("--batch-size", type=int, default=512)
    parser.add_argument("--lr", type=float, default=0.002)
    parser.add_argument("--calibration", type=int, default=None, help="images in the counterfactual set")
    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, default=None)
    args = parser.parse_args()
    args.seeds = args.seeds or ([42] if args.quick else [42, 43, 44])
    args.epochs = args.epochs if args.epochs is not None else (2 if args.quick else 8)
    args.calibration = args.calibration or (500 if args.quick else 2000)
    out = args.out or HERE / ("out-colored-quick" if args.quick else "out-colored")
    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)
    order = np.random.default_rng(0).permutation(len(train_images))
    n_train, n_cal, n_val = (5000, 1000, 1000) if args.quick else (50000, 5000, 5000)
    tr, ca, va = order[:n_train], order[n_train:n_train + n_cal], order[n_train + n_cal:n_train + n_cal + n_val]
    if args.quick:
        test_images, test_labels = test_images[:2000], test_labels[:2000]
    to_x = lambda a: torch.from_numpy(a.reshape(-1, 784).astype(np.float32) / 255)
    to_y = lambda a: torch.from_numpy(a.astype(np.int64))
    data = (to_x(train_images[tr]), to_y(train_labels[tr]), to_x(train_images[ca]),
            to_x(train_images[va]), to_y(train_labels[va]), to_x(test_images), to_y(test_labels))
    print(f"Splits: train {n_train:,} / calibration {n_cal:,} / validation {n_val:,} / test {len(test_images):,}",
          flush=True)

    per_seed = [run_seed(seed, args, data, out) for seed in args.seeds]
    plot_summary(per_seed, out)

    def aggregate(key):
        table = {}
        for d in DOMAINS:
            vals = [(np.mean([t[d]["accuracy"] for t in s[key]]) if key == "random_control_test"
                     else s[key][d]["accuracy"]) for s in per_seed]
            table[d] = {"mean": float(np.mean(vals)), "sd": float(np.std(vals)), "per_seed": [float(v) for v in vals]}
        return table

    payload = {
        "config": {"variant": "10-class Colored MNIST variant (not IRM binary)", "palette": PALETTE.tolist(),
                   "palette_names": PALETTE_NAMES, "train_domain": TRAIN_DOMAIN, "domains": DOMAINS,
                   "colour_seeds": COLOUR_SEEDS, "split_seed": 0,
                   "colour_rule": "with probability p use the mapping colour, else a uniform colour from all 10; "
                                  "so the designated colour appears with probability p + (1 - p)/10",
                   "effective_designated_colour_probability": {d: p + (1 - p) / 10 for d, (_, p) in DOMAINS.items()},
                   "rank_candidates_note": "rank 3 was added to the candidate set after inspecting the spectrum "
                                           "(three dominant directions); selection still uses validation only",
                   "splits": {"train": n_train, "calibration": n_cal, "validation": n_val, "test": len(test_images)},
                   "calibration_images_in_counterfactual_set": args.calibration,
                   "model": "2352 -> 128 ReLU -> 10", "epochs": args.epochs, "batch_size": args.batch_size,
                   "learning_rate": args.lr, "training_seeds": args.seeds, "ranks": RANKS,
                   "strengths": STRENGTHS, "random_draws": RANDOM_DRAWS,
                   "selection": "validation only; objective = mean accuracy over the three domains",
                   "dataset_sha256": data_hash, "torch": torch.__version__, "quick": args.quick},
        "definitions": {
            "colour_subspace": "top right-singular vectors of within-image hidden differences over all 10 "
                               "colours on calibration images; an ESTIMATED colour-associated subspace",
            "suppression": "A = I - strength * U U^T inserted between the frozen extractor and frozen head",
            "null_space_control": "A = I - W+W removes everything the head reads; logits become the bias",
            "row_space_identity": "A = W+W keeps everything the head reads; logits unchanged",
            "replacement": "h_patch = h + ((h_donor - h) U) U^T; donor is the SAME source image recolored. "
                           "Uses the suppression-selected rank without further tuning. All ten donor colours, "
                           "including the source colour, are evaluated equally. Aligned/shifted donors use "
                           "known labels as diagnostic conditions, not deployable accuracy improvements. "
                           "Random controls use the same rank, full replacement, and identical donors.",
        },
        "aggregate_test": {k: aggregate(k) for k in ("baseline_test", "selected_test", "selected_test_full_strength",
                                                     "random_control_test",
                                                     "null_space_control_test", "row_space_identity_test")},
        "per_seed": per_seed,
        "elapsed_seconds": time.perf_counter() - started,
    }
    (out / "results.json").write_text(json.dumps(payload, indent=2) + "\n")
    print(f"\nWrote {out.resolve()} in {payload['elapsed_seconds']:.1f}s", flush=True)


if __name__ == "__main__":
    main()
