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.