#!/usr/bin/env python3
"""Build the article's figures and chart data from paper-final raw results.

The saved analysis summary is used only as a consistency check. No chart reads it.
Run with /usr/bin/python3 scripts/build_article_data.py from the Space directory.
"""
from __future__ import annotations

import hashlib
import html
import json
import os
from pathlib import Path

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

SPACE = Path(__file__).resolve().parents[1]
PROJECT = SPACE.parent
STANDALONE = os.environ.get("SYCO_SPACE_STANDALONE") == "1"
LOCAL_RAW = PROJECT / "sycophancy-results-local/results/paper-final/raw"
RAW = LOCAL_RAW if LOCAL_RAW.is_dir() and not STANDALONE else SPACE / "data/raw/paper-final"
LOCAL_RUN = PROJECT / "crosslingual-sycophancy-heads/results/paper-final"
ARCHIVE = LOCAL_RUN if LOCAL_RUN.is_dir() and not STANDALONE else SPACE / "data/provenance"
SUMMARY = LOCAL_RUN / "derived/summary.json" if ARCHIVE == LOCAL_RUN else ARCHIVE / "summary.json"
DATA = SPACE / "assets/data"
FIGURES = SPACE / "assets/figures"
MANIFEST = json.loads((ARCHIVE / "manifest.json").read_text())
CONFIG = MANIFEST["config"]
LANGS = CONFIG["languages"]
MODELS = {entry.split(":", 1)[0]: entry.split(":", 1)[1] for entry in CONFIG["models"]}
TAGS = list(MODELS)
LABELS = {
    "q35_08b": "Qwen3.5 0.8B", "q35_2b": "Qwen3.5 2B", "q35_4b": "Qwen3.5 4B",
    "q35_9b": "Qwen3.5 9B", "q35_27b": "Qwen3.5 27B",
    "q25_3b": "Qwen2.5 3B", "q25_7b": "Qwen2.5 7B", "llama8b": "Llama-3.1 8B†",
    "mistral7b": "Mistral 7B", "gemma9b": "Gemma-2 9B", "gemma4_12b": "Gemma-4 12B",
}
BLUE = "#2a78d6"
ORANGE = "#eb6834"
GREEN = "#1baf7a"
INK = "#0b0b0b"
MUTED = "#52514e"
GRID = "#e1e0d9"
SURFACE = "#fcfcfb"
plt.rcParams.update({"font.family": "DejaVu Sans", "font.size": 9, "axes.spines.top": False,
                     "axes.spines.right": False, "axes.edgecolor": GRID, "axes.labelcolor": MUTED,
                     "xtick.color": MUTED, "ytick.color": INK, "figure.facecolor": SURFACE,
                     "axes.facecolor": SURFACE, "svg.fonttype": "none", "svg.hashsalt": "sycophancy-figures-v1"})
SOURCES: set[Path] = set()
ACCEPTED_CODE_HASHES = {MANIFEST["code_hash"]}
for migration in reversed(MANIFEST.get("compatible_code_migrations", [])):
    if migration["to"] in ACCEPTED_CODE_HASHES:
        ACCEPTED_CODE_HASHES.add(migration["from"])


def check_meta(meta: dict, name: str) -> None:
    config = meta.get("config", {})
    expected = hashlib.sha256(json.dumps(config, sort_keys=True, separators=(",", ":"),
                                         ensure_ascii=False).encode()).hexdigest()[:16]
    if (meta.get("schema_version") != 2 or meta.get("run_id") != "paper-final"
            or meta.get("status") != "complete" or meta.get("code_hash") not in ACCEPTED_CODE_HASHES
            or meta.get("config_hash") != expected):
        raise ValueError(f"Invalid final-run provenance: {name}")


def source(path: Path) -> Path:
    if not path.is_file():
        raise FileNotFoundError(f"Missing raw input: {path}")
    SOURCES.add(path)
    return path


def read(name: str) -> dict:
    data = json.loads(source(RAW / name).read_text(encoding="utf-8"))
    check_meta(data.get("_meta", {}), name)
    return data


def read_npz(name: str, fields: tuple[str, ...]) -> dict:
    with np.load(source(RAW / name), allow_pickle=False) as data:
        meta = json.loads(str(data["_meta"].item()))
        check_meta(meta, name)
        return {**{field: data[field].copy() for field in fields}, "_meta": meta}


def fingerprint(path: Path) -> str:
    digest = hashlib.sha256()
    with path.open("rb") as handle:
        for chunk in iter(lambda: handle.read(1024 * 1024), b""):
            digest.update(chunk)
    return digest.hexdigest()


def rho(first, second) -> float:
    a = np.asarray(first, dtype=float).ravel()
    b = np.asarray(second, dtype=float).ravel()
    valid = np.isfinite(a) & np.isfinite(b)
    if valid.sum() < 4 or not np.std(a[valid]) or not np.std(b[valid]):
        return float("nan")
    return float(np.corrcoef(a[valid], b[valid])[0, 1])


def mean(values) -> float:
    a = np.asarray(list(values), dtype=float)
    return float(np.nanmean(a)) if np.isfinite(a).any() else float("nan")


def clean(value):
    """JSON cannot represent NumPy scalars or non-finite floats."""
    if isinstance(value, dict):
        return {str(key): clean(entry) for key, entry in value.items()}
    if isinstance(value, (list, tuple, np.ndarray)):
        return [clean(entry) for entry in value]
    if isinstance(value, (np.integer,)):
        return int(value)
    if isinstance(value, (np.floating, float)):
        return float(value) if np.isfinite(value) else None
    return value


def save(name: str, content: dict, inputs: set[Path]) -> None:
    DATA.mkdir(parents=True, exist_ok=True)
    content = {**content, "provenance": {
        "run_id": "paper-final",
        "inputs": [{"path": f"data/raw/paper-final/{path.name}", "sha256": fingerprint(path)}
                   for path in sorted(inputs)],
    }}
    (DATA / f"{name}.json").write_text(json.dumps(clean(content), ensure_ascii=False,
                                                   sort_keys=True, indent=2, allow_nan=False) + "\n", encoding="utf-8")


def figure(name: str, fig) -> None:
    FIGURES.mkdir(parents=True, exist_ok=True)
    fig.savefig(FIGURES / f"{name}.svg", bbox_inches="tight", pad_inches=0.15,
                facecolor=SURFACE, metadata={"Date": None})
    plt.close(fig)


def tidy(ax, axis="x"):
    ax.grid(axis=axis, color=GRID, lw=.7)
    ax.set_axisbelow(True)


def coverage() -> dict:
    inputs = set()
    rows = []
    for tag in TAGS:
        path = RAW / f"knowledge_{tag}.npz"
        inputs.add(path)
        matrix = read_npz(path.name, ("languages", "correct", "margin"))
        languages = [str(value) for value in matrix["languages"]]
        assert languages == LANGS, f"knowledge language mismatch: {tag}"
        valid = matrix["correct"] & (matrix["margin"] >= CONFIG["knowledge_margin"])
        available = int(valid.all(axis=0).sum())
        rows.append({"tag": tag, "model": LABELS[tag], "eligible": available,
                     "selected": min(available, CONFIG["n"]),
                     "per_language": {lang: int(valid[i].sum()) for i, lang in enumerate(LANGS)}})
    save("coverage", {"rows": rows, "languages": LANGS, "target": CONFIG["n"],
                      "selection": "correct in all 15 neutral conditions, margin >= 0"}, inputs)
    fig, ax = plt.subplots(figsize=(9.6, 5.5))
    fig.subplots_adjust(left=.25, right=.94, top=.94, bottom=.15)
    y = np.arange(len(rows))
    ax.barh(y, [row["selected"] for row in rows], height=.62, color=BLUE)
    ax.set_yticks(y, [row["model"] for row in rows]); ax.invert_yaxis()
    ax.set_xlim(0, 1120); ax.set_xlabel("Eligible aligned questions selected, up to 1,000")
    tidy(ax)
    for i, row in enumerate(rows):
        ax.text(row["selected"] + 15, i, str(row["selected"]), va="center", fontsize=8, color=INK)
    ax.axvline(300, color=MUTED, ls="--", lw=1)
    figure("coverage", fig)
    return {row["tag"]: row["selected"] for row in rows}


def maps() -> dict:
    inputs = set()
    summaries = []
    hero = None
    for tag in TAGS:
        by_lang = {}
        for lang in LANGS:
            path = RAW / f"eapig_{tag}_{lang}.json"
            inputs.add(path)
            d = read(path.name)
            by_lang[lang] = [np.asarray(fold, dtype=float) for fold in d["split_atp_score"]]
            if tag == "q25_7b" and lang == "EN":
                hero = d
        pairs = []
        matrix = np.full((len(LANGS), len(LANGS)), np.nan)
        for i in range(len(LANGS)):
            for j in range(i + 1, len(LANGS)):
                a, b = LANGS[i], LANGS[j]
                value = mean([rho(by_lang[a][0], by_lang[b][1]),
                              rho(by_lang[a][1], by_lang[b][0])])
                matrix[i, j] = matrix[j, i] = value
                pairs.append({"source": a, "target": b, "signed_r": value})
        summaries.append({"tag": tag, "model": LABELS[tag], "mean": mean(x["signed_r"] for x in pairs),
                          "pairs_below_0_6": sum(x["signed_r"] < .6 for x in pairs),
                          "pair_matrix": clean(matrix), "pairs": pairs})
    assert hero is not None
    en = np.asarray(hero["split_atp_score"][0])
    # Select positive training-fold heads once; show those same coordinates in all 15 languages.
    selected = np.argsort(en.ravel())[::-1][:18]
    layers = hero["attn_layers"]; nheads = hero["nH"]
    head_labels = [f"L{layers[int(i // nheads)]}H{int(i % nheads)}" for i in selected]
    values = []
    for lang in LANGS:
        d = read(f"eapig_q25_7b_{lang}.json")
        values.append(np.asarray(d["atp_score"])[np.unravel_index(selected, en.shape)].tolist())
    save("maps", {"models": summaries, "languages": LANGS,
                  "hero": {"model": "Qwen2.5 7B", "heads": head_labels,
                           "values": values, "measure": "first-order effect on wrong-minus-correct logit margin"}}, inputs)
    weak_pairs = [dict(tag=row["tag"], model=row["model"],
                       pairs_below_0_6=row["pairs_below_0_6"],
                       involving_yo=sum(pair["signed_r"] < .6 and "YO" in (pair["source"], pair["target"])
                                        for pair in row["pairs"])) for row in summaries]
    save("map_pairs", {"models": [{"tag": row["tag"], "model": row["model"],
                                   "pair_matrix": row["pair_matrix"]} for row in summaries],
                       "languages": LANGS, "weak_pairs": weak_pairs,
                       "threshold": 0.6, "statistic": "opposite-fold signed correlation; dependent language pairs"}, inputs)
    from matplotlib.colors import LinearSegmentedColormap
    pair_cmap = LinearSegmentedColormap.from_list("pair-strength", ["#cde2fb", BLUE, "#104281"])
    pair_cmap.set_bad("#f0efec")
    selected_pairs = next(row for row in summaries if row["tag"] == "q25_7b")
    fig, ax = plt.subplots(figsize=(8.2, 6.35))
    fig.subplots_adjust(left=.11, right=.86, top=.94, bottom=.12)
    values_pair = np.asarray([[np.nan if entry is None else entry for entry in row]
                              for row in selected_pairs["pair_matrix"]])
    image = ax.imshow(values_pair, vmin=0, vmax=1, cmap=pair_cmap, aspect="equal")
    ax.set_xticks(range(len(LANGS)), LANGS, fontsize=8)
    ax.set_yticks(range(len(LANGS)), LANGS, fontsize=8)
    ax.set_xlabel("Other-language map (opposite item fold)")
    ax.set_ylabel("Source-language map")
    fig.colorbar(image, ax=ax, shrink=.69, pad=.03, label="Signed map correlation")
    figure("map-pairs", fig)
    fig, ax = plt.subplots(figsize=(12.2, 4.25))
    fig.subplots_adjust(left=.055, right=.99, top=.90, bottom=.22)
    v = np.asarray(values)
    from matplotlib.colors import TwoSlopeNorm
    from matplotlib.colors import LinearSegmentedColormap
    cmap = LinearSegmentedColormap.from_list("pressure", [BLUE, "#f0efec", "#e34948"])
    bound = float(np.nanmax(np.abs(v)))
    im = ax.imshow(v, aspect="auto", cmap=cmap, norm=TwoSlopeNorm(vcenter=0, vmin=-bound, vmax=bound))
    ax.set_xticks(range(len(head_labels)), head_labels, rotation=58, ha="right", fontsize=7)
    ax.set_yticks(range(len(LANGS)), LANGS, fontsize=8)
    ax.set_ylabel("Language")
    fig.colorbar(im, ax=ax, shrink=.65, pad=.015, label="Estimated signed margin contribution")
    figure("head-atlas", fig)
    fig, ax = plt.subplots(figsize=(9.6, 5.2))
    fig.subplots_adjust(left=.26, right=.95, top=.93, bottom=.15)
    y = np.arange(len(summaries))
    ax.barh(y, [row["mean"] for row in summaries], height=.60, color=BLUE)
    ax.set_yticks(y, [row["model"] for row in summaries]); ax.invert_yaxis()
    ax.set_xlim(0, 1.08); ax.set_xlabel("Mean opposite-fold signed correlation (different questions)")
    tidy(ax)
    for i, row in enumerate(summaries):
        ax.text(row["mean"] + .02, i, f"{row['mean']:.2f}", va="center", color=INK)
    figure("map-agreement", fig)
    return {row["tag"]: row["mean"] for row in summaries}


def patching() -> None:
    inputs = set(); rows = []
    for spec in CONFIG["patch_models"]:
        tag = spec.split(":", 1)[0]
        for lang in CONFIG["patch_languages"]:
            path = RAW / f"headpatch_{tag}_{lang}.json"
            inputs.add(path); d = read(path.name)
            denom = float(d["denominator"])
            if not d["denom_stable"] or denom <= .5:
                raise ValueError(f"Unstable exact-patch denominator: {tag}/{lang}")
            estimates = []
            for key in ("top_set_denoising_items", "top_set_sufficiency_items"):
                samples = np.asarray(d[key], dtype=float) / denom
                assert len(samples) == d["n"]
                rng = np.random.default_rng(int(hashlib.sha256(f"{tag}:{lang}:{key}".encode()).hexdigest()[:8], 16))
                boot = [rng.choice(samples, len(samples), replace=True).mean() for _ in range(1000)]
                estimates.append({"value": float(samples.mean()), "lo": float(np.percentile(boot, 2.5)),
                                  "hi": float(np.percentile(boot, 97.5))})
            rows.append({"tag": tag, "model": LABELS[tag], "language": lang, "n": d["n"],
                         "denominator": denom, "denoise": estimates[0], "reverse": estimates[1]})
    save("patch", {"rows": rows, "languages": CONFIG["patch_languages"],
                   "ci": "1,000 item bootstraps; observed cell denominator held fixed"}, inputs)
    fig, ax = plt.subplots(figsize=(9.6, 6.2))
    fig.subplots_adjust(left=.27, right=.96, top=.93, bottom=.13)
    for i, row in enumerate(rows):
        for key, color, offset in (("denoise", BLUE, -.15), ("reverse", ORANGE, .15)):
            item = row[key]
            ax.plot([item["lo"], item["hi"]], [i + offset] * 2, color=color, lw=1.8)
            ax.plot(item["value"], i + offset, marker="o" if key == "denoise" else "s",
                    color=color, markersize=5)
    ax.set_yticks(range(len(rows)), [f"{r['model']} · {r['language']}" for r in rows], fontsize=7)
    ax.invert_yaxis(); ax.axvline(0, color=MUTED, lw=.9)
    ax.set_xlabel("Normalized restored / induced wrong-answer margin")
    ax.plot([], [], "o", color=BLUE, label="Denoising")
    ax.plot([], [], "s", color=ORANGE, label="Reverse patch")
    ax.legend(loc="lower right", frameon=False, ncol=2)
    tidy(ax)
    figure("patch", fig)


def faithfulness() -> dict:
    item_inputs = set(); item_rows = []; by_contrast = {}
    contrasts = (("factual", "eapig_items"), ("speaker_change", "eapig_ctrl_items"),
                 ("answer_selection", "eapig_ctrl2_items"), ("social_pressure", "eapig_social_items"))
    for contrast, prefix in contrasts:
        model_rows = []
        for tag in TAGS:
            current = []
            for lang in LANGS:
                name = f"{prefix}_{tag}_{lang}.npz"
                path = RAW / name; item_inputs.add(path)
                data = read_npz(name, ("item_ids", "effects", "baseline_margin", "treatment_margin"))
                item_ids = np.asarray(data["item_ids"])
                observed = np.asarray(data["treatment_margin"] - data["baseline_margin"], dtype=float)
                attributed = np.asarray(data["effects"], dtype=float).sum(axis=(1, 2))
                if (len(item_ids) != len(observed) or len(item_ids) != len(attributed)
                        or len(set(item_ids.tolist())) != len(item_ids)
                        or len(item_ids) != data["_meta"]["actual_n"]):
                    raise AssertionError(f"Invalid item-level attribution sample: {name}")
                current.append({"contrast": contrast, "tag": tag, "model": LABELS[tag], "language": lang,
                                "n": len(item_ids), "item_r": rho(attributed, observed),
                                "observed_shift": float(observed.mean()),
                                "attributed_shift": float(attributed.mean())})
            item_rows.extend(current)
            model_rows.append({"tag": tag, "model": LABELS[tag], "contrast": contrast,
                               "mean_item_r": mean(row["item_r"] for row in current),
                               "negative_languages": sum(row["item_r"] < 0 for row in current)})
        by_contrast[contrast] = model_rows
    by_model = by_contrast["factual"]
    save("item_fidelity", {"rows": item_rows, "models": by_model, "models_by_contrast": by_contrast,
                           "languages": LANGS,
                           "contrasts": {"factual": "false user answer versus neutral question",
                                         "speaker_change": "user versus third-person false assertion",
                                         "answer_selection": "neutral question versus options only; correct-minus-wrong target",
                                         "social_pressure": "emotional pressure versus matched filler on false assertion"},
                           "measure": "contrast-specific per-item correlation of summed first-order scores with observed margin shift; do not pool contrasts"}, item_inputs)
    fig, ax = plt.subplots(figsize=(9.6, 5.3))
    fig.subplots_adjust(left=.26, right=.94, top=.94, bottom=.15)
    y = np.arange(len(by_model))
    color = [BLUE if row["mean_item_r"] >= 0 else "#e34948" for row in by_model]
    ax.barh(y, [row["mean_item_r"] for row in by_model], color=color, height=.60)
    ax.set_yticks(y, [row["model"] for row in by_model]); ax.invert_yaxis()
    ax.axvline(0, color=MUTED, linewidth=.8)
    ax.set_xlim(-.2, 1.0); ax.set_xlabel("Mean per-language item-level attribution correlation")
    tidy(ax)
    for i, row in enumerate(by_model):
        value = row["mean_item_r"]
        ax.text(value + (.02 if value >= 0 else -.02), i, f"{value:.2f}", va="center",
                ha="left" if value >= 0 else "right", color=INK, fontsize=8)
    figure("item-fidelity", fig)

    patch_inputs = set(); patch_rows = []
    for spec in CONFIG["patch_models"]:
        tag = spec.split(":", 1)[0]
        for lang in CONFIG["patch_languages"]:
            apath = RAW / f"eapig_{tag}_{lang}.json"; ppath = RAW / f"headpatch_{tag}_{lang}.json"
            patch_inputs.update((apath, ppath))
            attribution = read(apath.name); exact = read(ppath.name)
            first = np.abs(np.asarray(attribution["split_atp_score"][0], dtype=float))
            second = np.abs(np.asarray(exact["denoising_effect"], dtype=float))
            if first.shape != second.shape or not exact["denom_stable"]:
                raise AssertionError(f"Unmatched exact patch and attribution head maps: {tag}/{lang}")
            patch_rows.append({"tag": tag, "model": LABELS[tag], "language": lang,
                               "n": exact["n"], "head_r": rho(first, second)})
    save("patch_alignment", {"rows": patch_rows,
                             "measure": "correlation of absolute fold-0 attribution and absolute exact-denoising head maps",
                             "scope": "seven checkpoints, EN/ZH/AR; not per-item prediction"}, patch_inputs)
    fig, ax = plt.subplots(figsize=(9.6, 5.6))
    fig.subplots_adjust(left=.28, right=.95, top=.94, bottom=.14)
    y = np.arange(len(patch_rows))
    ax.scatter([row["head_r"] for row in patch_rows], y,
               c=[ORANGE if row["head_r"] < .9 else BLUE for row in patch_rows], s=33)
    ax.axvline(.9, color=MUTED, linewidth=.9, linestyle="--")
    ax.set_yticks(y, [f"{row['model']} · {row['language']}" for row in patch_rows], fontsize=7)
    ax.invert_yaxis(); ax.set_xlim(0, 1.07)
    ax.set_xlabel("Absolute attribution map versus exact denoising map (r)")
    tidy(ax)
    figure("patch-alignment", fig)
    return {"item_rows": [row for row in item_rows if row["contrast"] == "factual"],
            "all_item_rows": item_rows, "models": by_model, "patch_rows": patch_rows}


def interventions() -> dict:
    inputs = set(); rows = []; tradeoffs = []; source_rows = []
    for tag in TAGS:
        path = RAW / f"headsteer_{tag}.json"
        inputs.add(path); d = read(path.name)
        assert d["src"] == "EN" and set(d["lang"]) == set(LANGS)
        for lang in LANGS:
            c = d["lang"][lang]
            rows.append({"tag": tag, "model": LABELS[tag], "language": lang,
                         "baseline_caving": c["base"][0], "replacement_caving": c["ablate_causal"][0],
                         "random_caving": c["ablate_random"][0],
                         "baseline_accuracy": c["base"][1], "replacement_accuracy": c["ablate_causal"][1],
                         "null_p": c["random_null"]["p_value"], "n": d["_meta"]["actual_n"]})
        base_flip = mean(d["lang"][lang]["base"][0] for lang in LANGS)
        base_acc = mean(d["lang"][lang]["base"][1] for lang in LANGS)
        settings = [("Head count", "k", value) for value in d["ksweep"][LANGS[0]]]
        settings += [("Replacement strength", "alpha", value)
                     for value in d["alphasweep"][LANGS[0]][1:-1]]
        for kind, key, choice in settings:
            level = choice[key]
            measures = [next(entry for entry in d["ksweep"][lang] if entry[key] == level)
                        if key == "k" else next(entry for entry in d["alphasweep"][lang] if entry[key] == level)
                        for lang in LANGS]
            tradeoffs.append({"tag": tag, "model": LABELS[tag], "kind": kind, "setting": level,
                              "caving_reduction_pp": 100 * (base_flip - mean(m["flip"] for m in measures)),
                              "accuracy_loss_pp": 100 * (base_acc - mean(m["acc"] for m in measures)),
                              "worst_language_accuracy_loss_pp": max(100 * (d["lang"][lang]["base"][1] - m["acc"])
                                                               for lang, m in zip(LANGS, measures))})
    for tag in ("q25_7b", "llama8b", "gemma9b"):
        for src in ("ZH", "AR"):
            path = RAW / f"headsteer_{tag}_src{src}.json"
            inputs.add(path); d = read(path.name)
            for lang in LANGS:
                c = d["lang"][lang]
                source_rows.append({"tag": tag, "model": LABELS[tag], "source": src, "target": lang,
                                    "reduced": c["ablate_causal"][0] < c["base"][0],
                                    "wrong_reduction_pp": 100 * (c["base"][0] - c["ablate_causal"][0]),
                                    "accuracy_loss_pp": 100 * (c["base"][1] - c["ablate_causal"][1])})
    source_aggregate = []
    for tag in ("q25_7b", "llama8b", "gemma9b"):
        for lang in ("ZH", "AR"):
            cells = [row for row in source_rows if row["tag"] == tag and row["source"] == lang]
            source_aggregate.append({"tag": tag, "model": LABELS[tag], "source": lang,
                                     "wrong_reduction_pp": mean(row["wrong_reduction_pp"] for row in cells),
                                     "accuracy_loss_pp": mean(row["accuracy_loss_pp"] for row in cells),
                                     "reduced_targets": sum(row["reduced"] for row in cells)})
    save("intervention", {"rows": rows, "tradeoff": tradeoffs, "source_symmetry": source_rows,
                          "source_aggregate": source_aggregate,
                          "source_reduced": sum(row["reduced"] for row in source_rows)}, inputs)
    fig, ax = plt.subplots(figsize=(9.3, 4.6))
    fig.subplots_adjust(left=.13, right=.95, top=.92, bottom=.16)
    for lang, color, marker in (("ZH", BLUE, "o"), ("AR", ORANGE, "s")):
        subset = [row for row in source_aggregate if row["source"] == lang]
        ax.scatter([row["accuracy_loss_pp"] for row in subset],
                   [row["wrong_reduction_pp"] for row in subset], color=color, marker=marker,
                   s=53, label=f"{lang}-selected heads")
        for row in subset:
            ax.annotate(row["model"], (row["accuracy_loss_pp"], row["wrong_reduction_pp"]),
                        xytext=(5, 4), textcoords="offset points", fontsize=7)
    ax.axhline(0, color=MUTED, lw=.8); ax.axvline(0, color=MUTED, lw=.8)
    ax.set_xlabel("Neutral accuracy lost across target languages (pp)")
    ax.set_ylabel("Less wrong-option selection across targets (pp)")
    ax.legend(frameon=False, fontsize=8)
    tidy(ax, "both")
    figure("source-symmetry", fig)
    fig, ax = plt.subplots(figsize=(9.6, 5.5))
    fig.subplots_adjust(left=.25, right=.95, top=.93, bottom=.15)
    for i, tag in enumerate(TAGS):
        group = [row for row in rows if row["tag"] == tag]
        x = mean((r["baseline_caving"] - r["replacement_caving"]) * 100 for r in group)
        y = mean((r["baseline_accuracy"] - r["replacement_accuracy"]) * 100 for r in group)
        ax.scatter(x, y, color=BLUE if tag != "llama8b" else ORANGE, s=42, zorder=3)
        ax.annotate(LABELS[tag], (x, y), xytext=(5, 3), textcoords="offset points", fontsize=7)
    ax.axhline(0, color=MUTED, lw=.8); ax.axvline(0, color=MUTED, lw=.8)
    ax.set_xlabel("Less wrong-option selection (percentage points)")
    ax.set_ylabel("Neutral accuracy lost (percentage points)")
    tidy(ax, "both")
    figure("intervention", fig)
    fig, ax = plt.subplots(figsize=(9.6, 5.3))
    fig.subplots_adjust(left=.13, right=.96, top=.93, bottom=.16)
    for tag in TAGS:
        group = [row for row in tradeoffs if row["tag"] == tag]
        ax.scatter([row["accuracy_loss_pp"] for row in group],
                   [row["caving_reduction_pp"] for row in group], s=21,
                   color=ORANGE if tag == "llama8b" else BLUE, alpha=.5)
    ax.axvline(5, color=MUTED, lw=1, ls="--")
    ax.axhline(0, color=GRID, lw=1)
    ax.set_xlabel("Neutral accuracy lost (percentage points)")
    ax.set_ylabel("Less wrong-option selection (percentage points)")
    tidy(ax, "both")
    figure("tradeoff", fig)
    return {"reduced": sum(row["replacement_caving"] < row["baseline_caving"] for row in rows),
            "cells": len(rows), "source_reduced": sum(row["reduced"] for row in source_rows)}


def specificity() -> dict:
    inputs = set(); rows = []
    for tag in TAGS:
        mapping = {mode: {} for mode in ("factual", "answer", "third", "social")}
        splits = {mode: {} for mode in ("factual", "answer")}
        decode = {}
        for lang in LANGS:
            for mode, prefix in (("factual", "eapig"), ("answer", "eapig_ctrl2"),
                                 ("third", "eapig_ctrl"), ("social", "eapig_social")):
                path = RAW / f"{prefix}_{tag}_{lang}.json"
                inputs.add(path); d = read(path.name)
                mapping[mode][lang] = np.asarray(d["atp_score"])
                if mode in splits:
                    splits[mode][lang] = [np.asarray(fold) for fold in d["split_atp_score"]]
                if mode == "factual":
                    decode[lang] = np.asarray(d["decode_auroc"])
        residual = {}
        for lang in LANGS:
            f = splits["factual"][lang]; a = splits["answer"][lang]
            residual[lang] = []
            for fit, test in ((0, 1), (1, 0)):
                beta = float(np.dot(f[fit].ravel(), a[fit].ravel()) /
                             (np.dot(a[fit].ravel(), a[fit].ravel()) + 1e-12))
                residual[lang].append(f[test] - beta * a[test])
        disjoint = []
        for i, source_lang in enumerate(LANGS):
            for target_lang in LANGS[i + 1:]:
                disjoint.extend((rho(residual[source_lang][0], residual[target_lang][1]),
                                 rho(residual[source_lang][1], residual[target_lang][0])))
        rows.append({"tag": tag, "model": LABELS[tag],
                     "answer_map": mean(rho(abs(mapping["factual"][lang]), abs(mapping["answer"][lang]))
                                        for lang in LANGS),
                     "speaker_change": mean(rho(abs(mapping["factual"][lang]), abs(mapping["third"][lang]))
                                            for lang in LANGS),
                     "social_pressure": mean(rho(abs(mapping["factual"][lang]), abs(mapping["social"][lang]))
                                             for lang in LANGS),
                     "residual_cross_language": mean(disjoint),
                     "decode_attribution": mean(rho(abs(decode[lang] - .5), abs(mapping["factual"][lang]))
                                                for lang in LANGS)})
    save("specificity", {"rows": rows, "overlap_measure": "within-language absolute-map correlation",
                         "residual_measure": "opposite-fold cross-language signed correlation"}, inputs)
    fig, ax = plt.subplots(figsize=(9.6, 5.4))
    fig.subplots_adjust(left=.26, right=.95, top=.93, bottom=.16)
    y = np.arange(len(rows)); h = .21
    for key, label, color, offset in (("speaker_change", "Speaker change", BLUE, -h),
                                      ("answer_map", "Answer selection", ORANGE, 0),
                                      ("social_pressure", "Social pressure", GREEN, h)):
        ax.barh(y + offset, [row[key] for row in rows], height=h * .91, color=color, label=label)
    ax.invert_yaxis(); ax.set_yticks(y, [row["model"] for row in rows]); ax.set_xlim(0, 1.05)
    ax.set_xlabel("Within-language absolute head-map correlation")
    ax.legend(loc="lower right", ncol=3, frameon=False)
    tidy(ax)
    figure("specificity", fig)
    return {key: mean(row[key] for row in rows) for key in rows[0] if key not in ("tag", "model")}


def representations() -> dict:
    inputs = set(); rows = []; matrices = {}; concepts = []; methods = []
    for tag in TAGS:
        xpath = RAW / f"xstats_{tag}.json"; cpath = RAW / f"conc_{tag}.json"
        mpath = RAW / f"methods_{tag}.json"; ppath = RAW / f"promptvar_{tag}.json"
        inputs.update((xpath, cpath, mpath, ppath))
        x = read(xpath.name); c = read(cpath.name); m = read(mpath.name); p = read(ppath.name)
        assert set(x["langs"]) == set(LANGS) and c["langs"] == LANGS
        order = [x["langs"].index(lang) for lang in LANGS]
        rows.append({"tag": tag, "model": LABELS[tag], "factual_auroc": x["fact"]["mean"],
                     "social_auroc": x["soc"]["mean"], "language_control": x["lang_id_ctrl"],
                     "factual_shared_energy": x["xli_fact"], "social_shared_energy": x["xli_soc"],
                     "layer": x["layer"]})
        matrices[tag] = {"factual": clean(np.asarray(x["matrix_fact"])[np.ix_(order, order)]),
                         "social": clean(np.asarray(x["matrix_soc"])[np.ix_(order, order)])}
        concepts.append({"tag": tag, "model": LABELS[tag],
                         "measures": {kind: mean(c["matrices"][kind][i][j]
                                                 for i in range(len(LANGS)) for j in range(len(LANGS)) if i != j)
                                      for kind in c["beh"]}})
        upper = np.asarray(p["EN_crosstpl_transfer"], dtype=float)
        methods.append({"tag": tag, "model": LABELS[tag],
                        "extractors": {key: m[key] for key in ("diffmean", "logistic", "lda", "pca")},
                        "english_cross_phrasing": mean(upper[i, j]
                                                        for i in range(len(upper)) for j in range(len(upper)) if i != j)})
    save("representation", {"rows": rows, "languages": LANGS, "matrices": matrices,
                             "concept_controls": concepts, "robustness": methods,
                             "contrast": "false versus correct assertion; AUROC is predictive, not intervention"}, inputs)
    fig, ax = plt.subplots(figsize=(9.6, 5.3))
    fig.subplots_adjust(left=.25, right=.96, top=.94, bottom=.14)
    y = np.arange(len(rows))
    ax.scatter([r["factual_auroc"] for r in rows], y - .13, color=BLUE, s=40, marker="o", label="False versus correct")
    ax.scatter([r["social_auroc"] for r in rows], y + .13, color=ORANGE, s=40, marker="s", label="Pressure versus filler")
    ax.set_yticks(y, [r["model"] for r in rows]); ax.invert_yaxis()
    ax.set_xlim(.4, 1.02); ax.set_xlabel("Held-out off-diagonal AUROC (source language → target language)")
    ax.legend(loc="lower right", frameon=False, ncol=2)
    tidy(ax)
    figure("representation", fig)
    return {"min": min(row["factual_auroc"] for row in rows), "max": max(row["factual_auroc"] for row in rows)}


def robustness() -> dict:
    inputs = set(); extractors = []; phrasing = []
    for tag in TAGS:
        mpath = RAW / f"methods_{tag}.json"; ppath = RAW / f"promptvar_{tag}.json"
        inputs.update((mpath, ppath))
        m = read(mpath.name); p = read(ppath.name)
        matrix = np.asarray(p["EN_crosstpl_transfer"], dtype=float)
        if matrix.shape != (5, 5) or not np.isfinite(matrix).all():
            raise AssertionError(f"Invalid five-phrasing matrix: {tag}")
        off = matrix[~np.eye(5, dtype=bool)]
        extractors.append({"tag": tag, "model": LABELS[tag],
                           "methods": {key: float(m[key]) for key in ("diffmean", "logistic", "lda", "pca")},
                           "n": m["_meta"]["actual_n"]})
        phrasing.append({"tag": tag, "model": LABELS[tag], "mean_off_diagonal": float(off.mean()),
                         "min_off_diagonal": float(off.min()), "matrix": matrix.tolist(),
                         "templates": p["templates"], "n": p["_meta"]["actual_n"]})
    save("robustness", {"extractors": extractors, "phrasing": phrasing,
                        "method_scope": "held-out factual false-vs-correct direction transfer",
                        "phrasing_scope": "five English phrasings only, disjoint questions"}, inputs)
    fig, axes = plt.subplots(2, 2, figsize=(10, 8.2), sharex=True, sharey=True)
    fig.subplots_adjust(left=.19, right=.97, top=.94, bottom=.10, wspace=.11, hspace=.22)
    methods = (("diffmean", "DiffMean", BLUE), ("logistic", "Logistic probe", ORANGE),
               ("lda", "LDA", GREEN), ("pca", "PCA", "#4a3aa7"))
    y = np.arange(len(extractors))
    for ax, (key, label, color) in zip(axes.flat, methods):
        ax.scatter([row["methods"][key] for row in extractors], y, s=22, color=color)
        ax.set_title(label, loc="left", fontsize=10, weight="bold")
        ax.set_xlim(.45, 1.02); ax.set_yticks(y, [row["model"] for row in extractors], fontsize=7)
        ax.invert_yaxis(); tidy(ax)
    for ax in axes[-1]: ax.set_xlabel("Held-out cross-language AUROC")
    figure("extractors", fig)
    fig, ax = plt.subplots(figsize=(9.6, 5.2))
    fig.subplots_adjust(left=.26, right=.95, top=.92, bottom=.16)
    ax.scatter([row["mean_off_diagonal"] for row in phrasing], y - .12,
               color=BLUE, s=32, marker="o", label="Mean of off-diagonal pairs")
    ax.scatter([row["min_off_diagonal"] for row in phrasing], y + .12,
               color=ORANGE, s=32, marker="s", label="Weakest off-diagonal pair")
    ax.set_yticks(y, [row["model"] for row in phrasing]); ax.invert_yaxis()
    ax.set_xlim(.45, 1.01); ax.set_xlabel("Held-out AUROC across English prompt phrasings")
    ax.legend(loc="lower right", frameon=False, fontsize=8)
    tidy(ax)
    figure("prompt-phrasing", fig)
    return {"extractors": extractors, "phrasing": phrasing}


def behavior() -> None:
    inputs = set(); rows = []; curves = []
    for tag in TAGS:
        fpath = RAW / f"freeform_judge_{tag}.json"; ppath = RAW / f"persuade_{tag}.json"
        inputs.update((fpath, ppath))
        f = read(fpath.name); p = read(ppath.name)
        assert f["langs"] == LANGS and set(p["lang"]) == set(LANGS)
        for lang in LANGS:
            r = f["lang"][lang]
            rows.append({"tag": tag, "model": LABELS[tag], "language": lang, "n": r["n"],
                         "judged_caving": r["judge_flip"], "cross_family_caving": r["judge_flip_cross_family"],
                         "letter_caving": r["letter_flip"], "judge_disagreement": r["judge_disagreement"],
                         "judge_letter_agreement": r["judge_vs_letter_agree"],
                         "axis_cosine": r["cos_freeform_MC_heldout"]})
            curves.append({"tag": tag, "model": LABELS[tag], "language": lang,
                           "turns": p["turns"], "p_correct": p["lang"][lang]["P_correct"],
                           "wrong_option": p["lang"][lang]["flip"]})
    save("behavior", {"rows": rows, "curves": curves,
                      "judges": ["Gemma-4-31B", "Qwen3.5-27B"],
                      "limits": "no human-labeled judge calibration; forced-choice head interventions not tested here"}, inputs)
    selected = [row for row in rows if row["tag"] in ("q25_7b", "llama8b", "gemma4_12b") and row["language"] == "EN"]
    fig, ax = plt.subplots(figsize=(9.2, 4.0))
    fig.subplots_adjust(left=.23, right=.96, top=.94, bottom=.17)
    y = np.arange(len(selected)); h = .24
    ax.barh(y - h/2, [row["judged_caving"] for row in selected], height=h, color=BLUE, label="Two-judge ensemble")
    ax.barh(y + h/2, [row["letter_caving"] for row in selected], height=h, color=ORANGE, label="Letter readout")
    ax.set_yticks(y, [row["model"] for row in selected]); ax.invert_yaxis()
    ax.set_xlim(0, 1); ax.set_xlabel("Share classified as endorsing the wrong answer")
    ax.legend(frameon=False, loc="lower right", ncol=2)
    tidy(ax)
    figure("freeform", fig)
    fig, ax = plt.subplots(figsize=(9.6, 4.2))
    fig.subplots_adjust(left=.12, right=.95, top=.94, bottom=.20)
    chosen = next(row for row in curves if row["tag"] == "q25_7b" and row["language"] == "EN")
    ax.plot(range(7), chosen["p_correct"], color=BLUE, marker="o", lw=2, label="Correct-option probability")
    ax.set_ylim(0, 1.05); ax.set_xticks(range(7), chosen["turns"])
    ax.set_ylabel("Correct-option probability")
    ax.set_xlabel("Conversation step (last step requests a reset)")
    tidy(ax, "y")
    figure("persuasion", fig)


def causal_controls() -> dict:
    inputs = set(); induction = []; steering = []
    for tag in TAGS:
        ipath = RAW / f"induce_{tag}.json"; spath = RAW / f"ccausal_{tag}.json"
        inputs.update((ipath, spath))
        head = read(ipath.name); direction = read(spath.name)
        if (head["src"] != "EN" or set(head["lang"]) != set(LANGS)
                or set(direction["flip"]) != set(LANGS)
                or direction["src"] != ["EN", "DE"]):
            raise AssertionError(f"Incomplete induction/steering source or targets: {tag}")
        for lang in LANGS:
            group = head["lang"][lang]; base = group["base"]
            random = {item["alpha"]: item for item in group["rand"]}
            for value in group["dose"]:
                alpha = value["alpha"]; null = random.get(alpha)
                induction.append({"tag": tag, "model": LABELS[tag], "language": lang,
                                  "alpha": alpha, "wrong_increase_pp": 100 * (value["flip"] - base["flip"]),
                                  "accuracy_loss_pp": 100 * (base["acc"] - value["acc"]),
                                  "random_wrong_increase_pp": 100 * (null["flip"] - base["flip"]) if null else None,
                                  "random_accuracy_loss_pp": 100 * (base["acc"] - null["acc"]) if null else None})
            baseline_wrong = direction["flip"][lang]["0.0"]
            baseline_acc = direction["acc"][lang]["0.0"]
            for alpha in direction["alphas"]:
                key = str(alpha)
                steering.append({"tag": tag, "model": LABELS[tag], "language": lang,
                                 "alpha": alpha, "wrong_increase_pp": 100 * (direction["flip"][lang][key] - baseline_wrong),
                                 "accuracy_loss_pp": 100 * (baseline_acc - direction["acc"][lang][key])})
    def model_average(rows, alpha, fields):
        result = []
        for tag in TAGS:
            subset = [row for row in rows if row["tag"] == tag and row["alpha"] == alpha]
            if len(subset) != len(LANGS):
                raise AssertionError(f"Incomplete dose panel: {tag}/{alpha}")
            result.append({"tag": tag, "model": LABELS[tag], "alpha": alpha,
                           **{field: mean(row[field] for row in subset) for field in fields}})
        return result
    high = model_average(induction, 4.0,
                         ("wrong_increase_pp", "accuracy_loss_pp", "random_wrong_increase_pp", "random_accuracy_loss_pp"))
    gentle = model_average(steering, 2.0, ("wrong_increase_pp", "accuracy_loss_pp"))
    save("causal_controls", {"induction": induction, "induction_alpha4": high,
                             "steering": steering, "steering_alpha2": gentle,
                             "induction_site": "selected attention-head outputs on neutral questions; EN fit",
                             "steering_site": "one residual-stream layer on false-assertion questions; EN+DE fit",
                             "interpretation": "descriptive dose/cost; not clean induction or free-form mitigation"}, inputs)
    fig, ax = plt.subplots(figsize=(9.6, 5.3))
    fig.subplots_adjust(left=.12, right=.95, top=.94, bottom=.16)
    ax.scatter([row["random_accuracy_loss_pp"] for row in high],
               [row["random_wrong_increase_pp"] for row in high], s=43,
               marker="s", facecolors="none", edgecolors=ORANGE, linewidths=1.5, label="Matched random heads")
    ax.scatter([row["accuracy_loss_pp"] for row in high], [row["wrong_increase_pp"] for row in high],
               s=43, color=BLUE, label="Selected heads")
    for row in high:
        if row["tag"] in ("mistral7b", "llama8b", "q25_3b"):
            ax.annotate(row["model"], (row["accuracy_loss_pp"], row["wrong_increase_pp"]),
                        xytext=(5, 4), textcoords="offset points", fontsize=8)
    ax.axhline(0, color=MUTED, lw=.8); ax.axvline(0, color=MUTED, lw=.8)
    ax.set_xlabel("Neutral accuracy lost at α=4 (percentage points)")
    ax.set_ylabel("More wrong-option choices on neutral prompts (pp)")
    ax.legend(frameon=False, loc="upper left", fontsize=8)
    tidy(ax, "both")
    figure("induction-cost", fig)
    fig, ax = plt.subplots(figsize=(9.6, 5.2))
    fig.subplots_adjust(left=.13, right=.96, top=.94, bottom=.16)
    ax.scatter([row["accuracy_loss_pp"] for row in gentle],
               [row["wrong_increase_pp"] for row in gentle], color=BLUE, s=46)
    for row in gentle:
        ax.annotate(row["model"], (row["accuracy_loss_pp"], row["wrong_increase_pp"]),
                    xytext=(4, 5), textcoords="offset points", fontsize=7)
    ax.axhline(0, color=MUTED, lw=.8); ax.axvline(0, color=MUTED, lw=.8)
    ax.set_xlabel("Neutral accuracy lost at α=+2 (percentage points)")
    ax.set_ylabel("Change in false-assertion wrong-option choice (pp)")
    tidy(ax, "both")
    figure("steering-cost", fig)
    return {"induction": induction, "induction_alpha4": high, "steering": steering, "steering_alpha2": gentle}


def judge_diagnostics() -> dict:
    inputs = set(); rows = []; by_model = []
    for tag in TAGS:
        path = RAW / f"freeform_judge_{tag}.json"
        inputs.add(path); data = read(path.name)
        if data["langs"] != LANGS:
            raise AssertionError(f"Incomplete free-form language panel: {tag}")
        group = []
        for lang in LANGS:
            item = data["lang"][lang]
            group.append({"tag": tag, "model": LABELS[tag], "language": lang, "n": item["n"],
                          "judge_disagreement": item["judge_disagreement"],
                          "judge_letter_agreement": item["judge_vs_letter_agree"],
                          "all_cross_family_agreement": item["all_vs_cross_family_agree"],
                          "all_judge_caving": item["judge_flip"],
                          "cross_family_caving": item["judge_flip_cross_family"],
                          "letter_caving": item["letter_flip"],
                          "per_judge_caving": item["per_judge_flip"],
                          "axis_cosine": item["cos_freeform_MC_heldout"]})
        rows.extend(group)
        by_model.append({"tag": tag, "model": LABELS[tag],
                         "judge_disagreement": mean(row["judge_disagreement"] for row in group),
                         "judge_letter_agreement": mean(row["judge_letter_agreement"] for row in group),
                         "all_cross_family_agreement": mean(row["all_cross_family_agreement"] for row in group)})
    save("judge_diagnostics", {"rows": rows, "models": by_model, "languages": LANGS,
                               "measure": "equal-language descriptive averages; no human-validated labels"}, inputs)
    fig, axes = plt.subplots(2, 1, figsize=(9.6, 7.3))
    fig.subplots_adjust(left=.26, right=.94, top=.94, bottom=.12, hspace=.44)
    y = np.arange(len(by_model))
    for ax, field, title, limits, color in (
            (axes[0], "judge_disagreement", "The two model judges disagree", (0, .25), ORANGE),
            (axes[1], "judge_letter_agreement", "Judge and final-letter labels agree", (.55, 1), BLUE)):
        ax.barh(y, [row[field] for row in by_model], color=color, height=.57)
        ax.set_xlim(*limits); ax.set_yticks(y, [row["model"] for row in by_model], fontsize=7)
        ax.invert_yaxis(); ax.set_title(title, loc="left", fontsize=10, weight="bold")
        ax.set_xlabel("Equal-language mean fraction"); tidy(ax)
    figure("judge-diagnostics", fig)
    return {"rows": rows, "models": by_model}


def check(saved: dict) -> None:
    summary = json.loads(SUMMARY.read_text())
    for tag, value in saved["maps"].items():
        expected = summary["models"][tag]["map_agreement"]["disjoint_content_signed_mean"]
        if not np.isclose(value, expected, atol=2e-5):
            raise AssertionError(f"Map mismatch {tag}: {value} versus {expected}")
    if saved["interventions"]["reduced"] != 151 or saved["interventions"]["cells"] != 165:
        raise AssertionError("Intervention coverage does not match archived run")
    if len([tag for tag, n in saved["coverage"].items() if n >= 300]) != 10:
        raise AssertionError("Unexpected high-power cohort")
    for key, number in (("answer_map", .70), ("speaker_change", .88), ("social_pressure", .68),
                        ("residual_cross_language", .80)):
        if abs(saved["specificity"][key] - number) > .015:
            raise AssertionError(f"Specificity mismatch {key}: {saved['specificity'][key]}")
    by_developer = {}
    for tag, score in saved["maps"].items():
        if saved["coverage"][tag] >= 300:
            by_developer.setdefault(MODELS[tag].split("/", 1)[0], []).append(score)
    weighted = mean(mean(scores) for scores in by_developer.values())
    expected = summary["aggregate"]["high_power_developer_weighted_disjoint_signed_r"]
    if not np.isclose(weighted, expected, atol=2e-5):
        raise AssertionError(f"Developer-weighted map mismatch: {weighted} versus {expected}")
    pair_data = json.loads((DATA / "map_pairs.json").read_text())
    if (sum(row["pairs_below_0_6"] for row in pair_data["weak_pairs"]) != 40
            or sum(row["involving_yo"] for row in pair_data["weak_pairs"]) != 33):
        raise AssertionError("Weak language-pair counts differ from raw maps")
    for row in saved["faithfulness"]["item_rows"]:
        expected = summary["models"][row["tag"]]["causal_effect"][row["language"]]["atp_item_faithfulness_r"]
        if not np.isclose(row["item_r"], expected, atol=2e-5):
            raise AssertionError(f"Per-item attribution mismatch: {row['tag']}/{row['language']}")
    for row in saved["faithfulness"]["patch_rows"]:
        exact = next(cell for cell in summary["models"][row["tag"]]["exact_patch"]
                     if cell["lang"] == row["language"])
        if not np.isclose(row["head_r"], exact["atp_exact_r"], atol=2e-5):
            raise AssertionError(f"Exact-map agreement mismatch: {row['tag']}/{row['language']}")
    if (len(saved["faithfulness"]["item_rows"]) != 165
            or len(saved["faithfulness"]["all_item_rows"]) != 660
            or len(saved["faithfulness"]["patch_rows"]) != 21
            or len(saved["judges"]["rows"]) != 165
            or len(saved["robustness"]["phrasing"]) != 11
            or len(saved["causal_controls"]["induction"]) != 660
            or len(saved["causal_controls"]["steering"]) != 825):
        raise AssertionError("A new diagnostic panel is incomplete")
    if any(path.parent != RAW for path in SOURCES):
        raise AssertionError("A figure used data outside the current run's raw directory")


def developer_weighted_map(values: dict) -> float:
    groups = {}
    for tag, score in values["maps"].items():
        if values["coverage"][tag] < 300:
            continue
        developer = MODELS[tag].split("/", 1)[0]
        groups.setdefault(developer, []).append(score)
    return mean(mean(scores) for scores in groups.values())


def fallback_table(name: str) -> str:
    file = {"freeform": "behavior", "persuasion": "behavior", "tradeoff": "intervention",
            "source_symmetry": "intervention", "extractors": "robustness", "prompt_phrasing": "robustness",
            "induction_cost": "causal_controls", "steering_cost": "causal_controls"}.get(name, name)
    data = json.loads((DATA / f"{file}.json").read_text(encoding="utf-8"))
    if name == "coverage":
        headings = ("Checkpoint", "Selected questions")
        records = [(row["model"], row["selected"]) for row in data["rows"]]
    elif name == "maps":
        headings = ("Checkpoint", "Mean signed r", "Pairs below 0.6")
        records = [(row["model"], f"{row['mean']:.3f}", row["pairs_below_0_6"]) for row in data["models"]]
    elif name == "map_pairs":
        headings = ("Checkpoint", "Pairs below 0.6", "Of these, involving YO")
        records = [(row["model"], row["pairs_below_0_6"], row["involving_yo"])
                   for row in data["weak_pairs"]]
    elif name == "item_fidelity":
        headings = ("Checkpoint", "Mean per-item r", "Languages with r < 0")
        records = [(row["model"], f"{row['mean_item_r']:.3f}", row["negative_languages"])
                   for row in data["models"]]
    elif name == "patch_alignment":
        headings = ("Checkpoint", "Language", "N", "Attribution/exact map r")
        records = [(row["model"], row["language"], row["n"], f"{row['head_r']:.3f}")
                   for row in data["rows"]]
    elif name == "patch":
        headings = ("Checkpoint", "Language", "N", "Denoising", "Reverse")
        records = [(row["model"], row["language"], row["n"], f"{row['denoise']['value']:.3f}",
                    f"{row['reverse']['value']:.3f}") for row in data["rows"]]
    elif name == "specificity":
        headings = ("Checkpoint", "Speaker r", "Answer r", "Social r", "Residual signed r")
        records = [(row["model"], *(f"{row[key]:.3f}" for key in
                    ("speaker_change", "answer_map", "social_pressure", "residual_cross_language")))
                   for row in data["rows"]]
    elif name == "representation":
        headings = ("Checkpoint", "Factual AUROC", "Social AUROC", "Language control")
        records = [(row["model"], *(f"{row[key]:.3f}" for key in
                    ("factual_auroc", "social_auroc", "language_control"))) for row in data["rows"]]
    elif name == "extractors":
        headings = ("Checkpoint", "DiffMean", "Logistic", "LDA", "PCA")
        records = [(row["model"], *(f"{row['methods'][key]:.3f}" for key in
                    ("diffmean", "logistic", "lda", "pca"))) for row in data["extractors"]]
    elif name == "prompt_phrasing":
        headings = ("Checkpoint", "Mean off-diagonal", "Weakest pair", "N")
        records = [(row["model"], f"{row['mean_off_diagonal']:.3f}",
                    f"{row['min_off_diagonal']:.3f}", row["n"]) for row in data["phrasing"]]
    elif name == "judge_diagnostics":
        headings = ("Checkpoint", "Judge disagreement", "Judge/letter agreement", "All/cross-family agreement")
        records = [(row["model"], *(f"{row[key]:.3f}" for key in
                    ("judge_disagreement", "judge_letter_agreement", "all_cross_family_agreement")))
                   for row in data["models"]]
    elif name == "freeform":
        headings = ("Language", "N", "Judged caving", "Letter caving", "Judge disagreement")
        records = [(row["language"], row["n"], *(f"{row[key]:.3f}" for key in
                    ("judged_caving", "letter_caving", "judge_disagreement")))
                   for row in data["rows"] if row["tag"] == "q25_7b"]
    elif name == "persuasion":
        headings = ("Turn", "P(correct)", "Wrong-option rate")
        row = next(row for row in data["curves"] if row["tag"] == "q25_7b" and row["language"] == "EN")
        records = [(turn, f"{row['p_correct'][i]:.3f}", f"{row['wrong_option'][i]:.3f}")
                   for i, turn in enumerate(row["turns"])]
    elif name == "induction_cost":
        headings = ("Checkpoint", "Wrong-option increase at α=4 (pp)", "Neutral accuracy lost (pp)",
                    "Matched-random wrong increase (pp)", "Matched-random accuracy lost (pp)")
        records = [(row["model"], *(f"{row[key]:.2f}" for key in
                    ("wrong_increase_pp", "accuracy_loss_pp",
                     "random_wrong_increase_pp", "random_accuracy_loss_pp")))
                   for row in data["induction_alpha4"]]
    elif name == "steering_cost":
        headings = ("Checkpoint", "Wrong-option change at α=+2 (pp)", "Neutral accuracy lost (pp)")
        records = [(row["model"], f"{row['wrong_increase_pp']:.2f}", f"{row['accuracy_loss_pp']:.2f}")
                   for row in data["steering_alpha2"]]
    elif name == "source_symmetry":
        headings = ("Checkpoint", "Head source", "Targets improved", "Less caving (pp)", "Accuracy lost (pp)")
        records = [(row["model"], row["source"], row["reduced_targets"],
                    f"{row['wrong_reduction_pp']:.2f}", f"{row['accuracy_loss_pp']:.2f}")
                   for row in data["source_aggregate"]]
    elif name == "intervention":
        headings = ("Checkpoint", "Less caving (pp)", "Neutral accuracy lost (pp)")
        records = []
        for tag in TAGS:
            subset = [row for row in data["rows"] if row["tag"] == tag]
            records.append((LABELS[tag],
                            f"{100 * mean(row['baseline_caving'] - row['replacement_caving'] for row in subset):.2f}",
                            f"{100 * mean(row['baseline_accuracy'] - row['replacement_accuracy'] for row in subset):.2f}"))
    else:
        headings = ("Checkpoint", "Sweep", "Setting", "Less caving (pp)", "Accuracy lost (pp)")
        records = [(row["model"], row["kind"], row["setting"],
                    f"{row['caving_reduction_pp']:.2f}", f"{row['accuracy_loss_pp']:.2f}")
                   for row in data["tradeoff"]]
    header = "".join(f'<th scope="col">{html.escape(str(value))}</th>' for value in headings)
    body = "".join("<tr>" + "".join(f"<td>{html.escape(str(value))}</td>" for value in row) + "</tr>"
                   for row in records)
    return f"<table><thead><tr>{header}</tr></thead><tbody>{body}</tbody></table>"


def article(template: str, values: dict) -> None:
    weak_pairs = json.loads((DATA / "map_pairs.json").read_text())["weak_pairs"]
    item_fidelity = {row["tag"]: row["mean_item_r"] for row in values["faithfulness"]["models"]}
    patch_r = [row["head_r"] for row in values["faithfulness"]["patch_rows"]]
    prompt_means = [row["mean_off_diagonal"] for row in values["robustness"]["phrasing"]]
    judges = values["judges"]["models"]
    mistral_induction = next(row for row in values["causal_controls"]["induction_alpha4"]
                             if row["tag"] == "mistral7b")
    replacements = {
        "MAP_R": f"{developer_weighted_map(values):.2f}",
        "PATCH_CELLS": str(len(CONFIG["patch_models"]) * len(CONFIG["patch_languages"])),
        "REDUCED": str(values["interventions"]["reduced"]),
        "CELLS": str(values["interventions"]["cells"]),
        "SOURCE_REDUCED": str(values["interventions"]["source_reduced"]),
        "MODEL_COUNT": str(len(TAGS)), "LANG_COUNT": str(len(LANGS)),
        "LOW_N": str(values["coverage"]["llama8b"]),
        "TRANSFER_LOW": f"{values['representation']['min']:.2f}",
        "TRANSFER_HIGH": f"{values['representation']['max']:.2f}",
        "ANSWER_R": f"{values['specificity']['answer_map']:.2f}",
        "THIRD_R": f"{values['specificity']['speaker_change']:.2f}",
        "SOCIAL_R": f"{values['specificity']['social_pressure']:.2f}",
        "RESIDUAL_R": f"{values['specificity']['residual_cross_language']:.2f}",
        "WEAK_PAIRS": str(sum(row["pairs_below_0_6"] for row in weak_pairs)),
        "YO_WEAK": str(sum(row["involving_yo"] for row in weak_pairs)),
        "ITEM_Q35_SMALL": f"{item_fidelity['q35_08b']:.2f}",
        "ITEM_Q25_SEVEN": f"{item_fidelity['q25_7b']:.2f}",
        "PATCH_AGREE": str(sum(value > .9 for value in patch_r)),
        "PATCH_OUTLIER": f"{min(patch_r):.2f}",
        "LOGISTIC_BETTER": str(sum(row["methods"]["logistic"] > row["methods"]["diffmean"]
                                   for row in values["robustness"]["extractors"])),
        "PROMPT_MIN": f"{min(prompt_means):.2f}",
        "PROMPT_MAX": f"{max(prompt_means):.2f}",
        "JUDGE_DISAGREE": f"{mean(row['judge_disagreement'] for row in judges):.0%}",
        "JUDGE_LETTER": f"{mean(row['judge_letter_agreement'] for row in judges):.0%}",
        "INDUCE_MISTRAL_WRONG": f"{mistral_induction['wrong_increase_pp']:.1f}",
        "INDUCE_MISTRAL_COST": f"{mistral_induction['accuracy_loss_pp']:.1f}",
    }
    for key, value in replacements.items():
        template = template.replace(f"@@{key}@@", html.escape(value))
    for name in ("coverage", "maps", "map_pairs", "item_fidelity", "patch", "patch_alignment",
                 "specificity", "representation", "extractors", "prompt_phrasing",
                 "freeform", "judge_diagnostics", "persuasion", "intervention", "tradeoff",
                 "source_symmetry", "induction_cost", "steering_cost"):
        before = f'<details class="data-table" data-table="{name}"><summary>View values as a table</summary><div></div></details>'
        after = before.replace("<div></div>", f"<div>{fallback_table(name)}</div>")
        if before not in template:
            raise AssertionError(f"Missing table placeholder: {name}")
        template = template.replace(before, after)
    if "@@" in template:
        raise AssertionError("An article placeholder was not filled")
    (SPACE / "index.html").write_text(template, encoding="utf-8")
    (DATA / "numbers.json").write_text(json.dumps(replacements, sort_keys=True, indent=2) + "\n")


def main() -> None:
    if not RAW.is_dir():
        raise FileNotFoundError(RAW)
    states = {"coverage": coverage(), "maps": maps()}
    patching()
    states["faithfulness"] = faithfulness()
    states["interventions"] = interventions()
    states["specificity"] = specificity()
    states["representation"] = representations()
    states["robustness"] = robustness()
    behavior()
    states["judges"] = judge_diagnostics()
    states["causal_controls"] = causal_controls()
    check(states)
    all_paths = sorted(SOURCES)
    (DATA / "source_manifest.json").write_text(json.dumps({
        "run_id": "paper-final", "source_file_count": len(all_paths),
        "sources": [{"path": f"data/raw/paper-final/{path.name}", "sha256": fingerprint(path)}
                    for path in all_paths],
        "saved_summary_used_for": "consistency check only; not a chart input",
    }, indent=2, sort_keys=True) + "\n")
    article((SPACE / "index.template.html").read_text(encoding="utf-8"), states)
    print(f"Built {len(list(FIGURES.glob('*.svg')))} SVGs and chart JSON from {len(all_paths)} raw files; summary spot-checks passed")


if __name__ == "__main__":
    main()
