Download code/scripts/train.py from darioooooo0o/tiny-agent-112m: direct link, hf CLI and curl.
- Browser
- Download file 10.9 kB
-
https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/scripts/train.py
- Command line
-
hf download hf://darioooooo0o/tiny-agent-112m/code/scripts/train.py
-
curl -L -o train.py https://huggingface.co/darioooooo0o/tiny-agent-112m/resolve/main/code/scripts/train.py
10.9 kB
| """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 <out>/DECAY -> start the LR decay now (lasts --decay_frac of the steps done so far) | |
| touch <out>/STOP -> checkpoint and exit | |
| Snapshots for RL excursions are written every --snapshot_tokens to <out>/snap_<Btok>.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) | |
| 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() | |