ProCreations's picture
download
raw
7.06 kB
"""GPT-2 training arm on real FineWeb10B tokens (Claims 4 & 6, scaled).
One arm per invocation. Arms share identical data order (fixed generator) and,
for seed-matched arms, identical fp32 initialization (flash arms downcast it,
exactly as the paper prescribes at training start).
ref : fp32 params, torch.optim.AdamW, bf16 autocast
flash : cast_model bf16, FlashAdamW(master_weight_bits=24, quantize=True)
linear : as flash but QuantizedTensorSpec overridden to plain linear
group-absmax quantization (no softsign, no sqrt) -- Fig. 5 ablation
Data: modded-nanogpt shard format (256 int32 header: [magic, version, ntok],
then uint16 GPT-2 BPE tokens).
"""
import argparse
import csv
import json
import math
import os
import sys
import time
import numpy as np
import torch
sys.path.insert(0, os.path.dirname(os.path.abspath(__file__)))
from models import GPT, SIZES, param_groups # noqa: E402
def load_shard(path):
header = np.fromfile(path, dtype=np.int32, count=256)
assert header[0] == 20240520, f"bad magic in {path}"
ntok = int(header[2])
toks = np.memmap(path, dtype=np.uint16, mode="r", offset=1024)[:ntok]
return toks
def make_optimizer(arm, model, lr):
if arm == "ref":
return torch.optim.AdamW(param_groups(model), lr=lr, betas=(0.9, 0.95))
from flashoptim import FlashAdamW
if arm == "flash":
return FlashAdamW(param_groups(model), lr=lr, betas=(0.9, 0.95),
master_weight_bits=24, quantize=True)
if arm == "linear":
from flashoptim.optimizers import QuantizedTensorSpec
class LinearQuantAdamW(FlashAdamW):
def _quantized_state_spec(self):
return {"exp_avg": QuantizedTensorSpec(signed=True, softsign=False),
"exp_avg_sq": QuantizedTensorSpec(signed=False, sqrt=False,
softsign=False)}
return LinearQuantAdamW(param_groups(model), lr=lr, betas=(0.9, 0.95),
master_weight_bits=24, quantize=True)
raise ValueError(arm)
def main():
ap = argparse.ArgumentParser()
ap.add_argument("--arm", required=True, choices=["ref", "flash", "linear"])
ap.add_argument("--size", default="gpt-124m")
ap.add_argument("--init-seed", type=int, default=0)
ap.add_argument("--data-seed", type=int, default=1234)
ap.add_argument("--train-shard", required=True)
ap.add_argument("--val-shard", required=True)
ap.add_argument("--batch", type=int, default=8)
ap.add_argument("--ctx", type=int, default=1024)
ap.add_argument("--lr", type=float, default=6e-4)
ap.add_argument("--warmup", type=int, default=300)
ap.add_argument("--max-steps", type=int, default=100000)
ap.add_argument("--time-budget-s", type=float, required=True)
ap.add_argument("--compile", type=int, default=1)
ap.add_argument("--clip", type=float, default=1.0)
ap.add_argument("--out", required=True)
args = ap.parse_args()
dev = torch.device("cuda")
torch.manual_seed(args.init_seed)
cfg = SIZES[args.size]
cfg.ctx = max(cfg.ctx, args.ctx)
model = GPT(cfg) # fp32 init, identical across arms for equal init-seed
n_params = sum(p.numel() for p in model.parameters())
model = model.to(dev)
if args.arm in ("flash", "linear"):
from flashoptim import cast_model
cast_model(model, dtype=torch.bfloat16)
opt = make_optimizer(args.arm, model, args.lr)
if args.compile:
try:
model = torch.compile(model)
except Exception as e: # fall back, but record it
print(f"compile failed: {e}", flush=True)
train = load_shard(args.train_shard)
val = load_shard(args.val_shard)
g = torch.Generator().manual_seed(args.data_seed)
max_start = len(train) - args.ctx - 1
def get_batch(source, idx):
xs = torch.stack([torch.from_numpy(
source[i:i + args.ctx].astype(np.int64)) for i in idx])
ys = torch.stack([torch.from_numpy(
source[i + 1:i + args.ctx + 1].astype(np.int64)) for i in idx])
return xs.pin_memory().to(dev, non_blocking=True), \
ys.pin_memory().to(dev, non_blocking=True)
use_autocast = args.arm == "ref"
rows = []
t0 = time.time()
step = 0
diverged = False
torch.cuda.reset_peak_memory_stats()
while step < args.max_steps and time.time() - t0 < args.time_budget_s:
idx = torch.randint(0, max_start, (args.batch,), generator=g)
x, y = get_batch(train, idx)
lr_now = args.lr * min(1.0, (step + 1) / args.warmup)
for pg in opt.param_groups:
pg["lr"] = lr_now
with torch.autocast("cuda", dtype=torch.bfloat16, enabled=use_autocast):
_, loss = model(x, y)
loss.backward()
if args.clip:
torch.nn.utils.clip_grad_norm_(
[p for grp in opt.param_groups for p in grp["params"]], args.clip)
opt.step()
opt.zero_grad(set_to_none=True)
li = float(loss.detach())
rows.append((step, round(time.time() - t0, 2), li, lr_now))
if step % 50 == 0:
print(f"[{args.arm}] step {step} loss {li:.4f} "
f"({(step+1)*args.batch*args.ctx/(time.time()-t0):.0f} tok/s)", flush=True)
if not math.isfinite(li) or li > 12.0 and step > 200:
diverged = True
print(f"[{args.arm}] DIVERGED at step {step} (loss={li})", flush=True)
break
step += 1
peak = torch.cuda.max_memory_allocated()
# held-out eval on val shard (fixed batches)
model.eval()
gv = torch.Generator().manual_seed(9999)
vlosses = []
with torch.no_grad():
for _ in range(30):
idx = torch.randint(0, len(val) - args.ctx - 1, (args.batch,), generator=gv)
x, y = get_batch(val, idx)
with torch.autocast("cuda", dtype=torch.bfloat16, enabled=use_autocast):
_, l = model(x, y)
if math.isfinite(float(l)):
vlosses.append(float(l))
val_loss = sum(vlosses) / len(vlosses) if vlosses else float("nan")
os.makedirs(os.path.dirname(args.out), exist_ok=True)
with open(args.out + ".csv", "w", newline="") as f:
w = csv.writer(f)
w.writerow(["step", "elapsed_s", "loss", "lr"])
w.writerows(rows)
summary = dict(arm=args.arm, size=args.size, n_params=n_params,
init_seed=args.init_seed, steps=step,
tokens=step * args.batch * args.ctx,
final_loss=rows[-1][2] if rows else None, val_loss=val_loss,
diverged=diverged, peak_bytes=peak,
elapsed_s=round(time.time() - t0, 1),
batch=args.batch, ctx=args.ctx, lr=args.lr,
torch=torch.__version__)
with open(args.out + ".json", "w") as f:
json.dump(summary, f, indent=2)
print("SUMMARY " + json.dumps(summary), flush=True)
if __name__ == "__main__":
main()

Xet Storage Details

Size:
7.06 kB
·
Xet hash:
298ee783f76338bcb01f89b086307008212ab4a6ea8b4c14e875b5350e400557

Xet efficiently stores files, intelligently splitting them into unique chunks and accelerating uploads and downloads. More info.