"""Phase 1: Train the BASE model only (next-token LM loss). Goal: produce coherent English sentences. No Medusa heads, no compression loss — just the 270M base transformer. Phase 2 (separate): freeze base, train Medusa heads + compression params. Speed optimizations: - Base-only loss: no 4095-head loop (was OOMing at 210 GB activations) - torch.compile: fuses FFN + attention kernels, eliminates launch overhead - Large batch + long seq: amortize weight reads over more tokens - Gradient checkpointing: trade compute for memory (fit bigger batches) - TF32: use Ampere tensor cores for matmuls - Persistent workers: overlap data loading with compute - Checkpointing: save every N steps + best loss Usage: python train.py --data corpus.txt --steps 10000 --seq 1024 --batch 16 python train.py --steps 10 # smoke test (bundled text) """ from __future__ import annotations import argparse import math import os import time import torch import torch.nn.functional as F import tiktoken from config import Config from model import build_model from data_pipeline import load_tokens # -------------------------------------------------------------------------------------- # Data pipeline # -------------------------------------------------------------------------------------- def load_text(path: str) -> str: with open(path, "r", encoding="utf-8", errors="ignore") as f: return f.read() def tokenize(text: str) -> torch.Tensor: enc = tiktoken.get_encoding("gpt2") return torch.tensor(enc.encode_ordinary(text), dtype=torch.long) def batched(data: torch.Tensor, seq_len: int, batch: int, device): """Yield random [batch, seq_len] slices from a 1D token tensor. Keeps data on the target device — no CPU→GPU copy per step. 2.4M tokens = 19 MB, trivial to keep on GPU. """ n = data.numel() while True: starts = torch.randint(0, n - seq_len - 1, (batch,), device=data.device) yield torch.stack([data[s:s + seq_len] for s in starts]) # -------------------------------------------------------------------------------------- # Training step (base-only, no Medusa, no compression) # -------------------------------------------------------------------------------------- def train_step_base(model, ids, cfg, use_checkpoint=False): """Base-only next-token loss. No Medusa, no compression. Returns (loss, base_loss). This is the cheapest possible training step — just the transformer forward + cross-entropy. ~10x cheaper than the joint train_step. """ h = model(ids, use_checkpoint=use_checkpoint) # [B, T, d] logits = model.lm_head(h) # [B, T, V] base_loss = F.cross_entropy( logits[:, :-1].reshape(-1, cfg.vocab_size), ids[:, 1:].reshape(-1), ) return base_loss, base_loss.detach() # -------------------------------------------------------------------------------------- # Training controller: adaptive LR + batch scheduling based on loss trajectory # -------------------------------------------------------------------------------------- class TrainingController: """Adaptive training controller inspired by ChronosLM-style trajectory prediction. Tracks loss history and adjusts: - Learning rate (cosine + reduce-on-plateau fallback) - Batch size (grow if stable, shrink if unstable) - Gradient clipping strength (based on grad norm history) - Reports predicted convergence step This is NOT a neural controller — it's a lightweight heuristic controller that adapts to the loss curve in real-time. A neural controller would need training data from many runs; this works from run 1. """ def __init__(self, base_lr, warmup_steps, total_steps, min_lr=1e-5): self.base_lr = base_lr self.warmup = warmup_steps self.total_steps = total_steps self.min_lr = min_lr self.loss_history = [] self.grad_norm_history = [] self.step = 0 self.best_loss = float('inf') self.patience = 0 self.max_patience = 100 def lr_at(self, step): """Cosine schedule with warmup.""" if step < self.warmup: return self.base_lr * (step + 1) / self.warmup prog = (step - self.warmup) / max(1, self.total_steps - self.warmup) return self.max(self.min_lr, self.base_lr * 0.5 * (1 + math.cos(math.pi * prog))) def update(self, loss, grad_norm, step): """Record loss + grad norm. Returns (lr, should_checkpoint, message).""" self.loss_history.append(loss) self.grad_norm_history.append(grad_norm) self.step = step + 1 # real step, not update count lr = self.lr_at(step + 1) # Track best loss should_checkpoint = False message = "" if loss < self.best_loss: self.best_loss = loss should_checkpoint = True message = " (new best)" self.patience = 0 else: self.patience += 1 # Detect plateau: if loss hasn't improved in max_patience steps, reduce LR if self.patience >= self.max_patience: lr = max(self.min_lr, lr * 0.5) self.patience = 0 message = " (plateau: LR reduced)" # Detect instability: if grad norm spikes > 5x recent average if len(self.grad_norm_history) > 10: recent_avg = sum(self.grad_norm_history[-10:]) / 10 if grad_norm > recent_avg * 5: message += " [unstable: high grad norm]" # Predict convergence: linear extrapolation of last 50 losses if len(self.loss_history) >= 50: recent = self.loss_history[-50:] x = torch.arange(50, dtype=torch.float) y = torch.tensor(recent) # Linear fit: y = a*x + b a = (y.mean() * x.mean() - (x * y).mean()) / (x.mean()**2 - (x**2).mean()) if a < 0: # loss is decreasing steps_to_converge = int((2.0 - y[-1].item()) / abs(a.item())) if steps_to_converge > 0 and steps_to_converge < 100000: message += f" [~{steps_to_converge} steps to loss=2.0]" return lr, should_checkpoint, message @staticmethod def max(a, b): return a if a > b else b # -------------------------------------------------------------------------------------- # Main training loop # -------------------------------------------------------------------------------------- def main(): ap = argparse.ArgumentParser() ap.add_argument("--data", type=str, default="wikitext2", help="dataset name (wikitext2, tinyshakespeare) or path to .txt") ap.add_argument("--steps", type=int, default=10000) ap.add_argument("--seq", type=int, default=1024) ap.add_argument("--batch", type=int, default=16) ap.add_argument("--lr", type=float, default=3e-4) ap.add_argument("--warmup", type=int, default=200) ap.add_argument("--out", type=str, default="checkpoints/base") ap.add_argument("--checkpoint_every", type=int, default=500) ap.add_argument("--compile", action="store_true", default=True, help="torch.compile the model (default: on)") ap.add_argument("--no-compile", dest="compile", action="store_false") ap.add_argument("--tf32", action="store_true", default=True, help="enable TF32 for Ampere+ (default: on)") ap.add_argument("--grad-clip", type=float, default=1.0) ap.add_argument("--grad-checkpoint", action="store_true", default=False, help="use gradient checkpointing (saves memory, ~30% slower)") ap.add_argument("--resume", type=str, default=None, help="checkpoint .pt to resume model weights from") args = ap.parse_args() device = "cuda" # Enable TF32 for Ampere+ (A6000 supports it) if args.tf32: torch.backends.cuda.matmul.allow_tf32 = True torch.backends.cudnn.allow_tf32 = True print("TF32: enabled (Ampere tensor cores)") cfg = Config.v5_500m() model = build_model(cfg, device) # Count base-only params (exclude spec heads for reporting) base_params = sum(p.numel() for n, p in model.named_parameters() if not n.startswith(("medusa_", "spec_"))) total_params = sum(p.numel() for p in model.parameters()) print(f"model: {total_params/1e6:.1f}M total base: {base_params/1e6:.1f}M " f"medusa: {(total_params-base_params)/1e6:.1f}M (frozen this phase)") print(f"config: d={cfg.d_model} L={cfg.n_layers} H={cfg.n_heads} " f"K={cfg.medusa_heads} seq={args.seq} batch={args.batch}") # Freeze spec-head params (phase 1: base only). Covers both naming schemes. for name, param in model.named_parameters(): if name.startswith(("medusa_", "spec_")): param.requires_grad = False trainable = sum(p.numel() for p in model.parameters() if p.requires_grad) print(f"trainable: {trainable/1e6:.1f}M params (base only)") # Resume from checkpoint (model weights only — optimizer state resets) if args.resume: ckpt = torch.load(args.resume, map_location=device) missing, unexpected = model.load_state_dict(ckpt["model"], strict=False) print(f"resumed: {args.resume} (step {ckpt.get('step')}, " f"loss {ckpt.get('loss'):.4f}) — {len(missing)} missing, " f"{len(unexpected)} unexpected keys") # torch.compile for fused kernels if args.compile: print("compiling model (torch.compile, default mode)...") # Use default mode (not max-autotune which hangs on autotuning) # Disable CUDA graphs (incompatible with RoPE precompute) torch._inductor.config.triton.cudagraph_trees = False t0 = time.perf_counter() model = torch.compile(model, mode="default", dynamic=False) print(f" compiled in {time.perf_counter()-t0:.1f}s") # Data — move to GPU once (103M tokens = ~400 MB int64) if args.data in ("wikitext2", "wikitext103", "mixture250m", "mixture500m", "mixture1b", "tinyshakespeare", "openwebtext_10k"): data = load_tokens(args.data).to(device) print(f"data: {args.data} ({data.numel():,} tokens, on {device})") elif args.data and os.path.exists(args.data): text = load_text(args.data) data = tokenize(text).to(device) print(f"data: {args.data} ({data.numel():,} tokens, on {device})") else: text = ("Lorem ipsum dolor sit amet, consectetur adipiscing elit. " * 4000) data = tokenize(text).to(device) print(f"data: bundled text ({data.numel():,} tokens, on {device}) [smoke test]") # Optimizer (only trainable params) opt = torch.optim.AdamW( [p for p in model.parameters() if p.requires_grad], lr=args.lr, betas=(0.9, 0.95), weight_decay=0.1, fused=True) # Controller controller = TrainingController(args.lr, args.warmup, args.steps) # Checkpoint dir os.makedirs(args.out, exist_ok=True) loader = batched(data, args.seq, args.batch, device) model.train() print(f"\n=== Training (phase 1: base only) ===") print(f" steps: {args.steps}") print(f" tokens/step: {args.batch * args.seq}") print(f" total tokens: {args.steps * args.batch * args.seq:,}") print(f" checkpoint: every {args.checkpoint_every} steps + best loss") print() t_start = time.perf_counter() tokens_total = 0 # Cache trainable params (avoid list comprehension every step) trainable_params = [p for p in model.parameters() if p.requires_grad] for step in range(args.steps): ids = next(loader) opt.zero_grad(set_to_none=True) loss, bl = train_step_base(model, ids, cfg, use_checkpoint=args.grad_checkpoint) loss.backward() grad_norm = torch.nn.utils.clip_grad_norm_(trainable_params, args.grad_clip) opt.step() # Controller update (every 20 steps to reduce CPU overhead) if step % 20 == 0 or step == args.steps - 1: lr, should_ckpt, msg = controller.update(bl.item(), grad_norm.item(), step) for pg in opt.param_groups: pg["lr"] = lr else: # Fast path: just update LR from cosine schedule, no CPU sync lr = controller.lr_at(step + 1) for pg in opt.param_groups: pg["lr"] = lr should_ckpt = False msg = "" tokens_total += args.batch * args.seq if step % 20 == 0 or step == args.steps - 1: elapsed = time.perf_counter() - t_start tok_s = tokens_total / elapsed if elapsed > 0 else 0 print(f"step {step:5d} lr {lr:.2e} loss {bl.item():.4f} " f"grad {grad_norm.item():.2f} {tok_s:,.0f} tok/s{msg}") # Periodic checkpoint if (step + 1) % args.checkpoint_every == 0: ckpt_path = os.path.join(args.out, f"step_{step+1}.pt") torch.save({ "model": model.state_dict(), "cfg": cfg.__dict__, "step": step + 1, "loss": bl.item(), }, ckpt_path) print(f" checkpoint: {ckpt_path}") # Best-loss checkpoint if should_ckpt and bl.item() < float('inf'): ckpt_path = os.path.join(args.out, "best.pt") torch.save({ "model": model.state_dict(), "cfg": cfg.__dict__, "step": step + 1, "loss": bl.item(), }, ckpt_path) # Final checkpoint ckpt_path = os.path.join(args.out, "final.pt") torch.save({ "model": model.state_dict(), "cfg": cfg.__dict__, "step": args.steps, "loss": bl.item(), }, ckpt_path) elapsed = time.perf_counter() - t_start print(f"\n=== Done ===") print(f" steps: {args.steps}") print(f" tokens: {tokens_total:,}") print(f" time: {elapsed:.1f}s") print(f" speed: {tokens_total/elapsed:,.0f} tok/s") print(f" best loss: {controller.best_loss:.4f}") print(f" final checkpoint: {ckpt_path}") if __name__ == "__main__": main()