"""Profile training speed to find the bottleneck. Current train.py: batch=8, seq=512 = 4096 tokens/step Goal: find what limits training throughput and how to 10x it. """ import time import torch import torch.nn.functional as F from config import Config from model import build_model from train import train_step, medusa_targets cfg = Config.v5_500m() model = build_model(cfg, "cuda") model.train() print(f"GPU: {torch.cuda.get_device_name(0)}") print(f"Model: {sum(p.numel() for p in model.parameters())/1e6:.1f}M params") print(f"Model VRAM: {torch.cuda.memory_allocated()/1e9:.2f} GB") print() # --- Test different batch/seq configurations --- configs = [ (8, 512), # current default (8, 1024), (8, 2048), (8, 4096), # match inference chunk size (16, 1024), (16, 2048), (32, 1024), (32, 2048), (4, 4096), (8, 4096), ] opt = torch.optim.AdamW(model.parameters(), lr=3e-4, betas=(0.9, 0.95), weight_decay=0.1, fused=True) print(f"{'batch':>6} {'seq':>6} {'tokens/step':>11} {'step(ms)':>9} {'tok/s':>10} {'VRAM(GB)':>9} {'fwd_bwd':>8}") print("-" * 70) for batch, seq in configs: try: ids = torch.randint(0, cfg.vocab_size, (batch, seq), device="cuda") # Warmup for _ in range(3): opt.zero_grad(set_to_none=True) loss, _, _ = train_step(model, ids, cfg) loss.backward() opt.step() torch.cuda.synchronize() # Measure torch.cuda.reset_peak_memory_stats() times = [] for _ in range(10): torch.cuda.synchronize() t0 = time.perf_counter() opt.zero_grad(set_to_none=True) loss, _, _ = train_step(model, ids, cfg) loss.backward() opt.step() torch.cuda.synchronize() times.append(time.perf_counter() - t0) times.sort() step_ms = times[5] * 1000 tokens = batch * seq tok_s = tokens / (step_ms / 1000) vram = torch.cuda.max_memory_allocated() / 1e9 # Separate fwd vs bwd torch.cuda.synchronize() t0 = time.perf_counter() opt.zero_grad(set_to_none=True) loss, _, _ = train_step(model, ids, cfg) torch.cuda.synchronize() fwd_ms = (time.perf_counter() - t0) * 1000 t0 = time.perf_counter() loss.backward() torch.cuda.synchronize() bwd_ms = (time.perf_counter() - t0) * 1000 print(f"{batch:>6} {seq:>6} {tokens:>11} {step_ms:>9.1f} {tok_s:>10,.0f} {vram:>9.2f} {fwd_ms:>4.0f}/{bwd_ms:.0f}") except RuntimeError as e: if "out of memory" in str(e).lower(): print(f"{batch:>6} {seq:>6} {batch*seq:>11} {'OOM':>9}") torch.cuda.empty_cache() else: raise print() print("=== Key metrics ===") print(f" Model weights: {sum(p.numel() for p in model.parameters())/1e6:.1f}M params = {sum(p.numel() for p in model.parameters())*2/1e9:.2f} GB (bf16)") print(f" AdamW state: 2x weights = {sum(p.numel() for p in model.parameters())*4/1e9:.2f} GB (fp32 moments)") print(f" Total optimizer: {sum(p.numel() for p in model.parameters())*6/1e9:.2f} GB") print(f" A6000 VRAM: 48 GB") print(f" Available for activations: {48 - sum(p.numel() for p in model.parameters())*6/1e9:.1f} GB")