dreddnafious's picture
Revision 2: step density matched; headline replicates
0662d8e verified
Raw History Blame Contribute Delete
11 kB
"""Forgetting probe: read one subject (or several) from a trained base and
measure every subject's held-out loss as reading goes on.
Rebuilds upstream's unpublished probe from the README description:
"reads 524,000 characters of chess and nothing else, at batch 1", with arms
swap + trunk LR = expert LR
frozen + trunk LR = expert LR (working set never re-chosen)
swap + trunk at 0.1x (what the run uses)
control: all subjects read
The training step is upstream's `cmd_read` inner loop, line for line in
behaviour: peek_experts on the first chunk of a visit and want_experts after,
lr = base_lr * plasticity factor, FileReader.step (whole-window re-forward),
backward, clip, AdamW. Optimiser moments are restored from the base.
Deliberate choices where upstream is silent (recorded in the output JSON):
* the plasticity controller is FROZEN at the base's scale - it does not
observe the probe's evaluations, so every arm reads at the same rate;
* no growth and no pruning during the probe (a dry `read` does neither);
* "frozen" = the working set is chosen once, by peeking the first chess
chunk, and restored after every evaluation (evaluation itself re-chooses
per subject, exactly as upstream's evaluator does);
* every arm evaluates the same held-out text at the same character counts.
"""
import argparse
import json
import math
import os
import shutil
import sys
import time
import numpy as np
import torch
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from upstream import use_config # noqa: E402
CHANCE = math.log(265)
def setup(args):
c = use_config(args.config)
import train as T
from minagi import store as weights_store
from minagi.plasticity import Plasticity
from minagi.precision import set_compute_dtype
os.makedirs(args.out, exist_ok=True)
wdir = os.path.join(args.out, "weights")
if os.path.exists(wdir):
shutil.rmtree(wdir)
shutil.copytree(args.base, wdir)
device = torch.device("cuda")
torch.manual_seed(args.seed)
set_compute_dtype(c["training"]["precision"])
model, cfg, pool, man = T.build_paged(wdir, device, c["pool"]["resident"],
c["pool"]["ram_cache"],
c["model"]["context_end"])
m = c["model"]
chunk = int(c["training"]["chunk"])
# as cmd_read: config wins over the manifest for the depth policy
cfg.train_steps_mean = float(m["train_steps_mean"])
cfg.min_steps = max(1, min(int(m["min_steps"]), cfg.max_steps))
cfg.bptt_window = min(int(m["bptt_window"]), cfg.max_steps)
cfg.ponder_beta = float(m["ponder_beta"])
cfg.halt_prior = float(m["halt_prior"])
cfg.halt_thresh = float(m["halt_thresh"])
pool.margin = float(c["pool"]["margin"])
pool.dwell = max(1, int(c["pool"]["dwell_chars"]) // chunk)
pool.explore = float(c["pool"]["explore"])
pool.dying_at = float(c["prune"]["dying_at"])
pool.trial = max(1, int(round(int(c["prune"]["survival_chars"]) / chunk)))
step0 = int(man.get("step", 0) or 0)
pool.now = step0
lr = float(c["training"]["lr"])
trunk, pool_ps = T._split_trunk_pool(model)
tg = {"params": trunk, "name": "trunk", "weight_decay": float(c["training"]["weight_decay"]),
"lr": lr * args.trunk_mult, "base_lr": lr * args.trunk_mult}
pg = {"params": pool_ps, "name": "pool", "weight_decay": float(c["training"]["weight_decay"]),
"lr": lr * args.pool_mult, "base_lr": lr * args.pool_mult}
opt = torch.optim.AdamW([tg, pg], lr=lr, betas=(0.9, 0.95), fused=True)
pool.attach_optimiser(opt)
weights_store._load_optim(opt, model, wdir)
plast = Plasticity.restore(man.get("plasticity"))
scale = plast.factor() if args.lr_scale is None else float(args.lr_scale)
for g in opt.param_groups:
g["lr"] = g["base_lr"] * scale
return c, T, model, cfg, pool, man, opt, chunk, scale, trunk
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--config", required=True)
ap.add_argument("--base", required=True, help="trained weights directory (copied, never written)")
ap.add_argument("--out", required=True)
ap.add_argument("--train-root", default="upstream/mini-AGI/data/train")
ap.add_argument("--held-out", default="upstream/mini-AGI/data/val")
ap.add_argument("--lanes", default="chess", help="comma list of subjects to read, or 'all'")
ap.add_argument("--arm", choices=["swap", "frozen"], default="swap")
ap.add_argument("--trunk-mult", type=float, default=0.1)
ap.add_argument("--pool-mult", type=float, default=1.0)
ap.add_argument("--lr-scale", type=float, default=None,
help="override the plasticity scale; default = the base's")
ap.add_argument("--chars", type=int, default=524_288)
ap.add_argument("--eval-every", type=int, default=65_536)
ap.add_argument("--eval-chunks", type=int, default=120, help="chunks per subject per evaluation")
ap.add_argument("--seed", type=int, default=0)
ap.add_argument("--accum", type=int, default=1,
help="chunks per optimiser step (gradient accumulation); 4 at chunk 512 matches "
"upstream's step density at chunk 2048")
ap.add_argument("--recover-chars", type=int, default=0,
help="after the probe, read ALL subjects this long (displacement test)")
ap.add_argument("--recover-eval-every", type=int, default=32_768)
args = ap.parse_args()
c, T, model, cfg, pool, man, opt, chunk, scale, trunk = setup(args)
from minagi.ingest import collect
from minagi.stream import FolderEvaluator
device = torch.device("cuda")
ctx = int(cfg.block)
ev = FolderEvaluator(model, args.held_out, chunk, ctx, device,
segment_chunks=max(1, int(c["pool"]["segment_chars"]) // chunk))
subjects = sorted(d for d in os.listdir(args.train_root)
if os.path.isdir(os.path.join(args.train_root, d)))
read = subjects if args.lanes == "all" else args.lanes.split(",")
base_chars = int(man.get("read_chars", 0) or 0)
def lanes_for(names, salt):
files = collect([os.path.join(args.train_root, s) for s in names])
return T._lanes(files, int(c["data"]["shuffle_seed"]) + 1000 * (args.seed + 1),
[args.train_root], resume=base_chars + salt)
record = {"args": vars(args), "lr_scale": scale, "base_step": int(man.get("step", 0) or 0),
"base_chars": base_chars, "n_experts": pool.n_experts(),
"resident": pool.resident, "chance": CHANCE, "evals": [], "recover": []}
def evaluate(at, phase, frozen_set=None):
t = time.time()
d = ev.run(args.eval_chunks)
se = d.pop("stderr", None)
if frozen_set is not None:
pool.swap_to(frozen_set)
e = {"chars": at, "phase": phase, "loss": d, "stderr": se, "secs": round(time.time() - t, 1)}
record[phase if phase == "recover" else "evals"].append(e)
shown = " ".join(f"{k[:6]} {v:.3f}" for k, v in d.items())
print(f"[{phase} {at/1e3:7.0f}k] {shown}", flush=True)
return e
trained_experts = set()
def read_stream(names, total, every, phase, frozen=False, salt=0):
lanes = lanes_for(names, salt)
turn = max(max(1, -(-ctx // chunk)), max(1, -(-int(c["data"]["passage"]) // chunk)))
seen, step = 0, record.get("_step", int(man.get("step", 0) or 0))
frozen_set = None
next_eval = every
while seen < total:
for lane in lanes:
r = lane.open(model, chunk, ctx, device, turn * chunk)
if r is None:
continue
pending = 0 # chunks accumulated since the last optimiser step
def apply_step():
nonlocal pending, step
torch.nn.utils.clip_grad_norm_(model.parameters(), float(c["training"]["clip"]))
opt.step()
opt.zero_grad(set_to_none=True)
pending = 0
step += 1
pool.now = step
for j in range(turn):
if r.done() or seen >= total:
break
nxt = r.peek()
# Experts are (re)chosen only at the start of an accumulation cycle: swapping
# mid-cycle would apply one expert's accumulated slot gradient to another.
if nxt is not None and pending == 0:
if frozen:
if frozen_set is None:
model.peek_experts(nxt, free=True)
frozen_set = [s for s in pool.slots if s >= 0]
record["frozen_set"] = [int(pool._f(i)) for i in frozen_set]
elif j == 0:
model.peek_experts(nxt, free=True)
else:
model.want_experts(nxt)
if pending == 0:
opt.zero_grad(set_to_none=True)
use0 = pool.use.clone()
loss = r.step(learn=True, aux_weight=cfg.pool_aux)
if loss is None:
break
(loss / args.accum).backward()
pending += 1
if pending == args.accum:
apply_step()
hit = (pool.use - use0 > 0).nonzero().flatten().tolist()
trained_experts.update(int(pool._f(i)) for i in hit)
seen += chunk
if pending == 0 and (seen >= next_eval or seen >= total):
evaluate(seen, phase, frozen_set)
next_eval += every
if pending:
apply_step() # a visit that ended mid-cycle still takes its step
lane.rest(model)
if seen >= total:
break
record["_step"] = step
t0 = time.time()
evaluate(0, "evals")
read_stream(read, args.chars, args.eval_every, "evals", frozen=(args.arm == "frozen"))
record["trained_experts"] = len(trained_experts)
record["opt_steps"] = record["_step"] - record["base_step"]
if args.recover_chars:
read_stream(subjects, args.recover_chars, args.recover_eval_every, "recover", salt=7_777_777)
record["minutes"] = round((time.time() - t0) / 60, 2)
record.pop("_step", None)
with open(os.path.join(args.out, "probe.json"), "w") as f:
json.dump(record, f, indent=1)
shutil.rmtree(os.path.join(args.out, "weights"))
print(f"done in {record['minutes']} min; {len(trained_experts)} of {pool.n_experts()} experts trained")
if __name__ == "__main__":
main()