dreddnafious's picture
Revision 2: step density matched; headline replicates
0662d8e verified
Raw History Blame Contribute Delete
13.5 kB
"""Figures for REPORT.md, from runs/probes/s_seed{0,1}/*/probe.json."""
import glob
import json
import math
import os
import matplotlib
matplotlib.use("Agg")
import matplotlib.pyplot as plt # noqa: E402
OUT = "report_figs"
SURF, INK, INK2, GRID = "#fcfcfb", "#0b0b0b", "#52514e", "#e4e3df"
BLUE, ORANGE, AQUA, GRAY = "#2a78d6", "#eb6834", "#1baf7a", "#8a8984"
SEED_STYLE = {0: "-", 1: "--"}
UPSTREAM = {"R1": 2.5871, "R2": 2.2300, "R3": 0.0067, "R4": -0.0077}
plt.rcParams.update({
"figure.facecolor": SURF, "axes.facecolor": SURF, "savefig.facecolor": SURF,
"axes.edgecolor": GRID, "axes.labelcolor": INK2, "xtick.color": INK2,
"ytick.color": INK2, "text.color": INK, "font.size": 11,
"axes.spines.top": False, "axes.spines.right": False, "axes.grid": True,
"grid.color": GRID, "grid.linewidth": 0.8, "lines.linewidth": 2,
"axes.titleweight": "bold", "axes.titlesize": 13, "axes.titlelocation": "left",
})
def load(seed, name):
p = f"runs/probes/s_seed{seed}/{name}/probe.json"
return json.load(open(p)) if os.path.exists(p) else None
def curve(p, phase="evals", unread_only=True):
ev = p[phase]
a = p["evals"][0]["loss"]
lanes = p["args"]["lanes"]
keys = [k for k in a if lanes == "all" or k not in lanes.split(",")] if unread_only else list(a)
xs = [e["chars"] / 1e3 for e in ev]
ys = [sum(e["loss"][k] - a[k] for k in keys) / len(keys) for e in ev]
return xs, ys, keys
def label_end(ax, x, y, text, color):
ax.annotate(text, (x, y), xytext=(6, 0), textcoords="offset points",
va="center", fontsize=10, color=INK2)
ax.plot([x], [y], "o", ms=5, color=color, mec=SURF, mew=2)
def fig_probe():
fig, ax = plt.subplots(figsize=(8, 4.6))
for s in (0, 1):
p = load(s, "R3_swap_t0.1")
x, y, _ = curve(p)
ax.plot(x, y, SEED_STYLE[s], color=BLUE)
label_end(ax, x[-1], y[-1], f"trunk 0.1x, seed {s}: +{y[-1]:.2f}", BLUE)
c = load(s, "R4_control_all")
x, y, _ = curve(c)
ax.plot(x, y, SEED_STYLE[s], color=GRAY, lw=1.5)
ax.annotate("control, all subjects read\n(both seeds)", (530, 0.075), fontsize=10,
color=INK2, va="bottom")
ax.plot([524.3], [UPSTREAM["R3"]], "D", ms=8, color=ORANGE, mec=SURF, mew=2)
ax.annotate("upstream README: +0.0067", (524.3, UPSTREAM["R3"]), xytext=(10, -12),
textcoords="offset points", ha="left", fontsize=10, color=INK2)
ax.set_xlim(0, 700)
ax.set_xlabel("characters of chess read (thousands)")
ax.set_ylabel("forgetting on the 7 unread subjects (nats/char)")
ax.set_title("At the inherited probe rate: 50-110x the README number")
fig.tight_layout()
fig.savefig(f"{OUT}/1_probe.png", dpi=150)
def fig_sweep():
arms = [("E4_swap_t0", 0), ("E3_swap_t0.03", 0.03), ("R3_swap_t0.1", 0.1),
("E2_swap_t0.3", 0.3), ("R2_swap_t1", 1.0)]
pos = {m: i for i, (_, m) in enumerate(arms)}
fig, ax = plt.subplots(figsize=(8, 4.6))
for s in (0, 1):
xs, ys = [], []
for name, m in arms:
p = load(s, name)
if p:
xs.append(pos[m]); ys.append(curve(p)[1][-1])
ax.plot(xs, ys, SEED_STYLE[s], marker="o", ms=7, color=BLUE, mec=SURF, mew=2,
label=f"trunk multiplier sweep, seed {s}")
g = load(s, "E1_swap_all0.1")
if g:
ax.plot([pos[0.1] + 0.12], [curve(g)[1][-1]], "s", ms=8, color=AQUA, mec=SURF, mew=2,
label="everything at 0.1x (no split)" if s == 0 else None)
ax.plot([pos[0.1], pos[1.0]], [UPSTREAM["R3"], UPSTREAM["R2"]], "D", ms=8, color=ORANGE,
mec=SURF, mew=2, ls="none", label="upstream README")
ax.set_yscale("log")
ax.set_xticks(range(len(arms)), [f"{m:g}x" for _, m in arms])
ax.set_xlabel("trunk learning rate, as a multiple of the experts'")
ax.set_ylabel("forgetting after 524k chars (nats/char, log)")
ax.set_title("Trunk LR is the lever; the trunk/expert split is not")
ax.legend(frameon=False, fontsize=10, loc="upper left")
fig.tight_layout()
fig.savefig(f"{OUT}/2_sweep.png", dpi=150)
def fig_long():
fig, ax = plt.subplots(figsize=(8, 4.6))
for s in (0, 1):
p = load(s, "E6_swap_t0.1_4M")
if p:
x, y, keys = curve(p)
ax.plot([v / 1e3 for v in x], y, SEED_STYLE[s], color=BLUE, marker="o", ms=5)
label_end(ax, x[-1] / 1e3, y[-1], f"trunk 0.1x, inherited LR, seed {s}: +{y[-1]:.1f}", BLUE)
a = p["evals"][0]["loss"]
chance = math.log(265) - sum(a[k] for k in keys) / len(keys)
ax.axhline(chance, color=INK2, lw=1, ls=SEED_STYLE[s])
c = load(s, "E7_control_4M")
if c:
x, y, _ = curve(c)
ax.plot([v / 1e3 for v in x], y, SEED_STYLE[s], color=GRAY, lw=1.5)
q = load(s, "L_swap_t0.1_scale0.1_4M")
if q:
x, y, _ = curve(q)
ax.plot([v / 1e3 for v in x], y, SEED_STYLE[s], color=AQUA, marker="s", ms=5)
label_end(ax, x[-1] / 1e3, y[-1], f"same, probe LR 0.1x, seed {s}: +{y[-1]:.2f}", AQUA)
ax.annotate("above this line: worse than uniform guessing",
(0.05, 4.42), fontsize=10, color=INK2, va="bottom")
ax.annotate("control, all subjects read", (2.6, -0.32), fontsize=10, color=INK2)
ax.set_xlim(0, 6.6)
ax.set_ylim(-0.45, 5.4)
ax.set_xlabel("characters of chess read (millions)")
ax.set_ylabel("forgetting on the 7 unread subjects (nats/char)")
ax.set_title("Read 8x longer, forgetting keeps growing, even at the low rate")
fig.tight_layout()
fig.savefig(f"{OUT}/3_long.png", dpi=150)
def fig_recover():
fig, ax = plt.subplots(figsize=(8, 4.6))
for s in (0, 1):
p = load(s, "R2_swap_t1")
if not p or not p.get("recover"):
continue
a, b = p["evals"][0]["loss"], p["evals"][-1]["loss"]
keys = [k for k in a if k != "chess"]
dmg = sum(b[k] - a[k] for k in keys) / len(keys)
xs = [0] + [e["chars"] / 1e3 for e in p["recover"]]
ys = [0] + [100 * (dmg - sum(e["loss"][k] - a[k] for k in keys) / len(keys)) / dmg
for e in p["recover"]]
ax.plot(xs, ys, SEED_STYLE[s], color=BLUE, label=f"seed {s} (damage +{dmg:.2f} nats)")
ax.plot([131], [75], "D", ms=8, color=ORANGE, mec=SURF, mew=2, label="upstream README: ~75% by 131k")
ax.axvline(262.144, color=GRID, lw=1)
ax.annotate("every subject visited once", (268, 8), fontsize=10, color=INK2)
ax.set_ylim(0, 100)
ax.set_xlabel("characters of mixed reading after the damage (thousands)")
ax.set_ylabel("damage recovered (%)")
ax.set_title("Recovery replicates, then stalls short of full")
ax.legend(frameon=False, fontsize=10, loc="lower right")
fig.tight_layout()
fig.savefig(f"{OUT}/4_recover.png", dpi=150)
def fig_lrscale():
scales = [0.1, 0.3, 0.6, 1.0]
fig, ax = plt.subplots(figsize=(8, 4.6))
any_ = False
for s in (0, 1):
xs, ys = [], []
for sc in scales:
p = load(s, f"L_swap_t0.1_scale{sc}")
if p:
xs.append(sc); ys.append(curve(p)[1][-1])
if xs:
any_ = True
ax.plot(xs, ys, SEED_STYLE[s], marker="o", ms=7, color=BLUE, mec=SURF, mew=2,
label=f"seed {s}")
if not any_:
plt.close(fig)
return
ax.axhline(UPSTREAM["R3"], color=ORANGE, lw=1.5, ls=":")
ax.annotate("upstream README: +0.0067", (0.62, UPSTREAM["R3"] * 1.25), fontsize=10, color=INK2)
ctl = [curve(load(s, "R4_control_all"))[1][-1] for s in (0, 1)]
ax.axhline(sum(ctl) / 2, color=GRAY, lw=1.5, ls=":")
ax.annotate("control, all subjects read: +0.017", (0.62, sum(ctl) / 2 * 1.2), fontsize=10, color=INK2)
ax.set_yscale("log")
ax.set_xlabel("probe learning rate, as a fraction of the configured 3e-4")
ax.set_ylabel("forgetting after 524k chars (nats/char, log)")
ax.set_title("Upstream's number reappears at a 10x lower probe learning rate")
ax.legend(frameon=False, fontsize=10, loc="upper left")
fig.tight_layout()
fig.savefig(f"{OUT}/5_lrscale.png", dpi=150)
# ---------------------------------------------------------------- v2: at upstream's step density
# v1 probes ran at 4x upstream's optimiser steps per character (chunk 512 vs 2048). v2 figures use the
# --accum 4 reruns (runs/probes/*/A4_*, S_*_accum4), which match upstream's step density.
def fig_v2_headline():
fig, ax = plt.subplots(figsize=(8, 4.6))
for s in (0, 1):
x, y, _ = curve(load(s, "S_swap_t0.1_accum4"))
ax.plot(x, y, SEED_STYLE[s], color=BLUE)
ax.plot([x[-1]], [y[-1]], "o", ms=5, color=BLUE, mec=SURF, mew=2)
ax.annotate(f"trunk 0.1x, seed {s}: +{y[-1]:.3f}", (x[-1], y[-1]), xytext=(8, 9 if s == 0 else -2),
textcoords="offset points", va="center", fontsize=10, color=INK2)
x, y, _ = curve(load(s, "A4_control_all"))
ax.plot(x, y, SEED_STYLE[s], color=GRAY, lw=1.5)
ax.annotate("control, all subjects read (both seeds)", (300, -0.075), fontsize=10, color=INK2)
ax.plot([524.3], [UPSTREAM["R3"]], "D", ms=8, color=ORANGE, mec=SURF, mew=2)
ax.annotate("upstream README: +0.0067", (524.3, UPSTREAM["R3"]), xytext=(10, -14),
textcoords="offset points", ha="left", fontsize=10, color=INK2)
ax.set_xlim(0, 720)
ax.set_ylim(-0.1, 0.12)
ax.set_xlabel("characters of chess read (thousands)")
ax.set_ylabel("forgetting on the 7 unread subjects (nats/char)")
ax.set_title("At upstream's step density, the headline replicates")
fig.tight_layout()
fig.savefig(f"{OUT}/1_headline.png", dpi=150)
def fig_v2_displacement():
"""Every trunk-0.1x, 524k-char probe vs lr_scale x optimiser steps (upstream-equivalent LR)."""
acc1 = ["L_swap_t0.1_scale0.1", "L_swap_t0.1_scale0.3", "L_swap_t0.1_scale0.6",
"L_swap_t0.1_scale1.0", "R3_swap_t0.1"]
accn = ["S_swap_t0.1_accum2", "S_swap_t0.1_accum4"]
fig, ax = plt.subplots(figsize=(8, 4.8))
def pts(s, names):
out = []
for n in names:
p = load(s, n)
if p:
out.append((p["lr_scale"] * (p.get("opt_steps") or 1024) / 256, curve(p)[1][-1]))
return sorted(out)
for s in (0, 1):
a = pts(s, acc1)
ax.plot([x for x, _ in a], [y for _, y in a], SEED_STYLE[s], marker="o", ms=7, color=BLUE,
mec=SURF, mew=2, label=f"learning rate varied, seed {s}")
b = pts(s, accn)
ax.plot([x for x, _ in b], [y for _, y in b], "s", ms=9, color=AQUA, mec=SURF, mew=2,
label="step count varied (gradient accumulation)" if s == 0 else None)
xs = [0.35, 1.3]
ax.plot(xs, [0.0065 * (x / 0.35) ** 2 for x in xs], ":", color=INK2, lw=1.5)
ax.annotate("slope 2: forgetting ∝ displacement²", (0.62, 0.012), fontsize=10, color=INK2, rotation=27)
ax.axvline(0.7, color=GRID, lw=1)
ax.annotate("upstream's setting (≈0.6-0.8)", (0.72, 0.0068), fontsize=9.5, color=INK2)
ax.set_xscale("log"); ax.set_yscale("log")
from matplotlib.ticker import FixedLocator, FixedFormatter, NullLocator
ticks = [0.4, 0.6, 1, 2, 3, 4]
ax.xaxis.set_major_locator(FixedLocator(ticks)); ax.xaxis.set_minor_locator(NullLocator())
ax.xaxis.set_major_formatter(FixedFormatter([f"{t:g}" for t in ticks]))
yt = [0.01, 0.03, 0.1, 0.3, 1]
ax.yaxis.set_major_locator(FixedLocator(yt)); ax.yaxis.set_minor_locator(NullLocator())
ax.yaxis.set_major_formatter(FixedFormatter([f"{t:g}" for t in yt]))
ax.set_xlabel("probe learning rate × optimiser steps (1.0 = configured rate at upstream's step density)")
ax.set_ylabel("forgetting after 524k chars (nats/char)")
ax.set_title("Forgetting tracks how far the trunk moves, roughly squared")
ax.legend(frameon=False, fontsize=10, loc="upper left")
fig.tight_layout()
fig.savefig(f"{OUT}/2_displacement.png", dpi=150)
def fig_v2_long():
fig, ax = plt.subplots(figsize=(8, 4.6))
for s in (0, 1):
p = load(s, "A4_swap_t0.1_4M")
x, y, keys = curve(p)
ax.plot([v / 1e3 for v in x], y, SEED_STYLE[s], color=BLUE, marker="o", ms=5)
label_end(ax, x[-1] / 1e3, y[-1], f"trunk 0.1x, seed {s}: +{y[-1]:.2f}", BLUE)
c = load(s, "E7_control_4M")
x, y, _ = curve(c)
ax.plot([v / 1e3 for v in x], y, SEED_STYLE[s], color=GRAY, lw=1.5)
ax.annotate("control, all subjects read", (2.4, -0.17), fontsize=10, color=INK2)
ax.axvspan(0, 0.53, color=GRID, alpha=0.5, lw=0)
ax.annotate("the README's\nprobe length", (0.04, 1.25), fontsize=9.5, color=INK2)
ax.set_xlim(0, 5.6)
ax.set_ylim(-0.25, 1.75)
ax.set_xlabel("characters of chess read (millions)")
ax.set_ylabel("forgetting on the 7 unread subjects (nats/char)")
ax.set_title("Read 8x longer: the floor gives way after about 0.5M characters")
fig.tight_layout()
fig.savefig(f"{OUT}/3_long.png", dpi=150)
if __name__ == "__main__":
import sys
if len(sys.argv) > 1 and sys.argv[1] == "v2":
OUT = "report_figs_v2"
os.makedirs(OUT, exist_ok=True)
for f in (fig_v2_headline, fig_v2_displacement, fig_v2_long, fig_recover):
f()
else:
os.makedirs(OUT, exist_ok=True)
for f in (fig_probe, fig_sweep, fig_long, fig_recover, fig_lrscale):
f()
print(sorted(glob.glob(f"{OUT}/*.png")))