File size: 10,880 Bytes
4397e12 | 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 | """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()
|