SnapJudge2.0 / make_chart_v2.py
Seedyai's picture
SnapJudge 2.0: ModernBERT-large joint + 5 per-game experts, benchmarks + Jev-compatible server
c8d9552 verified
Raw History Blame Contribute Delete
3.37 kB
"""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)