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