"""Pretraining loop: compiled model, Muon + Sinkhorn, WSD schedule, resumable, early decay on demand. source env.sh && $TA_PY scripts/train.py --size M --tokens 6e9 --out $TA_DATA/runs/m1 Control while running (no restart needed): touch /DECAY -> start the LR decay now (lasts --decay_frac of the steps done so far) touch /STOP -> checkpoint and exit Snapshots for RL excursions are written every --snapshot_tokens to /snap_.pt. """ import argparse import json import math import os import time from dataclasses import asdict import numpy as np import torch from tiny_agent.data import MixtureLoader from tiny_agent.model import ModelConfig, TinyAgentLM, make_block_mask from tiny_agent.optim import build_optimizers from tiny_agent.text import DATA SIZES = { "S": dict(d_model=512, n_layers=12, n_heads=8), "M": dict(d_model=640, n_layers=16, n_heads=10), "L": dict(d_model=768, n_layers=18, n_heads=12), "XL": dict(d_model=1024, n_layers=20, n_heads=16), } def lr_mult(step, total, warmup, decay_start, decay_steps): if step < warmup: return (step + 1) / warmup if step < decay_start: return 1.0 # 1 - sqrt decay (works well for WSD), floor at 0 frac = min(1.0, (step - decay_start) / max(1, decay_steps)) return max(0.0, 1 - math.sqrt(frac)) def save(path, model, opts, step, tokens, meta): tmp = path + ".tmp" torch.save({"model": model.state_dict(), "opts": [o.state_dict() for o in opts], "step": step, "tokens": tokens, "meta": meta}, tmp) os.replace(tmp, path) def save_snapshot(path, model, cfg, tokens): sd = {k: v.to(torch.bfloat16) if v.is_floating_point() else v for k, v in model.state_dict().items()} torch.save({"model": sd, "config": asdict(cfg), "tokens": tokens}, path) @torch.no_grad() def evaluate(cmodel, batches, cfg, device): out = {} for name, bs in batches.items(): ls = [] for inp, tgt, doc in bs: inp, tgt, doc = inp.to(device), tgt.to(device), doc.to(device) with torch.autocast("xpu", dtype=torch.bfloat16): ls.append(cmodel(inp, doc, make_block_mask(doc, cfg.swa_window), tgt).item()) out[name] = round(float(np.mean(ls)), 4) return out def main(): ap = argparse.ArgumentParser() ap.add_argument("--size", default="M") ap.add_argument("--engram", type=int, default=1) ap.add_argument("--tokens", type=float, default=6e9) ap.add_argument("--T", type=int, default=2048) ap.add_argument("--micro_B", type=int, default=8) ap.add_argument("--decay_T", type=int, default=0, help="context length in the decay phase (same tokens/step)") ap.add_argument("--batch_tokens", type=int, default=262144) ap.add_argument("--lr", type=float, default=3e-3) ap.add_argument("--wd", type=float, default=0.0) ap.add_argument("--warmup", type=int, default=200) ap.add_argument("--decay_frac", type=float, default=0.2) ap.add_argument("--mixture", default="stable") ap.add_argument("--decay_mixture", default="decay") ap.add_argument("--out", required=True) ap.add_argument("--eval_every", type=int, default=250) ap.add_argument("--ckpt_minutes", type=float, default=20) ap.add_argument("--snapshot_tokens", type=float, default=5e8) ap.add_argument("--max_minutes", type=float, default=0, help="stop (after decay) by wall clock; 0 = off") ap.add_argument("--stop_minutes", type=float, default=0, help="hard stop (no decay) for short A/B runs") ap.add_argument("--init", default="", help="start from a snapshot's weights (e.g. a pre-decay snap_*.pt)") ap.add_argument("--seed", type=int, default=0) a = ap.parse_args() os.makedirs(a.out, exist_ok=True) dev = "xpu" torch.manual_seed(a.seed) cfg = ModelConfig(**SIZES[a.size], engram_layers=(1,) if a.engram else (), max_seq_len=max(8192, a.T)) model = TinyAgentLM(cfg).to(dev) cid = f"{DATA}/cid_map.npy" if a.engram and os.path.exists(cid): model.cid_map.copy_(torch.from_numpy(np.load(cid).astype(np.int64))) opts = build_optimizers(model, lr=a.lr, weight_decay=a.wd) accum = max(1, a.batch_tokens // (a.micro_B * a.T)) step_tokens = accum * a.micro_B * a.T total_steps = int(a.tokens // step_tokens) decay_start, decay_steps = int(total_steps * (1 - a.decay_frac)), int(total_steps * a.decay_frac) step, tokens, elapsed0 = 0, 0, 0.0 ck = os.path.join(a.out, "ckpt.pt") meta = {"args": vars(a), "config": asdict(cfg)} if a.init and not os.path.exists(ck): # continue pretraining from a snapshot (weights only; optimizer state starts fresh) sd = torch.load(a.init, map_location=dev, weights_only=False)["model"] model.load_state_dict({k: v.float() if v.is_floating_point() else v for k, v in sd.items()}) print(f"initialized from {a.init}", flush=True) if os.path.exists(ck): st = torch.load(ck, map_location=dev, weights_only=False) model.load_state_dict(st["model"]) for o, s in zip(opts, st["opts"]): o.load_state_dict(s) step, tokens = st["step"], st["tokens"] decay_start = st["meta"].get("decay_start", decay_start) decay_steps = st["meta"].get("decay_steps", decay_steps) total_steps = st["meta"].get("total_steps", total_steps) elapsed0 = st["meta"].get("elapsed_s", 0.0) print(f"resumed at step {step}, {tokens/1e9:.2f}B tokens", flush=True) # static shapes: train (B, T), decay (B', T') and eval are separate graphs instead of one # slower dynamic-shape graph cmodel = torch.compile(model, dynamic=False) in_decay = step >= decay_start dT = a.decay_T or a.T dB = max(1, a.micro_B * a.T // dT) loader = (MixtureLoader("train", a.decay_mixture, dT, dB, seed=a.seed + step) if in_decay else MixtureLoader("train", a.mixture, a.T, a.micro_B, seed=a.seed + step)) val = MixtureLoader("val", a.mixture, a.T, a.micro_B, seed=99, stream=False).fixed_batches(6) print("mixture", loader.describe(), "| params", model.param_counts(), "| steps", total_steps, "| tokens/step", step_tokens, flush=True) log = open(os.path.join(a.out, "log.jsonl"), "a") last_ck, t_start = time.time(), time.time() - elapsed0 # wall clock survives resumes decay_t0, decay_s0 = time.time(), step t0, tok0 = time.time(), tokens next_snap = (tokens // a.snapshot_tokens + 1) * a.snapshot_tokens while step < total_steps: if os.path.exists(os.path.join(a.out, "DECAY")) and step < decay_start: decay_start, decay_steps = step, max(1, int(step * a.decay_frac / (1 - a.decay_frac))) total_steps = decay_start + decay_steps os.remove(os.path.join(a.out, "DECAY")) print(f"decay requested: steps {decay_start}..{total_steps}", flush=True) if a.max_minutes and step < decay_start: # start decay early enough to finish by the wall-clock limit el = (time.time() - t_start) / 60 rate = max(step, 1) / max(el, 1e-6) if el + a.decay_frac / (1 - a.decay_frac) * step / rate >= a.max_minutes * 0.98: decay_start, decay_steps = step, max(1, int(step * a.decay_frac / (1 - a.decay_frac))) total_steps = decay_start + decay_steps print(f"wall-clock decay: steps {decay_start}..{total_steps}", flush=True) if step >= decay_start and not in_decay: in_decay = True loader.close() loader = MixtureLoader("train", a.decay_mixture, dT, dB, seed=a.seed + step) print("decay mixture", loader.describe(), flush=True) decay_t0, decay_s0 = time.time(), step if a.max_minutes and in_decay and (step - decay_s0) in (50, 200, 500, 1000, 2000, 4000): # decay steps can be slower than stable ones (longer context, recompiles): re-fit the # end of the schedule to the remaining wall-clock budget per = (time.time() - decay_t0) / (step - decay_s0) remaining = max(0.0, a.max_minutes * 60 - (time.time() - t_start)) total_steps = step + max(1, int(remaining / per)) decay_steps = total_steps - decay_start print(f"decay re-fit: {per:.2f}s/step, end at step {total_steps}", flush=True) m = lr_mult(step, total_steps, a.warmup, decay_start, decay_steps) for o in opts: for g in o.param_groups: g["lr"] = g["base_lr"] * m loss_acc = 0.0 for _ in range(accum): inp, tgt, doc = loader.next(dev) with torch.autocast("xpu", dtype=torch.bfloat16): loss = cmodel(inp, doc, make_block_mask(doc, cfg.swa_window), tgt) (loss / accum).backward() loss_acc += loss.detach() for o in opts: o.step() o.zero_grad(set_to_none=True) step += 1 tokens += step_tokens if step % 10 == 0: l = (loss_acc / accum).item() if not math.isfinite(l): raise RuntimeError(f"non-finite loss at step {step}") dt = time.time() - t0 rec = {"step": step, "tokens": tokens, "loss": round(l, 4), "lr_mult": round(m, 4), "tok_s": round((tokens - tok0) / dt), "elapsed_min": round((time.time() - t_start) / 60, 1)} t0, tok0 = time.time(), tokens if step % a.eval_every == 0: rec["val"] = evaluate(cmodel, val, cfg, dev) log.write(json.dumps(rec) + "\n") log.flush() print(json.dumps(rec), flush=True) if tokens >= next_snap: save_snapshot(os.path.join(a.out, f"snap_{tokens/1e9:.2f}B.pt"), model, cfg, tokens) next_snap += a.snapshot_tokens stop = os.path.exists(os.path.join(a.out, "STOP")) or \ (a.stop_minutes and time.time() - t_start > a.stop_minutes * 60) if time.time() - last_ck > a.ckpt_minutes * 60 or stop: meta.update(decay_start=decay_start, decay_steps=decay_steps, total_steps=total_steps, elapsed_s=time.time() - t_start) save(ck, model, opts, step, tokens, meta) last_ck = time.time() if stop: if os.path.exists(os.path.join(a.out, "STOP")): os.remove(os.path.join(a.out, "STOP")) print("stopped on request", flush=True) return meta.update(decay_start=decay_start, decay_steps=decay_steps, total_steps=total_steps, elapsed_s=time.time() - t_start) save(ck, model, opts, step, tokens, meta) save_snapshot(os.path.join(a.out, "final.pt"), model, cfg, tokens) print("done", tokens, flush=True) if __name__ == "__main__": main()