darioooooo0o's picture
tiny-agent-112m: base + RL weights, tokenizer, code, model card
4397e12 verified
Raw History Blame Contribute Delete
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)
@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()