spec100m / code /AGENTS.md
Akahsizrr's picture
squash to reclaim LFS quota
a8f07a3
|
Raw History Blame Contribute Delete
5.57 kB

spec100m β€” ~435M speculative-decoding model (v6 hybrid heads)

Goal

~435M param transformer + hybrid speculative heads, verified decoding (exact greedy match) AND accept-all mode. Max tok/s on RTX 3060 (Ampere, 12GB) / A100. English quality now matters β€” base is a trained continuation model.

Stack

  • Python 3.13, torch 2.6.0+cu124, CUDA 12.4, RTX 3060 (Ampere, 12GB; bf16 yes, fp8 no)
  • Remote: A100 80GB pod (gpu-744c7d98, frp.gpu.ai:10114, key gpuai_key, ~$1.19/hr)
  • tiktoken GPT-2 BPE (50257 vocab, pre-trained, no install)
  • SDPA (torch built-in fused attention) β€” no flash-attn (Windows install hell)
  • NEVER run heavy compute locally β€” all training/benchmarking on the pod.

Architecture (v6 β€” current: 434M + conditioned spec heads)

  • Base ~271M: d=1024, 8 layers, 16 q-heads, 1 kv-head (MQA), SwiGLU(inter=8192), RMSNorm, tied embeddings. Trained on 825M tokens (75M wikitext β†’ 250M mixture β†’ 500M fineweb-edu/cosmopedia/owt/wt103/finepdfs). Best loss 3.41.
  • v6 hybrid spec heads (~164M), replaces rank-1 medusa:
    • spec_cond_w [256, 1280, 256] β€” conditioned heads: head k predicts offset k+1 given [h_t ; compressed embs of previous group's 8 tokens]. Groups of G=8 share conditioning β†’ 32 sequential group steps at inference, teacher-forced (real tokens) during training.
    • spec_par_w [256, 1024, 256] β€” parallel tail heads, offsets 257..512.
    • spec_vocab [256, 50257] + spec_bias [V] β€” SHARED vocab projection: every head emits a real context-dependent distribution (the fix for r=1's 2-token argmax/argmin limitation).
    • spec_tok_proj [1024, 32] β€” compresses token emb for conditioning.
  • Verified decoding: draft K tokens -> ONE causal forward (causal_extend=True) -> commit longest prefix matching base argmax + 1 free correction token. Output == base greedy exactly. Cache rewinds to last accepted position. Correction token becomes the next round's pending (head-0 slot).
  • Accept-all mode: same heads, no verification β€” raw throughput mode.
  • KV compression (m=4) + FP8 + MoBA top-3/4096 + sparse-stride-4 (v5 stack, used by V5Engine for 20M context; BOTH verified and accept-all paths exist on V5Engine β€” verified rewinds the compressed cache at m-entry granularity and carries committed-but-uncompressed tokens in a small queue).

Training

  • Phase 1 (done): train.py base-only, 825M tokens cumulative. Checkpoints: /root/checkpoints/base_{75m,250m,500m}/{best,final}.pt (also pulled locally).
  • Phase 2: train_medusa.py --resume <base ckpt> --data mixture500m
    • freezes base, trains spec_* params only
    • teacher forcing: head k cond = real tokens of previous group
    • subsamples 48c+16p heads/step x 512 positions; per-head backward (shared-E graph would OOM β€” E is recomputed per head)
    • eval_streak() reports free-running draft-vs-base-greedy accept length

Commands

  • python config.py # print param count
  • python test_v6.py # v6 head shape + verified-decode smoke (pod)
  • python train_medusa.py --resume ckpt --steps 6000 # phase-2 head training
  • python bench_v6.py <ckpt> [tokens] # AR vs verified vs accept-all benchmark
  • python generate.py <ckpt> --verify # verified speculative generation
  • python generate.py <ckpt> # base-model sampling reference

Key files

  • config.py β€” Config (v5_500m preset = v6 heads), param counter
  • model.py β€” SpecModel: transformer + spec_cond_w/spec_par_w/spec_vocab + spec_draft() + causal_extend verify path + legacy r=1 heads
  • inference.py β€” InferenceEngine.speculative_verified() (exact) + speculative() (accept-all) + FastEngine + V5Engine (20M MoBA)
  • train_medusa.py β€” phase-2 head trainer
  • train.py β€” phase-1 base trainer (freezes medusa_*/spec_* params)
  • ngram.py β€” NGramModel (free tail extension, accept-all mode)

Measured (will fill in after v6 training)

  • v5 accept-all (rank-1, untrained heads): 600k tok/s w/ n-gram β€” garbage output
  • v6 target: verified ~300-800 tok/s (exact base-greedy quality), accept-all ~20-50k tok/s (head-quality output)

Known limitations / next steps

  • Compression params (w_kv_comp, w_z_comp, comp_bias) still untrained β€” only used in V5Engine 20M path; add consistency loss in a later phase.
  • Sequential conditioned draft is ~32 group steps β€” CUDA-graph it if the draft loop dominates round time.
  • Per-head backward loop in train_medusa.py is simple but has Python overhead; fine at ~2s/step.
  • torch.load weights_only warnings remain (checkpoints are ours β€” cosmetic).

Legacy results (v1-v5, rank-1 accept-all era β€” pre-training)

  • v3 short-ctx (RTX 3060, 228M, K=4096, r=1, CUDA graph): 1.40M tok/s accept-all
  • v4 MoBA prefill (A6000, 228M, 20M ctx): 971k tok/s simulated
  • v5 real chained prefill (A6000, 480M, 20M ctx, MoBA+compress+FP8): 124k tok/s
  • v5.1 generation + n-gram M=32768 (A6000): 600k tok/s accept-all
  • All legacy numbers are accept-all with random/untrained heads = garbage text; they measure throughput plumbing only, not usable output.
  • v5 20M-context stack (compressed KV + FP8 + MoBA + sparse) is preserved in V5Engine.forward_moba_v5 β€” works with r=1 checkpoints; needs spec_draft plumbing for v6 heads at long context (TODO).

Crash safety

  • Config.estimate_vram_gb() pre-checks before model allocation.
  • saturation.py refuses to build if est VRAM > 85% free.
  • A previous 5.8B-param medusa sweep crashed the PC via Windows TDR β€” keep ALL heavy runs on the pod, not local.