summaryrefslogtreecommitdiff
path: root/scripts/train.py
blob: af4391816c371ac019fb079a6adda750f8c2d7f1 (plain)
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
"""torchrun-able BP/ZBP trainer.  Usage:
  torchrun --nproc_per_node=8 scripts/train.py --model configs/model/m124.yaml --train configs/train/zbp_n16.yaml \
      --data data/fineweb --out runs/m124_zbp16 [--set key=value ...]
Single-GPU: python scripts/train.py ... [--device cuda:0]"""
import os, sys, json, math, time, argparse
sys.path.insert(0, os.path.join(os.path.dirname(__file__), "..", "src"))
import yaml
import torch
import torch.nn as nn
import torch.distributed as dist

from zbp_scaling.model import ScalingLM
from zbp_scaling.data import Shards
from zbp_scaling.zbp import ZBPConfig, zbp_blocks
from zbp_scaling.zbp.autograd import seed_probes

p = argparse.ArgumentParser()
p.add_argument("--model", required=True)
p.add_argument("--train", required=True)
p.add_argument("--data", required=True)
p.add_argument("--out", default="runs/dev")
p.add_argument("--device", default=None)
p.add_argument("--steps", type=int, default=None, help="override: stop after this many steps (smoke)")
p.add_argument("--micro_bs", type=int, default=8)
p.add_argument("--set", nargs="*", default=[], help="key=value overrides for the yaml configs")
a = p.parse_args()
M = yaml.safe_load(open(a.model)); Tr = yaml.safe_load(open(a.train))
for kv in a.set:
    k, v = kv.split("=", 1)
    tgt = M if k in M else Tr
    tgt[k] = yaml.safe_load(v)

ddp = "WORLD_SIZE" in os.environ and int(os.environ["WORLD_SIZE"]) > 1
rank, world = (int(os.environ.get("RANK", 0)), int(os.environ.get("WORLD_SIZE", 1)))
if ddp:
    dist.init_process_group("nccl")
    dev = torch.device(f"cuda:{int(os.environ['LOCAL_RANK'])}")
    torch.cuda.set_device(dev)
else:
    dev = torch.device(a.device or ("cuda:0" if torch.cuda.is_available() else "cpu"))
torch.manual_seed(Tr.get("seed", 0) + rank)
seed_probes(Tr.get("seed", 0) * 1000 + rank, dev)

cfg = ZBPConfig(mode=Tr["mode"], n_probes=Tr.get("n_probes", 4), probe=Tr.get("probe", "rademacher"),
                eps=Tr.get("eps", 0.1), probe_chunk=Tr.get("probe_chunk", 8))
model = ScalingLM(M["vocab"], M["d"], M["layers"], M["heads"], M["seq_len"], M.get("ffn_mult", "8/3"), cfg).to(dev)
raw = model
if ddp:
    model = torch.nn.parallel.DistributedDataParallel(model, device_ids=[dev.index])
n_params = sum(q.numel() for q in raw.parameters())

seq = M["seq_len"]
accum = max(1, int(Tr["global_batch_tokens"]) // (seq * a.micro_bs * world))
total_steps = a.steps or int(float(Tr["tokens"]) // int(Tr["global_batch_tokens"]))
opt = torch.optim.AdamW(raw.parameters(), lr=float(Tr["lr"]), weight_decay=Tr.get("wd", 0.1), betas=tuple(Tr.get("betas", [0.9, 0.95])))
sched = torch.optim.lr_scheduler.LambdaLR(opt, lambda s: min(1.0, (s + 1) / Tr.get("warmup", 2000)) * 0.5 * (1 + math.cos(math.pi * min(s, total_steps) / total_steps)))
data = Shards(a.data, seq, dev)
gen = torch.Generator(); gen.manual_seed(Tr.get("seed", 0) * 77 + rank)

os.makedirs(a.out, exist_ok=True)
logf = open(os.path.join(a.out, f"log_rank{rank}.jsonl"), "a")
def log(**kw):
    if rank == 0:
        kw["time"] = time.time()
        logf.write(json.dumps(kw) + "\n"); logf.flush(); print(json.dumps(kw), flush=True)

@torch.no_grad()
def evaluate(nb=20):
    model.eval(); tot = 0.0
    for _ in range(nb):
        x, y = data.batch("val", a.micro_bs, gen)
        tot += nn.functional.cross_entropy(model(x).flatten(0, 1), y.flatten()).item()
    model.train(); return tot / nb

if rank == 0:
    log(kind="meta", params=n_params, accum=accum, world=world, total_steps=total_steps, cfg=cfg.asdict(), model=M, train={k: str(v) for k, v in Tr.items()})
t0 = time.time()
for step in range(total_steps + 1):
    if step % Tr.get("eval_every", 500) == 0 or step == total_steps:
        log(kind="eval", step=step, val_loss=evaluate(), elapsed=time.time() - t0,
            tok_per_s=step * int(Tr["global_batch_tokens"]) / max(time.time() - t0, 1e-9))
        if rank == 0 and step % (10 * Tr.get("eval_every", 500)) == 0:
            torch.save({"model": raw.state_dict(), "opt": opt.state_dict(), "step": step}, os.path.join(a.out, "ckpt.pt"))
    if step == total_steps:
        break
    opt.zero_grad(set_to_none=True)
    for micro in range(accum):
        x, y = data.batch("train", a.micro_bs, gen)
        loss = nn.functional.cross_entropy(model(x).flatten(0, 1), y.flatten()) / accum
        loss.backward()
    if Tr.get("grad_clip", 0):
        torch.nn.utils.clip_grad_norm_(raw.parameters(), Tr["grad_clip"])
    opt.step(); sched.step()
    if step % 50 == 0:
        log(kind="train", step=step, loss=loss.item() * accum, lr=sched.get_last_lr()[0])
if ddp:
    dist.destroy_process_group()