# 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 --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 [tokens]` # AR vs verified vs accept-all benchmark
- `python generate.py --verify` # verified speculative generation
- `python generate.py ` # 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.