spec100m / code /train.py
Akahsizrr's picture
squash to reclaim LFS quota
a8f07a3
Raw History Blame Contribute Delete
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
@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()