"""Pre-train Sprout: show it text, ask for the next token, nudge the weights. python train.py --out runs/base --tokens 400e6 Writes runs//log.jsonl (loss, learning rate, samples as it learns) and checkpoints ckpt_.pt + ckpt_final.pt. """ import argparse import json import math import os import time import numpy as np import torch from model import GPT, Config from muon import Muon from tokenizer import Tokenizer PROMPTS = ['Once upon a time', 'Lily looked at the sky and said', 'Tom: Hi! How are you today?\nAnna:', 'The best thing about summer is'] def get_device(): if torch.cuda.is_available(): return 'cuda' if torch.backends.mps.is_available(): return 'mps' return 'cpu' class Batches: """Random windows of context+1 tokens from a flat file of token ids.""" def __init__(self, path, device, seed=0): self.data = np.memmap(path, dtype=np.uint16, mode='r') self.device = device self.rng = np.random.default_rng(seed) def get(self, batch, context): starts = self.rng.integers(0, len(self.data) - context - 1, batch) rows = np.stack([self.data[s:s + context + 1] for s in starts]).astype(np.int64) rows = torch.from_numpy(rows).to(self.device, non_blocking=True) return rows[:, :-1], rows[:, 1:] def lr_factor(step, total, warmup, cooldown): """Warm up, hold, then decay linearly to zero over the last `cooldown` share.""" if step < warmup: return (step + 1) / warmup decay_start = total * (1 - cooldown) if step < decay_start: return 1.0 return max(0.0, (total - step) / (total - decay_start)) def context_at(step, total, context, warm): """Sequence-length warm-up: short windows first, full length later.""" if not warm: return context if step < 0.1 * total: return max(64, context // 4) if step < 0.3 * total: return context // 2 return context def make_optimizers(model, args): matrices = [p for n, p in model.named_parameters() if p.ndim == 2 and 'embed' not in n] others = [p for n, p in model.named_parameters() if not (p.ndim == 2 and 'embed' not in n)] embed = [p for p in others if p.ndim == 2] gains = [p for p in others if p.ndim < 2] adam_groups = [dict(params=embed, lr=args.lr_embed, weight_decay=0.0), dict(params=gains, lr=args.lr_embed, weight_decay=0.0)] if args.optim == 'muon': split = {blk.attn.qkv.weight: 3 for blk in model.blocks} muon = Muon(matrices, lr=args.lr, weight_decay=args.wd, split=split) adam = torch.optim.AdamW(adam_groups, betas=(0.9, 0.95), eps=1e-10) return [muon, adam] adam_groups.append(dict(params=matrices, lr=args.lr, weight_decay=args.wd)) return [torch.optim.AdamW(adam_groups, betas=(0.9, 0.95), eps=1e-10)] @torch.no_grad() def evaluate(model, val, args, ctx): model.eval() rng_state = val.rng val.rng = np.random.default_rng(1234) # the same batches every time losses = [] for _ in range(args.eval_batches): x, y = val.get(args.batch_tokens // args.context, args.context) with ctx: _, loss = model(x, y) losses.append(loss.item()) val.rng = rng_state model.train() return sum(losses) / len(losses) @torch.no_grad() def samples(model, tok, device, n_tokens=80): model.eval() out = [] g = torch.Generator(device='cpu').manual_seed(7) for prompt in PROMPTS: torch.manual_seed(int(torch.randint(0, 2**31, (1,), generator=g))) idx = torch.tensor([[tok['<|endoftext|>']] + tok.encode(prompt)], device=device) ids = model.generate(idx, n_tokens, temperature=0.8, top_p=0.95, stop={tok['<|endoftext|>']}) out.append(tok.decode([i for i in ids[0, 1:].tolist() if i != tok['<|endoftext|>']])) model.train() return out def main(): ap = argparse.ArgumentParser() ap.add_argument('--data', default='data') ap.add_argument('--out', default='runs/base') ap.add_argument('--d-model', type=int, default=384) ap.add_argument('--n-layer', type=int, default=8) ap.add_argument('--n-head', type=int, default=6) ap.add_argument('--d-ff', type=int, default=1024) ap.add_argument('--context', type=int, default=512) ap.add_argument('--batch-tokens', type=int, default=32768) ap.add_argument('--tokens', type=float, default=400e6, help='how many tokens to train on') ap.add_argument('--optim', choices=['muon', 'adamw'], default='muon') ap.add_argument('--lr', type=float, default=0.02, help='Muon lr (or AdamW lr for matrices)') ap.add_argument('--lr-embed', type=float, default=0.006) ap.add_argument('--wd', type=float, default=0.0) ap.add_argument('--warmup', type=int, default=200) ap.add_argument('--cooldown', type=float, default=0.3) ap.add_argument('--seq-warmup', action='store_true') ap.add_argument('--eval-every', type=int, default=250) ap.add_argument('--eval-batches', type=int, default=20) ap.add_argument('--sample-every', type=int, default=500) ap.add_argument('--sample-at', default='', help='extra steps with samples, e.g. 25,50,100') ap.add_argument('--save-at', default='', help='extra checkpoint steps, e.g. 100,1000,5000') ap.add_argument('--no-compile', action='store_true') ap.add_argument('--seed', type=int, default=0) ap.add_argument('--device', default='', help='cuda / mps / cpu (default: the best available)') ap.add_argument('--threads', type=int, default=0, help='CPU threads (0 = default)') args = ap.parse_args() torch.manual_seed(args.seed) if args.threads: torch.set_num_threads(args.threads) device = args.device or get_device() os.makedirs(args.out, exist_ok=True) tok = Tokenizer.load(f'{args.data}/tokenizer.json') cfg = Config(vocab_size=tok.vocab_size, context=args.context, n_layer=args.n_layer, n_head=args.n_head, d_model=args.d_model, d_ff=args.d_ff) model = GPT(cfg).to(device) print(f'{model.num_params() / 1e6:.2f}M parameters on {device}') train_model = model if args.no_compile else torch.compile(model) opts = make_optimizers(model, args) for opt in opts: for g in opt.param_groups: g['base_lr'] = g['lr'] ctx = torch.autocast(device_type=device, dtype=torch.bfloat16) if device != 'cpu' else torch.autocast('cpu', enabled=False) train = Batches(f'{args.data}/train.bin', device, seed=args.seed) val = Batches(f'{args.data}/val.bin', device) steps = int(args.tokens // args.batch_tokens) save_at = {int(s) for s in args.save_at.split(',') if s} sample_at = {int(s) for s in args.sample_at.split(',') if s} log = open(f'{args.out}/log.jsonl', 'a') json.dump({'config': cfg.to_dict(), 'args': vars(args), 'steps': steps, 'params': model.num_params()}, log) log.write('\n') t0, seen = time.time(), 0 for step in range(steps + 1): last = step == steps if step % args.eval_every == 0 or step in sample_at or last: vl = evaluate(train_model, val, args, ctx) rec = {'step': step, 'tokens': seen, 'val': round(vl, 4), 'time': round(time.time() - t0, 1)} if step % args.sample_every == 0 or step in sample_at or last: rec['samples'] = samples(model, tok, device) print(json.dumps(rec), flush=True) log.write(json.dumps(rec) + '\n') log.flush() if step in save_at or last: name = 'final' if last else step torch.save({'config': cfg.to_dict(), 'model': model.state_dict(), 'step': step, 'tokens': seen}, f'{args.out}/ckpt_{name}.pt') if last: break T = context_at(step, steps, args.context, args.seq_warmup) x, y = train.get(args.batch_tokens // T, T) with ctx: _, loss = train_model(x, y) loss.backward() norm = torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0) f = lr_factor(step, steps, args.warmup, args.cooldown) for opt in opts: for g in opt.param_groups: g['lr'] = g['base_lr'] * f opt.step() opt.zero_grad(set_to_none=True) seen += x.numel() if step % 20 == 0: rec = {'step': step, 'loss': round(loss.item(), 4), 'lr': round(f, 4), 'norm': round(norm.item(), 3), 'tokens': seen, 'tok_s': round(seen / (time.time() - t0)), 'T': T} log.write(json.dumps(rec) + '\n') if step % 100 == 0: print(json.dumps(rec), flush=True) log.close() if __name__ == '__main__': main()