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

CONTINUE β€” agent handoff for spec100m

Read secrets/ACCESS.md first for pod/API credentials.

Project goal

Build a fast speculative-decoding LM (435M params) for RTX 3060/A100-class hardware. A trained English base transformer (271M) + speculative heads (~164M). Two decode modes:

  • verified: output identical to base greedy, ~3-6x speedup (the real one)
  • accept-all: no verification, raw throughput (the fun number)

User's intent: "INSANE SPEED", heads that offer real multi-token predictions, maximize accepted tokens per base-model forward pass.

Current state (as of this commit)

  • Base model: TRAINED. 825M tokens cumulative (75M wikitext-103 β†’ 250M mixture β†’ 500M fineweb-edu/cosmopedia-v2/owt/wt103/finepdfs). Best loss 3.41. Grammatical multi-paragraph English, still drifts after ~50 tokens. Checkpoints on pod at /root/checkpoints/base_*/.
  • v6 spec heads: TRAINING NOW on the pod β€” see "What's running" below.
  • Pod gpu-744c7d98 is RUNNING and billing ~$1.19/hr. Keep an eye on cost.
  • 2026-09-11 fixes (commits 38da0c8, d1030c0) β€” verified decode is now correct (was emitting every correction token twice + RoPE ignored start_pos so all decoded tokens sat at position 0). The 20M MoBA-v5 path now produces output faithful to dense attention (cos=1.000): identity-init compressor (init_compressor_from_attn), RoPE on pooled keys/queries, dense causal attention within the current chunk, sparse compressed attention for history only. Prefill measured >100K tok/s at 20M ctx with prefill_chunk=4096-8192 (shared GPU). Generation uses CUDA-graphed spec_draft. See git log for details.

v6 architecture (the important part)

Old rank-1 Medusa heads could only emit 2 fixed tokens per head β€” replaced with a hybrid design informed by EAGLE-3/Hydra research:

  • spec_cond_w [256, 1280, 256] β€” 256 conditioned heads in groups of G=8. Head k predicts token at offset k+1 given h_t + compressed embeddings of the PREVIOUS group's 8 tokens. At inference: 32 sequential micro-steps of [B,8] matmuls. In training: teacher forcing on ground-truth tokens.
  • spec_par_w [256, 1024, 256] β€” 256 parallel tail heads (offsets 257..512), unconditioned, all fired in one einsum.
  • spec_vocab [256, 50257] + spec_bias [V] β€” SHARED vocab projection: the key fix, every head gets a real context-dependent distribution.
  • spec_tok_proj [1024, 32] β€” compresses token embeddings for conditioning.

Verified loop (InferenceEngine.speculative_verified in inference.py): draft K tokens β†’ one causal forward (causal_extend=True) over the chunk β†’ commit longest prefix matching base argmax + 1 free correction token β†’ the correction becomes next round's pending (occupies head-0 slot, heads 1..K shift). Cache write position rewinds past rejected tokens. Output is bit-identical to AR greedy β€” bench_v6.py asserts this.

What's running on the pod RIGHT NOW

nohup python3 train_medusa.py --resume /root/checkpoints/base_500m/final.pt \
    --data mixture500m --steps 6000 --batch 16 --seq 1024 \
    --pos_per_head 512 --eval_every 250 --checkpoint_every 500 \
    --out /root/checkpoints/spec_v6
  • PID was 3941, log: /root/train_spec_v6.log
  • ~1.8-2s/step β†’ 6000 steps β‰ˆ 3h
  • Watch for [eval] accept streak ~N lines β€” that's the key metric (free-running draft vs base-greedy match length). Random heads β‰ˆ 0; decent heads target streak ~4-10.
  • Checkpoints land in /root/checkpoints/spec_v6/ (best.pt + step_N.pt)

Next steps after training finishes

  1. ssh in (see ACCESS.md), check tail -30 /root/train_spec_v6.log
  2. Run benchmark: cd /root/spec100m && python3 bench_v6.py /root/checkpoints/spec_v6/best.pt β€” prints AR vs verified vs accept-all tok/s + exact-match flag. Optional 3rd arg = verify_k (e.g. 64) caps draft/verify length per round.
  3. If accept streak is good (~4+): pull spec_v6/best.pt to local checkpoints/spec_v6/, update README/AGENTS with numbers
  4. If streak is weak: tune (more steps, bigger pos_per_head, or try scheduled sampling β€” heads currently teacher-forced only)
  5. 20M-ctx path (V5Engine): call model.init_compressor_from_attn() after load (compressor params are still untrained β€” identity init makes compressed KV ~= mean-pooled real K). Prefill with prefill_chunk=4096+ hits >100K tok/s at 20M. Generation uses CUDA-graphed draft; n-gram extension (generate_with_ngram) stacks free tokens on top.
  6. DONE: V5Engine.speculative_verified (verified decode on the 20M path, ~35-50 tok/s β€” commits ~1/round at current streak, needs better heads). Follow-ups: scheduled-sampling run (--ss-prob in train_medusa.py) to fix the teacher-forcing gap; train compressor for a consistency loss; pad- static MoBA step + CUDA graph to push prefill toward 300K+.

Repo map

  • config.py β€” Config; v5_500m() = current v6 preset
  • model.py β€” SpecModel + spec heads + spec_draft + causal_extend
  • inference.py β€” InferenceEngine (verified + accept-all), FastEngine, V5Engine
  • train_medusa.py β€” phase-2 head trainer (RUNNING NOW)
  • train.py β€” phase-1 base trainer
  • bench_v6.py β€” the benchmark that matters
  • test_v6.py β€” shape/logic smoke test
  • generate.py β€” sampling; --verify flag = verified spec decode
  • data_pipeline.py β€” HF dataset streaming/caching (mixture500m recipe)
  • ngram.py β€” n-gram free-token extension (accept-all only)
  • provision.py β€” gpu.ai pod provisioning (has API token inline)
  • remote_*.sh β€” helpers executed ON the pod
  • secrets/ β€” ssh key + all tokens
  • AGENTS.md β€” deeper project notes incl. legacy v3-v5 numbers

Local machine note

The user's RTX 3060 PC must NOT run heavy training/inference β€” a previous 5.8B-param sweep crashed it via Windows TDR. All compute on the pod. Local scp/pull of checkpoints is fine.

Workflow tips

  • PowerShell chokes on && and nested quotes β€” use cmd /c 'ssh ... "..."' single-quoted style (see history) or the remote_*.sh helpers.
  • Pod-side long runs: nohup ... > /root/log 2>&1 & then poll with tail.
  • Old checkpoints have medusa_* (r=1) keys, new arch has spec_* β€” always load_state_dict(..., strict=False) and print missing/unexpected.