File size: 10,963 Bytes
0662d8e | 1 2 3 4 5 6 7 8 9 10 11 12 13 14 15 16 17 18 19 20 21 22 23 24 25 26 27 28 29 30 31 32 33 34 35 36 37 38 39 40 41 42 43 44 45 46 47 48 49 50 51 52 53 54 55 56 57 58 59 60 61 62 63 64 65 66 67 68 69 70 71 72 73 74 75 76 77 78 79 80 81 82 83 84 85 86 87 88 89 90 91 92 93 94 95 96 97 98 99 100 101 102 103 104 105 106 107 108 109 110 111 112 113 114 115 116 117 118 119 120 121 122 123 124 125 126 127 128 129 130 131 132 133 134 135 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 190 191 192 193 194 195 196 197 198 199 200 201 202 203 204 205 206 207 208 209 210 211 212 213 214 215 216 217 218 219 220 221 222 223 224 225 226 227 228 229 230 231 232 233 234 235 236 | """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()
|