from __future__ import annotations import argparse import json import queue import random import sys import threading import time from pathlib import Path import torch import torch.nn.functional as F sys.path.insert(0, str(Path(__file__).resolve().parent)) from model import ( make_reduce_cell, make_add_cell, reduce_features, add_features, ) from data import make_reduce_batch, make_add_batch STAGES = [8, 12, 20, 36, 68, 132, 260, 516, 1027] TOKEN_BUDGET = 40_000 EVAL_EVERY = 500 ADVANCE_THRESHOLD = 0.9995 SOFT_STEPS = 1_500 NOISE = 0.01 MARGIN = 3.0 MARGIN_W = 0.1 AUX_W = 0.3 PLATEAU_STEPS = 12_000 KINDS = ("reduce", "add") def features_for(kind, b): if kind == "reduce": return reduce_features(b["x"], b["p"], b["p3"]) return add_features(b["x"], b["y"], b["g"]) def forward_loss(kind, cell, batch): outs = cell.forward_train(features_for(kind, batch)) tgt = batch["z"] loss_bits = F.binary_cross_entropy_with_logits(outs["bits"], tgt) sign = 2 * tgt - 1 loss_margin = F.relu(MARGIN - outs["bits"] * sign).mean() if kind == "reduce": loss_aux = ( F.binary_cross_entropy_with_logits(outs["borrow"], batch["borrow"]) + F.cross_entropy(outs["q"], batch["q"]) ) else: loss_aux = F.binary_cross_entropy_with_logits( outs["carry"], batch["carry"]) loss = loss_bits + MARGIN_W * loss_margin + AUX_W * loss_aux loss = loss + 1e-4 * sum(blk.last_h_l1 for blk in cell.blocks) ok_rows = ((outs["bits"] > 0) == (tgt > 0.5)).all(dim=1) return loss, { "loss": loss.item(), "bits": loss_bits.item(), "exact": ok_rows.float().mean().item(), }, ok_rows def bsz_for(n): return min(1024, max(8, TOKEN_BUDGET // n)) def fresh_batch(kind, rng, n, bsz, instances=None): if kind == "reduce": return make_reduce_batch(rng, n, bsz, instances) return make_add_batch(rng, n, bsz, instances) @torch.no_grad() def evaluate(kind, cell, rng, n, device, n_batches=4): cell.eval() total, good = 0, 0 for _ in range(n_batches): b = fresh_batch(kind, rng, n, bsz_for(n)) b.pop("raw") b = {k: v.to(device) for k, v in b.items()} logits = cell(features_for(kind, b)) ok = ((logits > 0) == (b["z"] > 0.5)).all(dim=1) good += int(ok.sum()) total += ok.numel() cell.train() return good / total def set_mode(cell, mode, noise): cell.gate.mode = mode for blk in cell.blocks: blk.scan_noise = noise def main(): ap = argparse.ArgumentParser() ap.add_argument("--resume", default=None) ap.add_argument("--steps", type=int, default=600_000) ap.add_argument("--lr", type=float, default=1e-3) ap.add_argument("--out", default="ckpt2") ap.add_argument("--seed", type=int, default=0) ap.add_argument("--max-stage", type=int, default=None) ap.add_argument("--stages", type=int, nargs="*", default=None) ap.add_argument("--force-stage", type=int, default=None) ap.add_argument("--eval-window", type=int, default=0, help="0 = min over all widths <= stage; K>0 = min over last K stages only") args = ap.parse_args() global STAGES if args.stages: STAGES = args.stages if args.max_stage is None: args.max_stage = len(STAGES) - 1 device = torch.device("mps" if torch.backends.mps.is_available() else "cpu") rng = random.Random(args.seed) torch.manual_seed(args.seed) cells = {"reduce": make_reduce_cell("soft").to(device), "add": make_add_cell("soft").to(device)} opts = {k: torch.optim.AdamW(c.parameters(), lr=args.lr, weight_decay=0.0) for k, c in cells.items()} import copy as _copy ema_cells = {k: _copy.deepcopy(c) for k, c in cells.items()} EMA_DECAY = 0.999 @torch.no_grad() def ema_update(kind): for pe, pr in zip(ema_cells[kind].parameters(), cells[kind].parameters()): pe.lerp_(pr, 1.0 - EMA_DECAY) stage, step, last_advance_step = 0, 0, 0 streak = 0 best_box: dict[int, float] = {} out_dir = Path(args.out) out_dir.mkdir(exist_ok=True) metrics_f = open(out_dir / "metrics.jsonl", "a") if args.resume: ck = torch.load(args.resume, map_location=device, weights_only=True) cells["reduce"].load_state_dict(ck["reduce_state_dict"], strict=False) if "add_state_dict" in ck: cells["add"].load_state_dict(ck["add_state_dict"], strict=False) for k in KINDS: if f"{k}_state_dict" in ck: sdk = ck.get(f"{k}_ema_state_dict", ck[f"{k}_state_dict"]) ema_cells[k].load_state_dict(sdk, strict=False) for k in KINDS: try: opts[k].load_state_dict(ck[f"{k}_opt"]) except Exception: print(f"{k}: fresh optimizer", flush=True) stage, step = ck["stage"], ck["step"] if args.force_stage is not None: stage = args.force_stage last_advance_step = step print(f"resumed at step {step}, stage {stage}", flush=True) def set_lr(v): for o in opts.values(): for g in o.param_groups: g["lr"] = v def get_lr(): return opts["reduce"].param_groups[0]["lr"] set_lr(args.lr) def pick_width(stage_now): if stage_now == 0 or rng.random() < 0.5: n = STAGES[stage_now] else: n = STAGES[rng.randint(0, stage_now - 1)] if rng.random() < 0.3: n = max(5, n - rng.randint(0, max(1, n // 8))) return n q: queue.Queue = queue.Queue(maxsize=4) stage_box = {"stage": stage} hard_buf: dict[tuple, list] = {} buf_lock = threading.Lock() def load_mined(kind): import json as _json path = out_dir / f"mined_{kind}.jsonl" if not path.exists(): return try: lines = path.read_text().splitlines()[-40000:] except Exception: return per_width: dict[int, list] = {} for ln in lines: try: row = _json.loads(ln) except Exception: continue if kind == "reduce": inst = (row["m"], row["x"]) else: inst = (row["x"], row["y"], row["g"]) per_width.setdefault(row["n"], []).append(inst) with buf_lock: for n, insts in per_width.items(): buf = hard_buf.setdefault((kind, n), []) buf.extend(insts) if len(buf) > 20_000: del buf[: len(buf) - 20_000] def producer(): prng = random.Random(args.seed + 1) last_mined = 0.0 while True: if time.time() - last_mined > 120: for k in KINDS: load_mined(k) last_mined = time.time() kind = "reduce" if prng.random() < 0.5 else "add" cur_w = STAGES[min(stage_box["stage"], len(STAGES) - 1)] + 8 with buf_lock: widths = [w for (k, w), v in hard_buf.items() if k == kind and len(v) >= 64 and w <= cur_w] if widths and prng.random() < 0.30: n = prng.choice(widths) bsz = bsz_for(n) with buf_lock: pool = hard_buf[(kind, n)] replay = [pool[prng.randrange(len(pool))] for _ in range(bsz // 2)] b = fresh_batch(kind, prng, n, bsz - len(replay)) br = fresh_batch(kind, prng, n, len(replay), instances=replay) merged = {k: torch.cat([b[k], br[k]]) for k in b if k != "raw"} merged["raw"] = b["raw"] + br["raw"] q.put((kind, n, merged)) else: n = pick_width(stage_box["stage"]) q.put((kind, n, fresh_batch(kind, prng, n, bsz_for(n)))) threading.Thread(target=producer, daemon=True).start() for c in cells.values(): c.train() t0 = time.time() while step < args.steps: step += 1 mode = "soft" if step <= SOFT_STEPS else "ste" noise = 0.0 if step <= SOFT_STEPS else NOISE for c in cells.values(): set_mode(c, mode, noise) stage_box["stage"] = stage kind, n, batch = q.get() raw = batch.pop("raw") batch = {k: v.to(device) for k, v in batch.items()} cell, opt = cells[kind], opts[kind] loss, stats, ok_rows = forward_loss(kind, cell, batch) opt.zero_grad(set_to_none=True) loss.backward() torch.nn.utils.clip_grad_norm_(cell.parameters(), 1.0) opt.step() ema_update(kind) if step > SOFT_STEPS: bad = (~ok_rows).nonzero().flatten().tolist() if bad: with buf_lock: buf = hard_buf.setdefault((kind, n), []) buf.extend(raw[i] for i in bad) if len(buf) > 20_000: del buf[: len(buf) - 20_000] if step % 100 == 0: stats.update(step=step, stage=stage, n=n, kind=kind, lr=get_lr(), sps=round(step / (time.time() - t0), 2)) print(json.dumps(stats), flush=True) metrics_f.write(json.dumps(stats) + "\n") metrics_f.flush() if step % EVAL_EVERY == 0 and step > SOFT_STEPS: accs = {} for k in KINDS: set_mode(ema_cells[k], "hard", 0.0) ema_cells[k].eval() lo = 0 if args.eval_window == 0 else max(0, stage + 1 - args.eval_window) accs[k] = min( evaluate(k, ema_cells[k], rng, STAGES[s], device, n_batches=2 if s < stage else 4) for s in range(lo, stage + 1)) print(json.dumps({"eval": accs, "stage": stage, "n": STAGES[stage], "step": step}), flush=True) metrics_f.write(json.dumps( {"eval": accs, "stage": stage, "step": step}) + "\n") metrics_f.flush() score = min(accs.values()) if score > best_box.get(stage, 0.0): best_box[stage] = score torch.save( {"reduce_state_dict": cells["reduce"].state_dict(), "add_state_dict": cells["add"].state_dict(), "reduce_ema_state_dict": ema_cells["reduce"].state_dict(), "add_ema_state_dict": ema_cells["add"].state_dict(), "stage": stage, "step": step, "score": score}, out_dir / f"best_stage{stage}.pt", ) ok = score >= ADVANCE_THRESHOLD streak = streak + 1 if ok else 0 if streak >= 2 and stage < min(args.max_stage, len(STAGES) - 1): stage += 1 streak = 0 last_advance_step = step set_lr(args.lr) print(f"=== ADVANCE to stage {stage} (N={STAGES[stage]}) ===", flush=True) elif step - last_advance_step >= PLATEAU_STEPS and get_lr() > 5.1e-5: set_lr(max(5e-5, get_lr() * 0.5)) last_advance_step = step print(f"=== LR DECAY to {get_lr():.2e} (plateau) ===", flush=True) torch.save( {"reduce_state_dict": cells["reduce"].state_dict(), "add_state_dict": cells["add"].state_dict(), "reduce_ema_state_dict": ema_cells["reduce"].state_dict(), "add_ema_state_dict": ema_cells["add"].state_dict(), "reduce_opt": opts["reduce"].state_dict(), "add_opt": opts["add"].state_dict(), "stage": stage, "step": step}, out_dir / "latest.pt", ) if step % 5000 == 0: import shutil shutil.copy(out_dir / "latest.pt", out_dir / f"snap_{step}.pt") snaps = sorted(out_dir.glob("snap_*.pt"), key=lambda p: int(p.stem.split("_")[1])) for old in snaps[:-6]: old.unlink() metrics_f.close() if __name__ == "__main__": main()