Download code/train.py from Akahsizrr/spec100m: direct link, hf CLI and curl.
- Browser
- Download file 14.7 kB
-
https://huggingface.co/Akahsizrr/spec100m/resolve/main/code/train.py
- Command line
-
hf download hf://Akahsizrr/spec100m/code/train.py
-
curl -L -o train.py https://huggingface.co/Akahsizrr/spec100m/resolve/main/code/train.py
14.7 kB
| """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 | |
| 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() | |