Buckets:
| """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.