Download code/harness/figures.py from dreddnafious/mini-agi-replication: direct link, hf CLI and curl.
- Browser
- Download file 13.5 kB
-
https://huggingface.co/spaces/dreddnafious/mini-agi-replication/resolve/main/code/harness/figures.py
- Command line
-
hf download hf://spaces/dreddnafious/mini-agi-replication/code/harness/figures.py
-
curl -L -o figures.py https://huggingface.co/spaces/dreddnafious/mini-agi-replication/resolve/main/code/harness/figures.py
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"))) | |