model / summary /scripts /plot_presentation_ablations.py
jyhptr's picture
Add project summary handoff
756cf82 verified
Raw
History Blame Contribute Delete
5.96 kB
#!/usr/bin/env python3
from __future__ import annotations
import math
import os
import re
from pathlib import Path
os.environ.setdefault("MPLCONFIGDIR", "/tmp/matplotlib")
import matplotlib.pyplot as plt
ROOT = Path(__file__).resolve().parents[1]
LOG_DIR = ROOT / "logs"
FIG_DIR = ROOT / "figures"
EXPERIMENTS = [
{
"key": "a0_nanogpt_fineweb_adamw",
"label": "Naive\nbaseline",
"short": "baseline",
"train_log": LOG_DIR / "pres_ablation_a0_nanogpt_fineweb_adamw_train.log",
"eval_log": LOG_DIR / "pres_ablation_a0_nanogpt_fineweb_adamw_course_val_eval.log",
},
{
"key": "a1_nanogpt_mixed_adamw",
"label": "Mixed\ndata",
"short": "data",
"train_log": LOG_DIR / "pres_ablation_a1_nanogpt_mixed_adamw_train.log",
"eval_log": LOG_DIR / "pres_ablation_a1_nanogpt_mixed_adamw_course_val_eval.log",
},
{
"key": "a2_nanogpt_fineweb_muon",
"label": "Muon\noptimizer",
"short": "Muon",
"train_log": LOG_DIR / "pres_ablation_a2_nanogpt_fineweb_muon_train.log",
"eval_log": LOG_DIR / "pres_ablation_a2_nanogpt_fineweb_muon_course_val_eval.log",
},
{
"key": "a3_lyra_fineweb_adamw",
"label": "Lyra\narchitecture",
"short": "arch",
"train_log": LOG_DIR / "pres_ablation_a3_lyra_fineweb_adamw_train.log",
"eval_log": LOG_DIR / "pres_ablation_a3_lyra_fineweb_adamw_course_val_eval.log",
},
{
"key": "a4_lyra_mixed_muon",
"label": "Combined\nshort run",
"short": "combined",
"train_log": LOG_DIR / "pres_ablation_a4_lyra_mixed_muon_train.log",
"eval_log": LOG_DIR / "pres_ablation_a4_lyra_mixed_muon_course_val_eval.log",
},
]
FINAL_LONG_RUN = {
"label": "Final\n37k",
"short": "long run",
"ppl": 19.7822,
}
STEP_RE = re.compile(r"step\s+(\d+): train loss ([0-9.]+), val loss ([0-9.]+)")
PPL_RE = re.compile(r"Perplexity:\s+([0-9.]+)")
def parse_last_training_run(path: Path) -> list[tuple[int, float, float]]:
if not path.exists():
return []
runs: list[list[tuple[int, float, float]]] = []
current: list[tuple[int, float, float]] = []
for line in path.read_text(errors="replace").splitlines():
match = STEP_RE.search(line)
if not match:
continue
step = int(match.group(1))
row = (step, float(match.group(2)), float(match.group(3)))
if step == 0 and current:
runs.append(current)
current = []
current.append(row)
if current:
runs.append(current)
return runs[-1] if runs else []
def parse_ppl(path: Path) -> float | None:
if not path.exists():
return None
matches = PPL_RE.findall(path.read_text(errors="replace"))
return float(matches[-1]) if matches else None
def main() -> None:
FIG_DIR.mkdir(exist_ok=True)
train_runs = {exp["key"]: parse_last_training_run(exp["train_log"]) for exp in EXPERIMENTS}
ppls = [parse_ppl(exp["eval_log"]) for exp in EXPERIMENTS]
plt.style.use("seaborn-v0_8-whitegrid")
fig, axes = plt.subplots(1, 2, figsize=(13.5, 5.2), constrained_layout=True)
ax = axes[0]
labels = [exp["label"] for exp in EXPERIMENTS]
vals = [p if p is not None else math.nan for p in ppls]
colors = ["#6b7280", "#2563eb", "#dc2626", "#059669", "#7c3aed"]
bars = ax.bar(labels, vals, color=colors, width=0.68)
ax.axhline(FINAL_LONG_RUN["ppl"], color="#111827", linewidth=1.8, linestyle="--")
ax.text(
0.02,
FINAL_LONG_RUN["ppl"] + 0.3,
f"final long run: {FINAL_LONG_RUN['ppl']:.2f}",
transform=ax.get_yaxis_transform(),
ha="left",
va="bottom",
fontsize=9,
color="#111827",
)
ax.set_title("Course Public Validation Perplexity", fontsize=13, weight="bold")
ax.set_ylabel("perplexity, lower is better")
ax.tick_params(axis="x", labelrotation=0)
for bar, val in zip(bars, vals):
if math.isnan(val):
ax.text(
bar.get_x() + bar.get_width() / 2,
1,
"pending",
ha="center",
va="bottom",
fontsize=9,
rotation=90,
color="#374151",
)
else:
ax.text(
bar.get_x() + bar.get_width() / 2,
bar.get_height(),
f"{val:.1f}",
ha="center",
va="bottom",
fontsize=9,
)
finite_vals = [v for v in vals if not math.isnan(v)]
if finite_vals:
low = min([FINAL_LONG_RUN["ppl"], *finite_vals])
high = max(finite_vals)
ax.set_ylim(max(0, low * 0.85), high * 1.12)
else:
ax.set_ylim(0, 100)
ax = axes[1]
plotted_curves = 0
for exp, color in zip(EXPERIMENTS, colors):
run = train_runs[exp["key"]]
if not run:
continue
plotted_curves += 1
steps = [row[0] for row in run]
val_losses = [row[2] for row in run]
ax.plot(steps, val_losses, marker="o", linewidth=2.0, markersize=4, color=color, label=exp["short"])
ax.set_title("Short-Run Heldout Loss Curves", fontsize=13, weight="bold")
ax.set_xlabel("training iteration")
ax.set_ylabel("validation loss")
if plotted_curves:
ax.legend(frameon=True, fontsize=9)
ax.text(
0.02,
-0.18,
"Each run changes one variable and uses the same 4,500-iteration budget; public PPL is the comparable metric.",
transform=ax.transAxes,
fontsize=8.5,
color="#4b5563",
)
fig.suptitle("Ablation Summary for the Presentation", fontsize=15, weight="bold")
for suffix in ("png", "pdf"):
out = FIG_DIR / f"presentation_ablation_summary.{suffix}"
fig.savefig(out, dpi=220)
print(out)
if __name__ == "__main__":
main()