"""
Run every experiment in the post and write the artifacts it uses.

    .venv/bin/python run_experiments.py

Writes assets/e1..e5-*.png (charts), results/summary.json (headline numbers,
curves, tables) and results/walkthrough.json (one complete GRPO update, with
every number). Everything is seeded: two runs produce identical numbers.
"""
import json
from pathlib import Path

import matplotlib

matplotlib.use("Agg")
import matplotlib.pyplot as plt  # noqa: E402
import numpy as np  # noqa: E402
from matplotlib.patches import PathPatch, Rectangle  # noqa: E402
from matplotlib.path import Path as BezierPath  # noqa: E402

import lava_leap as ll  # noqa: E402

ROOT = Path(__file__).resolve().parent
ASSETS, RESULTS = ROOT / "assets", ROOT / "results"

# ---- configuration (the hyperparameters reported in the post) ----------------
CFG = dict(lr=0.5, eps=0.2, G=8, passes=4, rounds=300, eval_every=10)
LR_CRITIC = 0.5                    # ppo_update's default; reported alongside CFG
SEEDS = list(range(5))             # five independent training runs per algorithm
TEST_SEED, N_TEST = 2024, 200      # the fixed test floors every policy is scored on
SHOWCASE_SEED = 8                  # E1: first seed whose floor has 6 channels, 2 of them 2-wide
WALK_SEED, WALK_AT = 22, 20        # E3a/E4: a hard floor where all 8 runs die and the clip
                                   # engages; GRPO policy frozen after 20 updates (seed 0)
BLIND_SEED, BLIND_DRAWS = 31, 40   # E3b: floors the trained PPO bot plays, 8 runs each
SIM_PLAYS = 50                     # Monte Carlo cross-check: plays per test floor
LR_SWEEP = {"reinforce": [1.0, 1.5], "ppo": [1.0, 1.5], "grpo": [1.0]}
ALGOS = ["reinforce", "ppo", "grpo"]

# ---- house style ---------------------------------------------------------------
TEAL, BLUE, GRAY, CEIL = "#2D7A6E", "#5A82B5", "#94A3B8", "#475569"
LAVA, SAFE, SOFT, INK, MUTED = "#C06A5A", "#E2E8F0", "#D3E8E3", "#0F172A", "#64748B"
COLOR = {"grpo": TEAL, "ppo": BLUE, "reinforce": GRAY}
NAME = {"grpo": "GRPO", "ppo": "PPO", "reinforce": "REINFORCE"}
PATTERN_NAME = {
    0b000: "no lava in view", 0b001: "lava 3 ahead", 0b010: "lava 2 ahead",
    0b011: "2-wide lava, 2 ahead", 0b100: "lava directly ahead",
    0b110: "2-wide lava, directly ahead", 0b101: "two channels, one tile apart",
    0b111: "a 3-wide channel",
}
NEVER = (0b101, 0b111)             # impossible on a real floor (gap >= 2, width <= 2)

plt.rcParams.update({
    "font.family": ["Noto Sans", "DejaVu Sans"],
    "figure.facecolor": "white", "axes.facecolor": "white", "savefig.facecolor": "white",
    "axes.spines.top": False, "axes.spines.right": False,
    "axes.edgecolor": MUTED, "axes.linewidth": 0.8,
    "axes.grid": True, "grid.color": INK, "grid.alpha": 0.15, "grid.linewidth": 0.8,
    "axes.titlesize": 15, "axes.titleweight": "bold", "axes.titlecolor": INK,
    "axes.titlelocation": "left", "axes.titlepad": 12,
    "axes.labelsize": 12, "axes.labelcolor": INK,
    "xtick.labelsize": 11, "ytick.labelsize": 11, "xtick.color": MUTED, "ytick.color": MUTED,
    "xtick.major.size": 0, "ytick.major.size": 0, "xtick.major.pad": 6, "ytick.major.pad": 6,
    "text.color": INK, "legend.frameon": False, "legend.fontsize": 10,
})


def words(n):
    return ["zero", "one", "two", "three", "four", "five", "six", "seven", "eight",
            "nine"][n] if 0 <= n <= 9 else str(n)


def header(fig, title, subtitle, x=0.06, y=0.965):
    """Left-aligned takeaway title plus a muted one-line subtitle, in figure coordinates."""
    fig.text(x, y, title, fontsize=15, fontweight="bold", color=INK, ha="left", va="top")
    fig.text(x, y - 0.068, subtitle, fontsize=10.5, color=MUTED, ha="left", va="top")


def save(fig, name):
    fig.savefig(ASSETS / name, dpi=200)
    plt.close(fig)
    print("wrote", ASSETS / name)


def pattern_str(s):
    return "".join(str(int(x)) for x in s[:3])


def rnd(x):
    """Numbers saved to the JSON files keep 6 decimals, so every rounding happens once."""
    return float(round(float(x), 6))


def floor_info(lava):
    """Lava tiles and channels (start, width) of a floor, for the JSON files."""
    tiles = [int(t) for t in np.flatnonzero(lava)]
    channels, t = [], 0
    while t < ll.FLOOR:
        if lava[t]:
            w = 1
            while t + w < ll.FLOOR and lava[t + w]:
                w += 1
            channels.append([t, w])
            t += w
        else:
            t += 1
    return dict(lava_tiles=tiles, channels=channels,
                layout="".join("#" if x else "." for x in lava))


def run_info(run, adv=None):
    d = dict(path=[int(p) for p in run.path], final_tile=int(run.path[-1]),
             reward=rnd(run.reward), steps=len(run.actions),
             moves="".join("SL"[a] for a in run.actions),
             patterns=[pattern_str(s) for s in run.states],
             reached_exit=bool(run.reward == 1.0))
    if adv is not None:
        d["advantage"] = rnd(adv)
    return d


def fall_kind(run, lava):
    """How a run ended: 'exit'; 'coin flip' (a LEAP from the tile just before a two-tile
    channel that landed on its far half, the one loss no policy can avoid); or 'misstep'."""
    end = run.path[-1]
    if run.reward == 1.0:
        return "exit"
    start = run.path[-2]
    wide = end >= 2 and lava[end - 1] and (end - 2 >= 0 and not lava[end - 2])
    if wide and run.actions[-1] == ll.LEAP and start == end - 2:
        return "coin flip"
    return "misstep"


def group_adv(runs):
    r = np.array([run.reward for run in runs])
    return r, (r - r.mean()) / (r.std() + 1e-8)


# ---- drawing helpers -------------------------------------------------------------
def draw_floor(ax, lava, y, h=0.36, x_of=lambda t: t, w=0.84):
    for t in range(ll.FLOOR):
        ax.add_patch(Rectangle((x_of(t) - w / 2, y - h / 2), w, h, lw=0,
                               color=LAVA if lava[t] else SAFE, zorder=1))


def draw_hops(ax, path, y, color, lw=1.3):
    """One arc per move: low for a STEP, taller for a LEAP."""
    for a, b in zip(path[:-1], path[1:]):
        lift = 0.16 + 0.11 * (b - a)
        verts = [(a, y + 0.2), ((a + b) / 2, y + 0.2 + 2 * lift), (b, y + 0.2)]
        codes = [BezierPath.MOVETO, BezierPath.CURVE3, BezierPath.CURVE3]
        ax.add_patch(PathPatch(BezierPath(verts, codes), fill=False, ec=color, lw=lw,
                               capstyle="round", zorder=3))


# ---- E1: before / after on the showcase floor -----------------------------------
def chart_e1(lava, before, after, title):
    fig, axes = plt.subplots(2, 1, figsize=(9.5, 4.6), sharex=True)
    fig.subplots_adjust(left=0.06, right=0.885, top=0.77, bottom=0.13, hspace=0.5)
    panels = [(axes[0], before, MUTED, "Untrained (all eight weights zero)"),
              (axes[1], after, TEAL, "Trained with GRPO (300 updates)")]
    for ax, runs, color, label in panels:
        n_exit = sum(run.reward == 1.0 for run in runs)
        for i, run in enumerate(runs):
            y = len(runs) - 1 - i
            draw_floor(ax, lava, y)
            draw_hops(ax, run.path, y, color)
            end = run.path[-1]
            if run.reward == 1.0:
                ax.scatter([ll.EXIT], [y], s=46, color=TEAL, edgecolor="white", lw=1.2, zorder=5)
                ax.text(40.4, y, "exit", va="center", fontsize=10, color=INK, fontweight="bold")
            else:
                ax.scatter([end], [y], marker="x", s=34, color=INK, lw=1.8, zorder=5)
                ax.text(40.4, y, f"fell at {end}", va="center", fontsize=10, color=MUTED)
        ax.set_title(f"{label}:  reached the exit {n_exit} of {len(runs)}",
                     fontsize=11.5, fontweight="bold", pad=4, color=INK)
        ax.set_ylim(-0.6, len(runs) - 0.25)
        ax.set_xlim(-0.8, 39.8)
        ax.set_yticks([])
        ax.grid(False)
        ax.spines["left"].set_visible(False)
        ax.spines["bottom"].set_visible(False)
    axes[1].set_xticks([0, 5, 10, 15, 20, 25, 30, 35, 39])
    axes[1].set_xlabel("Tile  (start = 0, exit dock = 39)", fontsize=11)
    header(fig, title, "The same floor, played eight times by each bot. Low arcs are steps, "
           "tall arcs are leaps; red tiles are lava.")
    save(fig, "e1-before-after.png")


# ---- E2: training curves -------------------------------------------------------------
def chart_e2(updates, curves, ceiling, untrained, title):
    fig, ax = plt.subplots(figsize=(9.5, 4.5))
    fig.subplots_adjust(left=0.09, right=0.80, top=0.80, bottom=0.14)
    finals = {}
    for algo in ALGOS:
        Y = np.array(curves[algo])
        for y in Y:
            ax.plot(updates, y, color=COLOR[algo], lw=0.8, alpha=0.3, zorder=2)
        mean = Y.mean(axis=0)
        ax.plot(updates, mean, color=COLOR[algo], lw=2.8 if algo == "grpo" else 2.0, zorder=3,
                solid_capstyle="round")
        ax.scatter([updates[-1]], [mean[-1]], s=30, color=COLOR[algo], zorder=4)
        finals[algo] = mean[-1]
    ax.axhline(ceiling, color=CEIL, lw=1.4, ls=(0, (5, 4)), zorder=1)
    # direct labels, nudged apart if they would collide
    label_y = {a: finals[a] for a in ALGOS}
    label_y["ceiling"] = ceiling
    order = sorted(label_y, key=label_y.get)
    for lo, hi in zip(order, order[1:]):
        if label_y[hi] - label_y[lo] < 0.024:
            label_y[hi] = label_y[lo] + 0.024
    for algo in ALGOS:
        ax.text(updates[-1] + 7, label_y[algo], f"{NAME[algo]}  {finals[algo]:.3f}",
                va="center", fontsize=11, color=INK, fontweight="bold" if algo == "grpo" else "normal",
                clip_on=False)
    ax.text(updates[-1] + 7, label_y["ceiling"], f"Best possible  {ceiling:.3f}", va="center",
            fontsize=11, color=CEIL, clip_on=False)
    ax.annotate(f"untrained {untrained:.3f}", xy=(0, untrained), xytext=(14, untrained - 0.03),
                fontsize=10, color=MUTED, arrowprops=dict(arrowstyle="-", color=MUTED, lw=0.8))
    ax.set_xlim(0, updates[-1])
    ax.set_ylim(0.24, 0.6)
    ax.set_xlabel("Training updates (each: one new floor, played 8 times)")
    ax.set_ylabel("Mean score, 200 test floors")
    header(fig, title, "Thick lines: mean of 5 training seeds; thin lines: each seed. "
           "Dashed: the best of all 256 fixed policies.", x=0.09)
    save(fig, "e2-training-curves.png")


# ---- E3a: grading on a curve ---------------------------------------------------------
def chart_e3a(runs, r, adv, title):
    order = sorted(range(len(runs)), key=lambda i: (r[i], i))
    mean, std = r.mean(), r.std()
    fig, ax = plt.subplots(figsize=(9.5, 4.4))
    fig.subplots_adjust(left=0.09, right=0.855, top=0.79, bottom=0.17)
    ax.axhspan(mean - std, mean + std, color=SOFT, lw=0, zorder=0)
    ax.axhline(mean, color=INK, lw=1.2, zorder=2)
    x = np.arange(len(runs))
    ax.bar(x, r[order], width=0.62, color=[TEAL if adv[i] > 0 else LAVA for i in order], zorder=3)
    for xi, i in zip(x, order):
        ax.text(xi, r[i] + 0.025, f"{adv[i]:+.2f}", ha="center", va="bottom", fontsize=11,
                color=INK, fontweight="bold" if adv[i] > 0 else "normal", zorder=4,
                bbox=dict(boxstyle="square,pad=0.15", fc="white", ec="none"))
    ax.set_xticks(x, [f"run {i + 1}\ntile {runs[i].path[-1]}" for i in order])
    ax.text(len(runs) - 0.45, mean, f"group mean {mean:.2f}", va="center", fontsize=10.5,
            color=INK, clip_on=False)
    ax.text(len(runs) - 0.45, mean + std, f"+1 std  {mean + std:.2f}", va="center",
            fontsize=10, color=MUTED, clip_on=False)
    ax.text(len(runs) - 0.45, mean - std, f"-1 std  {mean - std:.2f}", va="center",
            fontsize=10, color=MUTED, clip_on=False)
    ax.set_xlim(-0.6, len(runs) - 0.5)
    ax.set_ylim(0, max(r) + 0.16)
    ax.set_ylabel("Score  r = final tile / 39")
    ax.grid(axis="x", visible=False)
    header(fig, title, "One hard floor, eight runs, every one ended in lava. Above each bar: "
           "its advantage, (r - mean) / std.", x=0.09)
    save(fig, "e3a-group-advantage.png")


# ---- E3b: the critic's blind spot ----------------------------------------------------
def chart_e3b(easy, hard, easy_runs, hard_runs, forecast, title):
    fig, ax = plt.subplots(figsize=(9.5, 4.4))
    fig.subplots_adjust(left=0.18, right=0.97, top=0.78, bottom=0.14)
    rows = [(2.25, easy, easy_runs, "Easy floor"), (0.0, hard, hard_runs, "Hard floor")]
    for y0, lava, runs, label in rows:
        draw_floor(ax, lava, y0, h=0.3)
        ends = [run.path[-1] for run in runs]
        stack = {}
        for e in sorted(ends):
            k = stack.get(e, 0)
            stack[e] = k + 1
            ax.scatter([e], [y0 + 0.36 + 0.2 * k], s=36, color=INK, edgecolor="white", lw=1.0,
                       zorder=5)
        mean_tile = np.mean([run.reward for run in runs]) * ll.EXIT
        ax.scatter([mean_tile], [y0 - 0.34], marker="^", s=70, color=TEAL, zorder=5, clip_on=False)
        ax.text(mean_tile + 0.6, y0 - 0.36, f"group mean {mean_tile / ll.EXIT:.2f}", fontsize=10.5,
                color=INK, va="center")
        worst, best = min(ends), max(ends)
        ax.text(-1.4, y0 + 0.12, label, ha="right", va="center", fontsize=12, fontweight="bold",
                color=INK)
        ax.text(-1.4, y0 - 0.2, f"scores {worst / ll.EXIT:.2f} to {best / ll.EXIT:.2f}",
                ha="right", va="center", fontsize=10, color=MUTED)
    xf = forecast * ll.EXIT
    ax.plot([xf, xf], [-0.55, 3.55], color=BLUE, lw=1.6, ls=(0, (4, 3)), zorder=2)
    ax.text(xf - 0.5, 3.55, f"PPO critic's forecast from the start line: {forecast:.2f}",
            fontsize=10.5, color=INK, va="center", ha="right")
    ax.text(xf - 0.5, 3.25, "(the same number for every floor)", fontsize=10, color=MUTED,
            va="center", ha="right")
    ax.set_xlim(-0.8, 39.8)
    ax.set_ylim(-0.62, 3.7)
    ax.set_yticks([])
    ax.set_xticks([0, 5, 10, 15, 20, 25, 30, 35, 39])
    ax.set_xlabel("Final tile of each run  (score = tile / 39)", fontsize=11)
    ax.grid(False)
    ax.spines["left"].set_visible(False)
    ax.spines["bottom"].set_visible(False)
    header(fig, title, "The trained PPO bot plays each floor eight times; each dot is where a "
           "run ended. Both floors look the same from the start.", x=0.035)
    save(fig, "e3b-critic-blindspot.png")


# ---- E4: ratios across the four passes -----------------------------------------------
def chart_e4(trace, A, labels, title):
    fig, ax = plt.subplots(figsize=(9.5, 4.4))
    fig.subplots_adjust(left=0.09, right=0.79, top=0.79, bottom=0.11)
    eps = CFG["eps"]
    ax.axhspan(1 - eps, 1 + eps, color=SOFT, lw=0, zorder=0)
    jitter = np.random.default_rng(0).uniform(-0.22, 0.22, len(A))
    sign_color = np.array([TEAL if a > 0 else LAVA for a in A])
    for k, p in enumerate(trace):
        x = k + 1 + jitter
        free, cut = ~p["clipped"], p["clipped"]
        ax.scatter(x[free], p["rho"][free], s=18, lw=0, alpha=0.85, zorder=3, color=sign_color[free])
        ax.scatter(x[cut], p["rho"][cut], s=72, facecolor="white", edgecolor=sign_color[cut],
                   lw=2.0, zorder=4)
    for txt, (k, j) in labels:
        ax.annotate(txt, xy=(k + 1 + jitter[j], trace[k]["rho"][j]),
                    xytext=(4.5, trace[k]["rho"][j]), fontsize=10, color=INK, va="center",
                    arrowprops=dict(arrowstyle="-", color=MUTED, lw=0.8, shrinkA=0, shrinkB=6),
                    annotation_clip=False)
    ax.text(0.53, 1 + eps + 0.008, f"clip band: {1 - eps:.1f} to {1 + eps:.1f}", fontsize=10,
            color=TEAL, va="bottom")
    ax.set_xticks([1, 2, 3, 4], ["pass 1", "pass 2", "pass 3", "pass 4"])
    ax.set_xlim(0.5, 4.4)
    lo = min(p["rho"].min() for p in trace)
    hi = max(p["rho"].max() for p in trace)
    ax.set_ylim(min(0.76, lo - 0.03), max(1.24, hi + 0.03))
    ax.set_ylabel("ratio  π_new(a|s) / π_old(a|s)")
    ax.grid(axis="x", visible=False)
    header(fig, title, f"Each dot is one of the group's {len(A)} steps. Teal: its run beat "
           "the group mean; red: fell short. Rings: clipped, no push.", x=0.09)
    save(fig, "e4-ratio-spread.png")


# ---- E5: what the policy learned ---------------------------------------------------
def chart_e5(before, after, title):
    order = [0b100, 0b110, 0b000, 0b001, 0b010, 0b011, 0b101, 0b111]
    fig, ax = plt.subplots(figsize=(9.5, 4.6))
    fig.subplots_adjust(left=0.40, right=0.93, top=0.78, bottom=0.12)
    ys = np.arange(len(order))[::-1].astype(float)
    ys[-2:] -= 0.5                                    # set the impossible patterns apart
    for y, b in zip(ys, order):
        never = b in NEVER
        ax.barh(y + 0.17, before[b], height=0.3, color=SAFE, zorder=2)
        ax.barh(y - 0.17, after[b], height=0.3, color=GRAY if never else TEAL, zorder=2)
        ax.text(after[b] + 0.012, y - 0.17, f"{after[b]:.3f}", va="center", fontsize=10.5,
                color=MUTED if never else INK, fontweight="normal" if never else "bold")
        for j in range(3):                            # the three sensor tiles, drawn as tiles
            ax.scatter([-0.15 + 0.04 * j], [y], marker="s", s=110, clip_on=False, zorder=4,
                       color=LAVA if (b >> (2 - j)) & 1 else SAFE,
                       transform=ax.get_yaxis_transform())
        ax.text(-0.19, y, PATTERN_NAME[b], ha="right", va="center", fontsize=11,
                color=MUTED if never else INK, transform=ax.get_yaxis_transform())
    ax.text(before[0b100], ys[0] + 0.5, "0.5 before training, for every pattern", ha="center",
            va="center", fontsize=10, color=MUTED)
    ax.text(-0.19, (ys[-3] + ys[-2]) / 2, "never occur on a real floor:", ha="right",
            va="center", fontsize=10, color=MUTED, style="italic", transform=ax.get_yaxis_transform())
    ax.set_yticks([])
    ax.set_xlim(0, 1.0)
    ax.set_ylim(ys[-1] - 0.55, ys[0] + 0.85)
    ax.set_xlabel("P(LEAP): probability the bot leaps on this pattern", fontsize=11)
    ax.grid(axis="y", visible=False)
    ax.spines["left"].set_visible(False)
    ax.text(-0.11, ys[0] + 0.62, "+1  +2  +3", ha="center", va="center", fontsize=9.5,
            color=MUTED, transform=ax.get_yaxis_transform())
    header(fig, title, "Leap probability for each of the 8 sensor patterns: before training "
           "(light bars) and after 300 GRPO updates (teal).", x=0.035)
    save(fig, "e5-policy-learned.png")


# ---- the experiments -----------------------------------------------------------------
def main():
    ASSETS.mkdir(exist_ok=True)
    RESULTS.mkdir(exist_ok=True)
    test_rng = np.random.default_rng(TEST_SEED)
    test = [ll.new_floor(test_rng) for _ in range(N_TEST)]
    summary = dict(config=dict(CFG, lr_critic=LR_CRITIC, seeds=SEEDS, test_seed=TEST_SEED,
                               n_test_floors=N_TEST, adv_eps=1e-8, std="population (ddof=0)"))

    # Untrained baseline and the exact ceiling (best of all 256 deterministic policies).
    W0 = np.zeros((4, 2))
    untrained_score, untrained_exit = ll.evaluate(W0, test)
    scores, exits = ll.all_deterministic_policies(test)
    best = np.flatnonzero(np.isclose(scores, scores.max(), atol=1e-12))
    ceiling, ceiling_exit = float(scores.max()), float(exits[best[0]])
    leap_iff_lava_next = sum(1 << b for b in range(8) if b >> 2 & 1)
    sim_rng = np.random.default_rng(7)
    best_table = np.array([(best[0] >> b) & 1 for b in range(8)], dtype=float)
    summary["untrained"] = dict(score=rnd(untrained_score), exit_rate=rnd(untrained_exit),
                                sim=dict(zip(["score", "exit_rate"], map(rnd, ll.simulate(
                                    ll.leap_table(W0), test, SIM_PLAYS, sim_rng)))))
    summary["ceiling"] = dict(
        score=rnd(ceiling), exit_rate=rnd(ceiling_exit), n_policies=256, n_tied_best=int(len(best)),
        best_policy_ids=[int(b) for b in best],
        best_policies_leap_on=[[format(b, "03b") for b in range(8) if (pid >> b) & 1] for pid in best],
        leap_iff_lava_next_is_best=bool(leap_iff_lava_next in best),
        second_best_score=rnd(np.max(scores[scores < ceiling - 1e-9])),
        sim=dict(zip(["score", "exit_rate"], map(rnd, ll.simulate(best_table, test, SIM_PLAYS,
                                                                   sim_rng)))))
    print(f"untrained {untrained_score:.4f}  ceiling {ceiling:.4f}  ({len(best)} tied)")

    # E2: train all three algorithms on five seeds.
    curves, finals, Ws, clip_logs, snaps = {}, {}, {}, {}, {}
    for algo in ALGOS:
        curves[algo], finals[algo] = [], []
        for seed in SEEDS:
            log = [] if algo != "reinforce" else None
            W, v, curve, saved = ll.train(algo, seed, eval_floors=test, clip_log=log,
                                          snapshots=(10, WALK_AT), **CFG)
            curves[algo].append([c[1] for c in curve])
            finals[algo].append(dict(seed=seed, score=rnd(curve[-1][1]), exit_rate=rnd(curve[-1][2])))
            Ws[algo, seed] = (W, v)
            snaps[algo, seed] = saved
            if log is not None:
                clip_logs[algo, seed] = log
        print(algo, [f["score"] for f in finals[algo]])
    updates = [c for c in range(0, CFG["rounds"] + 1, CFG["eval_every"])]
    summary["curves"] = dict(updates=updates, **{a: dict(
        per_seed=[[rnd(x) for x in c] for c in curves[a]],
        mean=[rnd(x) for x in np.mean(curves[a], axis=0)]) for a in ALGOS})
    summary["final"] = {}
    for algo in ALGOS:
        s = np.array([f["score"] for f in finals[algo]])
        e = np.array([f["exit_rate"] for f in finals[algo]])
        Wfin = Ws[algo, 0][0]
        summary["final"][algo] = dict(
            score_mean=rnd(s.mean()), score_min=rnd(s.min()), score_max=rnd(s.max()),
            exit_mean=rnd(e.mean()), exit_min=rnd(e.min()), exit_max=rnd(e.max()),
            pct_of_ceiling=rnd(100 * s.mean() / ceiling), per_seed=finals[algo],
            seed0_sim=dict(zip(["score", "exit_rate"], map(rnd, ll.simulate(
                ll.leap_table(Wfin), test, SIM_PLAYS, sim_rng)))))
    # how many updates it takes each algorithm (5-seed mean curve) to reach 90% of the ceiling
    for algo in ALGOS:
        m = np.mean(curves[algo], axis=0)
        hit = np.flatnonzero(m >= 0.9 * ceiling)
        summary["final"][algo]["updates_to_90pct"] = int(updates[hit[0]]) if len(hit) else None
        summary["final"][algo]["score_at_50"] = rnd(m[updates.index(50)])
        summary["final"][algo]["score_at_100"] = rnd(m[updates.index(100)])

    # The slowest GRPO seed: what its policy looked like while it was stuck.
    at100 = updates.index(100)
    slow = int(np.argmin([c[at100] for c in curves["grpo"]]))
    sc = np.array(curves["grpo"][slow])
    plateau = sc[updates.index(10):updates.index(140) + 1]
    release = next(u for u, x in zip(updates, sc) if u > 10 and x >= plateau.mean() + 0.05)
    lt10 = ll.leap_table(snaps["grpo", slow][10][0])
    summary["grpo_slow_seed"] = dict(seed=slow, plateau_mean=rnd(plateau.mean()),
                                     plateau_min=rnd(plateau.min()), plateau_max=rnd(plateau.max()),
                                     release_update=int(release),
                                     leap_table_at_10=[rnd(x) for x in lt10],
                                     min_leap_at_10=rnd(lt10[[0, 1, 2, 3, 4, 6]].min()))

    # Learning-rate sensitivity (is GRPO's edge just its larger, std-scaled step?).
    sweep = {}
    for algo, lrs in LR_SWEEP.items():
        for lr in lrs:
            cfg = dict(CFG, lr=lr)
            res = [ll.train(algo, seed, eval_floors=test, **cfg)[2][-1][1] for seed in SEEDS]
            sweep[f"{algo}@{lr}"] = dict(algo=algo, lr=lr, score_mean=rnd(np.mean(res)),
                                         score_min=rnd(np.min(res)), score_max=rnd(np.max(res)))
    summary["lr_sensitivity"] = sweep

    # How often does the clip actually engage during training?
    clip_stats = {}
    for algo in ["grpo", "ppo"]:
        logs = [clip_logs[algo, s] for s in SEEDS]
        clipped = [np.array([u["clipped"] for u in lg]) for lg in logs]
        steps = [np.array([u["steps"] for u in lg]) for lg in logs]
        clip_stats[algo] = dict(
            updates_with_clip_per_seed=[int((c.sum(1) > 0).sum()) for c in clipped],
            share_per_seed=[rnd(c.sum() / (n.sum() * CFG["passes"])) for c, n in zip(clipped, steps)],
            share_of_step_passes_clipped=rnd(sum(c.sum() for c in clipped)
                                            / sum(n.sum() * CFG["passes"] for n in steps)),
            mean_abs_advantage=rnd(np.mean([u["mean_abs_adv"] for lg in logs for u in lg])),
            updates=CFG["rounds"])
    summary["clip_stats"] = clip_stats

    # E1: showcase floor, untrained vs trained GRPO (seed 0).
    W_grpo = Ws["grpo", 0][0]
    show = ll.new_floor(np.random.default_rng(SHOWCASE_SEED))
    rng_b, rng_a = np.random.default_rng([SHOWCASE_SEED, 1]), np.random.default_rng([SHOWCASE_SEED, 2])
    before = [ll.play(W0, show, rng_b) for _ in range(8)]
    after = [ll.play(W_grpo, show, rng_a) for _ in range(8)]
    sv_b, sp_b = ll.exact_outcomes(ll.leap_table(W0), [show])
    sv_a, sp_a = ll.exact_outcomes(ll.leap_table(W_grpo), [show])
    sv_o, sp_o = ll.exact_outcomes(best_table, [show])
    wide = {t for t, w in floor_info(show)["channels"] if w == 2}
    wide_tiles = {t + d for t in wide for d in (0, 1)}
    summary["e1"] = dict(
        floor=floor_info(show), floor_seed=SHOWCASE_SEED,
        before=[run_info(r) for r in before], after=[run_info(r) for r in after],
        before_exits=int(sum(r.reward == 1 for r in before)),
        after_exits=int(sum(r.reward == 1 for r in after)),
        before_mean_tile=rnd(np.mean([r.path[-1] for r in before])),
        after_mean_tile=rnd(np.mean([r.path[-1] for r in after])),
        after_falls_all_at_wide=bool(all(r.reward == 1 or r.path[-1] in wide_tiles for r in after)),
        after_fall_kinds=[fall_kind(r, show) for r in after],
        before_fall_kinds=[fall_kind(r, show) for r in before],
        exact=dict(untrained=dict(score=rnd(sv_b[0]), exit=rnd(sp_b[0])),
                   grpo=dict(score=rnd(sv_a[0]), exit=rnd(sp_a[0])),
                   best=dict(score=rnd(sv_o[0]), exit=rnd(sp_o[0]))))
    e1 = summary["e1"]
    e1_title = ("Untrained, it falls early; trained, it falls only at two-tile channels"
                if e1["after_falls_all_at_wide"] else "Training carries the bot much further")
    chart_e1(show, before, after, e1_title)

    # E2 chart.
    gr = summary["final"]["grpo"]
    e2_title = f"GRPO reaches {gr['pct_of_ceiling']:.0f}% of the best achievable score"
    chart_e2(updates, curves, ceiling, untrained_score, e2_title)

    # E3a + E4: the frozen mid-training GRPO policy on a hard floor.
    W_mid = snaps["grpo", 0][WALK_AT][0]
    walk_floor = ll.new_floor(np.random.default_rng(WALK_SEED))
    rng_w = np.random.default_rng([WALK_SEED, 0])
    group = [ll.play(W_mid, walk_floor, rng_w) for _ in range(CFG["G"])]
    r, adv = group_adv(group)
    n_pos = int((adv > 0).sum())
    e3a_title = (f"Every run fell, yet the {words(n_pos)} that got furthest earn a positive advantage"
                 if all(x.reward < 1 for x in group) else "The runs that got furthest earn a positive advantage")
    chart_e3a(group, r, adv, e3a_title)
    summary["e3a"] = dict(floor=floor_info(walk_floor), floor_seed=WALK_SEED, policy_updates=WALK_AT,
                          rewards=[rnd(x) for x in r], advantages=[rnd(a) for a in adv],
                          mean=rnd(r.mean()), std=rnd(r.std()), n_positive=n_pos,
                          all_died=bool(all(x.reward < 1 for x in group)))

    # E4 walkthrough: replay the update with the trace hook.
    trace = []
    W_after = ll.grpo_update(W_mid, group, trace=trace, lr=CFG["lr"], eps=CFG["eps"],
                             passes=CFG["passes"])
    S, acts, A = ll.batch(group, adv)
    N = len(acts)
    run_of = np.concatenate([[i] * len(x.actions) for i, x in enumerate(group)])
    step_of = np.concatenate([np.arange(len(x.actions)) for x in group])
    pi_old = trace[0]["pi"][np.arange(N), acts]
    # representative step: the most avoidable fatal move among the lowest-scoring runs
    # (the last move of a worst run, choosing the one the policy thought least likely)
    fatal = [int(np.flatnonzero(run_of == i)[-1]) for i in np.flatnonzero(r == r.min())]
    rep = min(fatal, key=lambda j: pi_old[j])
    s = S[rep]
    scores_rep = s @ W_mid
    pi_rep = ll.policy(W_mid, s)
    onehot = np.eye(2)[acts[rep]]
    coef1 = A[rep] * 1.0 / N
    push = np.outer(s, onehot - pi_rep) * coef1

    def step_ref(j):
        return dict(index=int(j), run=int(run_of[j]) + 1, step=int(step_of[j]) + 1,
                    tile=int(group[run_of[j]].path[step_of[j]]), pattern=pattern_str(S[j]),
                    move=["STEP", "LEAP"][acts[j]], advantage=rnd(A[j]), pi_old=rnd(pi_old[j]))

    passes = []
    for k, p in enumerate(trace):
        rho = p["rho"]
        clipped_idx = np.flatnonzero(p["clipped"])
        passes.append(dict(
            k=k + 1, W_before=np.round(p["W"], 4).tolist(),
            W_after=np.round(trace[k + 1]["W"] if k + 1 < len(trace) else W_after, 4).tolist(),
            delta_W=np.round(CFG["lr"] * p["grad"], 5).tolist(), grad=np.round(p["grad"], 5).tolist(),
            rho_min=rnd(rho.min()), rho_max=rnd(rho.max()),
            n_outside_band=int(((rho > 1 + CFG["eps"]) | (rho < 1 - CFG["eps"])).sum()),
            n_clipped=int(len(clipped_idx)),
            clipped=[dict(step_ref(j), rho=rnd(rho[j])) for j in clipped_idx],
            leap_table=[rnd(x) for x in ll.leap_table(p["W"])],
            rep_rho=rnd(rho[rep])))
    final_rho = ll.policy(W_after, S)[np.arange(N), acts] / pi_old
    # the steps whose ratio moved most by the last pass, for named examples
    far = np.argsort(-np.abs(np.log(trace[-1]["rho"])))[:4]
    walkthrough = dict(
        policy=dict(source=f"GRPO seed 0 after {WALK_AT} updates", W=np.round(W_mid, 4).tolist(),
                    leap_table=[rnd(x) for x in ll.leap_table(W_mid)],
                    test_score=rnd(ll.evaluate(W_mid, test)[0])),
        floor=dict(floor_info(walk_floor), seed=WALK_SEED),
        runs=[run_info(x, a) for x, a in zip(group, adv)],
        group=dict(mean=rnd(r.mean()), std=rnd(r.std()), N_steps=int(N),
                   mean_tile=rnd(r.mean() * ll.EXIT), std_tiles=rnd(r.std() * ll.EXIT)),
        hyper=dict(lr=CFG["lr"], eps=CFG["eps"], passes=CFG["passes"]),
        representative=dict(step_ref(rep), s=s.tolist(),
                            active_rows=[int(i) for i in np.flatnonzero(s)],
                            W_rows={str(int(i)): np.round(W_mid[i], 4).tolist() for i in np.flatnonzero(s)},
                            scores=np.round(scores_rep, 4).tolist(),
                            p_step=rnd(pi_rep[0]), p_leap=rnd(pi_rep[1]),
                            onehot_minus_pi=np.round(onehot - pi_rep, 4).tolist(),
                            coef_pass1=float(round(coef1, 6)), push=np.round(push, 5).tolist(),
                            push_lr=np.round(CFG["lr"] * push, 5).tolist(),
                            rho_by_pass=[rnd(p["rho"][rep]) for p in trace],
                            rho_final=rnd(final_rho[rep])),
        passes=passes,
        W_final=np.round(W_after, 4).tolist(),
        leap_table_final=[rnd(x) for x in ll.leap_table(W_after)],
        final_rho_min=rnd(final_rho.min()), final_rho_max=rnd(final_rho.max()),
        biggest_movers=[dict(step_ref(j), rho_pass4=rnd(trace[-1]["rho"][j]),
                             clipped_pass4=bool(trace[-1]["clipped"][j])) for j in far],
        total_clipped_step_passes=int(sum(p["n_clipped"] for p in passes)),
        test_score_after=rnd(ll.evaluate(W_after, test)[0]),
        leap_share=dict(above_mean=rnd(acts[A > 0].mean()), below_mean=rnd(acts[A < 0].mean()),
                        n_above=int((A > 0).sum()), n_below=int((A < 0).sum()),
                        leaps_above=int(acts[A > 0].sum()), leaps_below=int(acts[A < 0].sum())),
        steps=[dict(step_ref(j), rho=[rnd(p["rho"][j]) for p in trace],
                    clipped=[bool(p["clipped"][j]) for p in trace]) for j in range(N)])
    # sanity: the walkthrough's pass-1 delta equals what the trace recorded
    assert np.allclose(trace[1]["W"], W_mid + CFG["lr"] * trace[0]["grad"])
    n_clip_last = passes[-1]["n_clipped"]
    first_clip_pass = next((p["k"] for p in passes if p["n_clipped"]), None)
    e4_title = (f"By pass 4, {words(n_clip_last)} of {N} steps have hit the 20% clip and stopped pushing"
                if n_clip_last else f"All {N} ratios stay inside the 20% clip band")
    labels, last = [], len(trace) - 1

    def label(k, j):
        if all(j != jj for _, (_, jj) in labels):
            labels.append((f"run {run_of[j] + 1}, step {step_of[j] + 1}: {trace[k]['rho'][j]:.2f}",
                           (k, j)))

    if first_clip_pass is not None:
        for j in np.flatnonzero(trace[first_clip_pass - 1]["clipped"]):
            label(first_clip_pass - 1, j)
        cut = np.flatnonzero(trace[last]["clipped"])
        label(last, cut[np.argmax(trace[last]["rho"][cut])])
        if trace[last]["clipped"][rep]:
            label(last, rep)
    chart_e4(trace, A, labels, e4_title)

    # E3b: the critic's blind spot, using the trained PPO bot and its own critic.
    W_ppo, v_ppo = Ws["ppo", 0]
    start_view = ll.sensors(np.zeros(ll.FLOOR, dtype=bool), 0)
    forecast = float(start_view @ v_ppo)
    frng = np.random.default_rng(BLIND_SEED)
    drawn = []
    for i in range(BLIND_DRAWS):
        lava = ll.new_floor(frng)
        prng = np.random.default_rng([BLIND_SEED, i])
        runs = [ll.play(W_ppo, lava, prng) for _ in range(8)]
        drawn.append((np.mean([x.reward for x in runs]), i, lava, runs))
    hard_m, hard_i, hard, hard_runs = min(drawn, key=lambda d: d[0])
    easy_m, easy_i, easy, easy_runs = max(drawn, key=lambda d: d[0])
    easy_r = np.array([x.reward for x in easy_runs])
    hard_r = np.array([x.reward for x in hard_runs])
    # every run's start-line view is identical, whatever the floor
    assert all((x.states[0] == start_view).all() for x in easy_runs + hard_runs)
    _, easy_adv = group_adv(easy_runs)
    _, hard_adv = group_adv(hard_runs)
    hard_best = int(np.argmax(hard_r))
    summary["e3b"] = dict(
        policy="PPO seed 0 after 300 updates", critic_v=np.round(v_ppo, 4).tolist(),
        forecast_from_start=rnd(forecast), draws=BLIND_DRAWS, floor_seed=BLIND_SEED,
        easy=dict(draw=easy_i, floor=floor_info(easy), rewards=[rnd(x) for x in easy_r],
                  mean=rnd(easy_r.mean()), worst=rnd(easy_r.min()), best=rnd(easy_r.max()),
                  exits=int((easy_r == 1).sum()),
                  critic_adv_from_start=[rnd(x - forecast) for x in easy_r]),
        hard=dict(draw=hard_i, floor=floor_info(hard), rewards=[rnd(x) for x in hard_r],
                  mean=rnd(hard_r.mean()), worst=rnd(hard_r.min()), best=rnd(hard_r.max()),
                  exits=int((hard_r == 1).sum()),
                  critic_adv_from_start=[rnd(x - forecast) for x in hard_r]),
        separated=bool(easy_r.min() > hard_r.max()),
        grpo_adv_easy=[rnd(a) for a in easy_adv], grpo_adv_hard=[rnd(a) for a in hard_adv],
        hard_best_run=dict(run_info(hard_runs[hard_best]), fall=fall_kind(hard_runs[hard_best], hard),
                           grpo_adv=rnd(hard_adv[hard_best]), critic_adv=rnd(hard_r[hard_best] - forecast)),
        ppo_avg_score=summary["final"]["ppo"]["per_seed"][0]["score"])
    e3b_title = ("The easy floor's worst run beats the hard floor's best"
                 if summary["e3b"]["separated"] else "One forecast, two very different floors")
    chart_e3b(easy, hard, easy_runs, hard_runs, forecast, e3b_title)

    # E5: what the trained GRPO policy learned (seed 0), plus the other seeds' tables.
    before_t, after_t = ll.leap_table(W0), ll.leap_table(W_grpo)
    e5_title = "Trained, the bot leaps when lava is directly ahead and steps everywhere else"
    chart_e5(before_t, after_t, e5_title)
    summary["e5"] = dict(
        W=np.round(W_grpo, 4).tolist(),
        leap_minus_step=np.round(W_grpo[:, 1] - W_grpo[:, 0], 4).tolist(),
        leap_table={format(b, "03b"): rnd(after_t[b]) for b in range(8)},
        leap_table_all_seeds={format(b, "03b"): [rnd(ll.leap_table(Ws["grpo", s][0])[b]) for s in SEEDS]
                              for b in range(8)},
        reinforce_leap_table={format(b, "03b"): rnd(ll.leap_table(Ws["reinforce", 0][0])[b])
                              for b in range(8)},
        reinforce_leap_table_all_seeds={format(b, "03b"): [rnd(ll.leap_table(Ws["reinforce", s][0])[b])
                                                           for s in SEEDS] for b in range(8)},
        ppo_leap_table={format(b, "03b"): rnd(ll.leap_table(Ws["ppo", 0][0])[b]) for b in range(8)},
        never_occurs=[format(b, "03b") for b in NEVER])

    # E6: exit-reach rates on the test floors.
    summary["e6"] = dict(
        untrained=rnd(untrained_exit),
        **{a: dict(mean=summary["final"][a]["exit_mean"], min=summary["final"][a]["exit_min"],
                   max=summary["final"][a]["exit_max"]) for a in ALGOS},
        ceiling=rnd(ceiling_exit))

    # Floor statistics for the test set (the reason the ceiling sits where it does).
    n_ch = [len(floor_info(f)["channels"]) for f in test]
    n_wide = [sum(w == 2 for _, w in floor_info(f)["channels"]) for f in test]
    summary["test_floor_stats"] = dict(mean_channels=rnd(np.mean(n_ch)), mean_wide=rnd(np.mean(n_wide)),
                                       share_with_no_wide=rnd(np.mean(np.array(n_wide) == 0)))

    (RESULTS / "summary.json").write_text(json.dumps(summary, indent=2))
    (RESULTS / "walkthrough.json").write_text(json.dumps(walkthrough, indent=2))
    print("wrote", RESULTS / "summary.json", "and", RESULTS / "walkthrough.json")


if __name__ == "__main__":
    main()
