"""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")))