import argparse import itertools import json import math import os import random import time from pathlib import Path import numpy as np import torch import torch.nn as nn from datasets import load_dataset from tqdm import trange from model import MiniLM, ModelConfig, count_parameters def parse_args() -> argparse.Namespace: p = argparse.ArgumentParser(description="Train a ~10M parameter MiniLM on TinyStories") p.add_argument("--dataset", type=str, default="karpathy/tinystories-gpt4-clean") p.add_argument("--split", type=str, default="train") p.add_argument("--text-column", type=str, default="text") p.add_argument("--num-rows", type=int, default=10_000) p.add_argument("--cache-dir", type=str, default=".cache/huggingface") p.add_argument("--max-seq-len", type=int, default=16_384, help="Model context window") p.add_argument( "--train-seq-len", type=int, default=2_048, help="Actual training sequence length (<= max-seq-len)", ) p.add_argument("--batch-size", type=int, default=2) p.add_argument("--grad-accum", type=int, default=8) p.add_argument("--max-steps", type=int, default=500) p.add_argument("--eval-every", type=int, default=50) p.add_argument("--eval-batches", type=int, default=20) p.add_argument("--lr", type=float, default=3e-4) p.add_argument("--weight-decay", type=float, default=0.1) p.add_argument("--warmup-steps", type=int, default=50) p.add_argument("--grad-clip", type=float, default=1.0) p.add_argument("--dropout", type=float, default=0.0) p.add_argument("--seed", type=int, default=1337) p.add_argument("--out-dir", type=str, default="runs/tiny10m") p.add_argument("--save-every", type=int, default=100) p.add_argument("--device", type=str, default="auto") p.add_argument("--dtype", type=str, default="auto", choices=["auto", "float32", "bfloat16"]) p.add_argument("--compile", action="store_true", help="torch.compile model") p.add_argument( "--dry-run", action="store_true", help="Build model and run one forward pass on random data, then exit", ) return p.parse_args() def set_seed(seed: int): random.seed(seed) np.random.seed(seed) torch.manual_seed(seed) torch.cuda.manual_seed_all(seed) def detect_device(user_value: str) -> torch.device: if user_value != "auto": return torch.device(user_value) return torch.device("cuda" if torch.cuda.is_available() else "cpu") def detect_dtype(user_value: str, device: torch.device) -> torch.dtype: if user_value == "float32": return torch.float32 if user_value == "bfloat16": return torch.bfloat16 if device.type == "cuda" and torch.cuda.is_bf16_supported(): return torch.bfloat16 return torch.float32 def encode_byte_level(text: str) -> list[int]: return list(text.encode("utf-8", errors="ignore")) def load_tokens( dataset_name: str, split: str, text_column: str, num_rows: int, cache_dir: str, ) -> np.ndarray: ds = load_dataset(dataset_name, split=split, streaming=True, cache_dir=cache_dir) all_tokens: list[int] = [] sep = [10, 10] # two newlines rows_seen = 0 for row in itertools.islice(ds, num_rows): if text_column not in row: raise KeyError(f"Column '{text_column}' not found in dataset row keys={list(row.keys())}") text = row[text_column] all_tokens.extend(encode_byte_level(text)) all_tokens.extend(sep) rows_seen += 1 if len(all_tokens) < 2048: raise ValueError(f"Not enough tokens to train: got {len(all_tokens)}") if rows_seen == 0: raise ValueError("No rows were streamed from the dataset.") return np.asarray(all_tokens, dtype=np.uint16) def build_model(args: argparse.Namespace) -> MiniLM: cfg = ModelConfig( vocab_size=256, d_model=352, n_heads=8, n_layers=7, ffn_mult=4, max_seq_len=args.max_seq_len, dropout=args.dropout, ) model = MiniLM(cfg) return model def get_batch(tokens: torch.Tensor, seq_len: int, batch_size: int, device: torch.device): max_start = tokens.size(0) - seq_len - 1 idx = torch.randint(0, max_start, (batch_size,), device=tokens.device) x = torch.stack([tokens[i : i + seq_len] for i in idx]) y = torch.stack([tokens[i + 1 : i + seq_len + 1] for i in idx]) return x.to(device, non_blocking=True), y.to(device, non_blocking=True) def cosine_lr(step: int, max_steps: int, warmup_steps: int, base_lr: float) -> float: if step < warmup_steps: return base_lr * (step + 1) / max(1, warmup_steps) progress = (step - warmup_steps) / max(1, max_steps - warmup_steps) return base_lr * 0.5 * (1.0 + math.cos(math.pi * progress)) @torch.no_grad() def evaluate( model: nn.Module, val_tokens: torch.Tensor, seq_len: int, eval_batches: int, batch_size: int, device: torch.device, amp_dtype: torch.dtype, ) -> float: model.eval() losses = [] for _ in range(eval_batches): xb, yb = get_batch(val_tokens, seq_len, batch_size, device) with torch.autocast(device_type=device.type, dtype=amp_dtype, enabled=device.type == "cuda"): _, loss = model(xb, yb) losses.append(loss.item()) model.train() return float(np.mean(losses)) def save_checkpoint( out_dir: Path, model: nn.Module, optimizer: torch.optim.Optimizer, step: int, train_loss: float, val_loss: float | None, ): out_dir.mkdir(parents=True, exist_ok=True) ckpt_path = out_dir / f"step_{step:06d}.pt" torch.save( { "model": model.state_dict(), "optimizer": optimizer.state_dict(), "step": step, "train_loss": train_loss, "val_loss": val_loss, }, ckpt_path, ) def main(): args = parse_args() if args.train_seq_len > args.max_seq_len: raise ValueError("--train-seq-len must be <= --max-seq-len") set_seed(args.seed) device = detect_device(args.device) amp_dtype = detect_dtype(args.dtype, device) model = build_model(args).to(device) n_params = count_parameters(model) print(f"Device: {device}") print(f"AMP dtype: {amp_dtype}") print(f"Model parameters: {n_params:,}") print(f"Max context window: {args.max_seq_len}") if args.compile: model = torch.compile(model) if args.dry_run: x = torch.randint(0, 256, (2, min(128, args.train_seq_len)), device=device) y = torch.randint(0, 256, (2, min(128, args.train_seq_len)), device=device) _, loss = model(x, y) print(f"Dry-run loss: {loss.item():.4f}") return print( f"Streaming dataset {args.dataset} ({args.split}) and consuming first {args.num_rows} rows..." ) t0 = time.time() os.makedirs(args.cache_dir, exist_ok=True) tokens = load_tokens( args.dataset, args.split, args.text_column, args.num_rows, args.cache_dir, ) elapsed = time.time() - t0 print(f"Loaded {len(tokens):,} byte tokens in {elapsed:.1f}s") split_idx = int(0.95 * len(tokens)) train_tokens = torch.from_numpy(tokens[:split_idx]).long() val_tokens = torch.from_numpy(tokens[split_idx:]).long() if device.type == "cuda": train_tokens = train_tokens.pin_memory() val_tokens = val_tokens.pin_memory() min_needed = args.train_seq_len + 2 if len(train_tokens) < min_needed or len(val_tokens) < min_needed: raise ValueError( f"Tokenized dataset is too small for train_seq_len={args.train_seq_len}. " f"Need at least {min_needed} tokens in each split." ) optimizer = torch.optim.AdamW( model.parameters(), lr=args.lr, betas=(0.9, 0.95), weight_decay=args.weight_decay, fused=(device.type == "cuda"), ) out_dir = Path(args.out_dir) out_dir.mkdir(parents=True, exist_ok=True) with open(out_dir / "config.json", "w", encoding="utf-8") as f: json.dump(vars(args), f, indent=2) print( f"Starting training: steps={args.max_steps}, batch={args.batch_size}, " f"grad_accum={args.grad_accum}, seq_len={args.train_seq_len}" ) model.train() running_loss = 0.0 for step in trange(1, args.max_steps + 1): step_loss = 0.0 lr = cosine_lr(step - 1, args.max_steps, args.warmup_steps, args.lr) for pg in optimizer.param_groups: pg["lr"] = lr optimizer.zero_grad(set_to_none=True) for _ in range(args.grad_accum): xb, yb = get_batch(train_tokens, args.train_seq_len, args.batch_size, device) with torch.autocast(device_type=device.type, dtype=amp_dtype, enabled=device.type == "cuda"): _, loss = model(xb, yb) loss = loss / args.grad_accum loss.backward() step_loss += loss.item() if args.grad_clip > 0: torch.nn.utils.clip_grad_norm_(model.parameters(), args.grad_clip) optimizer.step() running_loss += step_loss if step % args.eval_every == 0 or step == 1: train_loss = running_loss / args.eval_every if step > 1 else running_loss running_loss = 0.0 val_loss = evaluate( model=model, val_tokens=val_tokens, seq_len=args.train_seq_len, eval_batches=args.eval_batches, batch_size=args.batch_size, device=device, amp_dtype=amp_dtype, ) ppl = math.exp(min(20.0, val_loss)) print( f"step={step:05d} lr={lr:.2e} train_loss={train_loss:.4f} " f"val_loss={val_loss:.4f} val_ppl={ppl:.2f}" ) if step % args.save_every == 0 or step == args.max_steps: latest_val = evaluate( model=model, val_tokens=val_tokens, seq_len=args.train_seq_len, eval_batches=max(1, args.eval_batches // 2), batch_size=args.batch_size, device=device, amp_dtype=amp_dtype, ) save_checkpoint(out_dir, model, optimizer, step, step_loss, latest_val) print(f"Training complete. Checkpoints saved in: {out_dir}") if __name__ == "__main__": main()