"""SnapJudge 2.0 performance charts (all numbers measured, seed 999).""" import matplotlib matplotlib.use("Agg") import matplotlib.pyplot as plt import numpy as np plt.rcParams.update({"font.size": 10, "axes.spines.top": False, "axes.spines.right": False}) games = ["tic-tac-toe", "snake", "temple-run", "connect-4", "maze"] joint = [0.7167, 0.9917, 1.0, 0.6917, 0.8433] expert = [0.8000, 0.9967, 1.0, 0.7450, 0.8450] fig, ax = plt.subplots(2, 2, figsize=(13, 8.5)) fig.suptitle("SnapJudge 2.0 — System-1 Game Decisions (421M active / 2.5B total, ModernBERT-large)", fontsize=14, fontweight="bold") fig.text(0.5, 0.93, "Fresh held-out: 1,000 games / 3,000 decisions, seed 999 • Tesla T4, fp16 • tie-aware (any optimal move counts)", ha="center", fontsize=9, color="#555") a = ax[0, 0] x = np.arange(len(games)); w = 0.36 b1 = a.bar(x - w/2, joint, w, label="Joint", color="#adb5bd") b2 = a.bar(x + w/2, expert, w, label="Expert", color="#264653") a.set_xticks(x); a.set_xticklabels(games, rotation=12); a.set_ylim(0.6, 1.05); a.set_ylabel("Accuracy") a.set_title("Joint vs per-game expert (fresh seed 999)") a.legend(frameon=False, fontsize=9) for i in range(len(games)): a.text(i - w/2, joint[i] + 0.008, f"{joint[i]:.3f}", ha="center", fontsize=7, color="#555") a.text(i + w/2, expert[i] + 0.008, f"{expert[i]:.3f}", ha="center", fontsize=7, fontweight="bold") a.grid(axis="y", alpha=0.3) b = ax[0, 1] stages = ["CE-1", "CE-2", "CE-3", "RLCD-1", "RLCD-2"] curve = [0.664, 0.680, 0.723, 0.780, 0.813] b.plot(stages, curve, "-o", color="#2a9d8f", lw=2, markersize=7) b.fill_between(stages, curve, alpha=0.12, color="#2a9d8f") b.set_ylim(0.6, 0.9); b.set_ylabel("Val accuracy (5 games)") b.set_title("Joint training curve (7,500 games, 2xT4 DDP)") for i, v in enumerate(curve): b.text(i, v + 0.008, f"{v:.3f}", ha="center", fontsize=8, fontweight="bold") b.grid(alpha=0.3) c = ax[1, 0] labels = ["SnapJudge 2.0\n(self-hosted)", "Jev API\n(published p50)"] lat = [53.1, 256.0] cols = ["#264653", "#e76f51"] bars = c.bar(labels, lat, color=cols, width=0.55) c.set_ylabel("ms per 3-question call") c.set_title("Latency: ~5x faster than API round-trip") for bar, v in zip(bars, lat): c.text(bar.get_x() + bar.get_width()/2, v + 5, f"{v:.0f} ms", ha="center", fontsize=10, fontweight="bold") c.set_ylim(0, 300); c.grid(axis="y", alpha=0.3) d = ax[1, 1] d.axis("off") rows = [ ["Routed overall", "0.877", "joint-only 0.849"], ["Hard-case slice", "0.719", "906 decisions"], ["ECE (scaled)", "0.190", "raw 0.392"], ["Latency p50", "53 ms", "batched 57 ms/state"], ["Params", "421M active", "2.5B total (6 ckpts)"], ["20-option stress", "pass", "108 ms"], ] tbl = d.table(cellText=rows, colLabels=["Metric", "Measured", "Note"], loc="center", colWidths=[0.34, 0.28, 0.38]) tbl.auto_set_font_size(False); tbl.set_fontsize(9); tbl.scale(1, 1.45) for (r, col), cell in tbl.get_celld().items(): if r == 0: cell.set_facecolor("#264653"); cell.set_text_props(color="white", fontweight="bold") elif r % 2 == 0: cell.set_facecolor("#f1faee") d.set_title("Summary — one forward pass, nothing to parse", pad=12) plt.tight_layout(rect=[0, 0, 1, 0.9]) out = "/kaggle/working/SnapJudge2.0/snapjudge2_perf.png" plt.savefig(out, dpi=150, bbox_inches="tight") print("saved", out)