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