feat(speculative): add prompt-lookup decoding to mlxcel generate - #2074
Conversation
Replies that edit or reformat the prompt (fix typos, add a JSON field, rename an identifier) repeat long runs of the prompt word for word, and speculative decoding can emit those runs several tokens per target forward without any draft model. mlxcel had no drafter-free speculation: the classic, MTP, and DFlash paths all need a second model or head. Add PromptLookupGenerator in mlxcel-core and wire it to `mlxcel generate --prompt-lookup`. Each round matches the tail of the sequence against earlier context and verifies up to seven proposed tokens in one target forward, trimming the rejected tail off the KV caches. - An incremental n-gram index answers each lookup with one hash probe per n-gram length instead of a scan over the whole context; a test checks it against the reference scan at every context length. - DraftGovernor shortens the block while few tokens land and pauses proposals after three misses in a row, doubling the pause up to 32 rounds. The default minimum n-gram is 2, since one-token matches slowed prose by a quarter. - Rounds without a proposal are pipelined the way plain decoding is: the next step is submitted from the lazy sample before the host reads it. Waiting on each token cost 12% of decode time on an M4 Pro. Pipelining starts after two plain rounds so short gaps inside edits keep their proposals. - The prompt is prefilled through the plain-decode routine, extracted as prefill_prompt_last_logits, so the caches and the first token match plain decoding exactly. KV caches come from make_caches with the configured --kv-cache-mode and Boundary-V policy, --seed is honored, and the repetition-loop guard runs after every emitted token. - Models whose state a cache trim cannot roll back are refused with the reason (recurrent or hybrid layers, model-owned sliding caches, or state that opts out of padded prefill). Multimodal prompts and interactive chat are refused. - The CLI rate line times the whole call, prefill included, as the plain path does; --profile prints the measured prefill/decode split. The acceptance summary goes to stderr. Measured on an M4 Pro, greedy, 400-token cap, whole-call timing: Qwen3-1.7B-4bit function edit 1.93x, JSON edit 1.73x, summary 1.13x, free writing 1.01x over plain decode; Qwen3-4B-4bit 1.73x, 1.70x, 1.52x, 0.95x. A Qwen3-0.6B classic drafter reached at most 0.99x on 1.7B and 1.40x on 4B on the same prompts. Edit and summary replies are byte-identical to plain decoding at greedy; a multi-token verify can still flip a near-tie, the same jitter class as the draft-model path. Validated with cargo test (mlxcel-core 179 including new index, governor, and loop-guard tests; mlxcel commands::generate 89) and clippy with no new warnings.
Prompt lookup rolls rejected proposals back with `KVCache::trim`, so a target must keep its whole sequence state in the external caches. The gate inferred that from `supports_batching` and `supports_padded_prefill`; Gemma 4 and Llama 4 batch yet own their caches, and were refused only through the padded-prefill flag, which answers a different question (lablup#1335). The gate now refuses a `ModelOwned` layout directly, and the generator asserts non-empty external caches (`can_trim_prompt_cache` is vacuously true on an empty slice). - Warm every verify width before the timed CLI run. The warmup generation reached only the widths its own proposals used, so on prose the timed run paid first-use kernel costs mid-decode: a 72-token Qwen3-1.7B reply on GB10 went from 0.65x to 1.07x of plain decoding. - Bound `ngram-max` at 16 and `max-draft` at 64; the n-gram index grows with the square of the range times the context. - Sample history-reading configs through the incremental `SamplerState`; refuse mirostat and adaptive-p, whose state the loop does not carry. - Tuning flags require `--prompt-lookup`; a bad configuration or multimodal input is refused before the checkpoint is resolved. - Rollback tests drive a cache-dependent target through accepted, rejected and pipelined rounds; six deliberate mutations of the trim and in-flight logic each fail them. The penalty case uses a sequential reference because `CxxGenerator` samples against a history one token stale (lablup#2090). - Document the acceptance rule in `docs/speculative-acceptance.md`. Refs lablup#2074
Hardening reviewReviewed against the description (no linked issue). Pushed Fixed
Found, not changed here
GB10 (CUDA): release build of
Edit and summary replies are byte-identical to plain decoding in every row; the 8B emails and all stories diverge mid-reply (the near-tie flip the description documents). Validation
Not run: Metal (no Apple machine on this host). |
#2092) ## Summary Prompt lookup slowed prose decoding 3 to 21% on GB10 (#2091). This adds `DraftPolicy::Gated`, the new default on CUDA builds, and keeps PR #2074's governor unchanged as `DraftPolicy::Graded` for Metal and ROCm. On GB10, stories go from 0.83x-0.93x to 0.98x-1.00x of plain decoding, edit and summary gains are kept or raised, and greedy parity is preserved. One target is not met: Qwen3-8B's email row stays at 0.97x (graded 0.98x in the same run). ## Why A synchronous verify forward on GB10 (affine 4-bit, MLX pin `81ba1c6a`, release build) costs, in pipelined steps for Qwen3-1.7B / Qwen3-8B: width 2 at 1.37 / 1.15, width 3 at 1.76 / 1.59, width 8 at 3.91 / 3.30, with widths 5 to 7 close to or above width 8 (`qmv`'s 8-row multirow instantiation, then `qmm_sm80` from 8 rows). A drafted round also drains the pipeline. Graded, tuned on an M4 Pro, re-probed prose with blocks up to width 8 every 4 to 32 rounds. ## What changed - `DraftPolicy::Gated`: blocks are narrow (2 proposals) or full (`max_draft`), widening after a narrow block lands whole. A reply starts narrow and on probation, and probation ends only on a round that lands a whole narrow block. A pause has no timer: while paused the loop keeps looking up and checks each proposal against the tokens decoding emits next (`ShadowProbe`); drafting resumes, on probation, only after two proposed tokens in a row came true. Paused rounds pipeline at once. - `DraftPolicy::Graded` makes exactly PR #2074's decisions and stays the default off CUDA: no Apple Silicon machine was reachable, so Metal keeps the rules it was tuned with. - `--prompt-lookup-policy auto|graded|gated` for A/B on any backend; `shadow_confirmations` on the `[Prompt lookup]` line and in the decode trace. - `examples/verify_width_cost.rs` measures the per-width costs above. The verify path and the in-flight transition are unchanged. ## Results (GB10) Final matrix on `34c627d7` (the tree this PR ships), greedy, median of 3, plain/graded/gated interleaved; full tables, counters, parity and method in `docs/benchmark_results/prompt-lookup-governor-gb10-2026-10-02.md`. | Model | write graded / gated | story graded / gated | worst edit/summary change vs graded | |---|---|---|---| | Qwen3-1.7B | 1.07x / 1.02x | 0.83x / 0.98x | +1% (edit_json), summary 0.82x to 0.99x | | Qwen3-4B | 0.94x / 0.98x | 0.87x / 0.98x | -4% (edit_json 1.71x to 1.64x) | | Qwen3-8B | 0.98x / 0.97x | 0.91x / 0.98x | +1% (edit_json) | | Llama-3.1-8B | 0.96x / 0.99x | 0.93x / 1.00x | -2% (edit_json), summary 0.94x to 1.02x | Every reply graded keeps byte-identical to plain decoding stays identical under gated. A third matrix with a stricter full-block threshold left Qwen3-8B's email at 0.97x and broke parity on two email rows, so it was reverted (`d678c940`); the record explains the remaining cost and the round-loop change that would address it. ## Changes during review No CRITICAL or HIGH findings. Applied: immediate pipelining while paused, an exact probe-alignment test, per-policy parity totals, two-token shadow lookups, budget-vs-width wording, and a trim guard in the width probe. ## Validation - `speculative::prompt_lookup`: 39 passed; deliberate mutations of the gated rules, probe alignment and probation each fail them. - CLI tests in the `mlxcel` bin: 96 passed; `dead_doc_pointers` passed; clippy `-D warnings` (lib, tests, bins, the new example) and fmt clean. - Not run: Metal (no Apple Silicon host). Refs #2091
Summary
Adds drafter-free speculative decoding to
mlxcel generate. With--prompt-lookup, each decode round matches the tail of the sequence against earlier context (prompt plus everything emitted so far), proposes the tokens that followed that earlier occurrence, and verifies up to seven of them in one target forward. Replies that edit or reformat the prompt repeat it word for word for most of their length, so they come out several tokens per weight read with no second model loaded.On an M4 Pro with Qwen3-1.7B-4bit, edit replies run 1.73–1.93x faster than plain decoding, summaries 1.13x, and free writing breaks even (1.01x). On Qwen3-4B-4bit the same prompts run 1.73x, 1.70x, 1.52x and 0.95x, ahead of a Qwen3-0.6B classic drafter everywhere except free writing.
Motivation
mlxcel already speculates through the classic draft-model path, MTP, and DFlash, but every one of them needs a second model or head, and MTP and DFlash exist only for specific families (Gemma 4, Qwen 3.5 line, Inkling, GLM-4.7-Flash, Muse Glimmer). Pure-attention models such as Qwen3, Llama and Mistral have no MTP or DFlash drafter, and a small classic drafter does not pay off on small targets (below 1.0x on Qwen3-1.7B, see the benchmarks). Prompt lookup needs nothing but the context, and editing replies are a common assistant workload where it does best.
What changed
PromptLookupGenerator(mlxcel-core/src/speculative/prompt_lookup.rs): the round loop. Each round forwards[current, d_0 … d_{k-1}], keeps the longest agreeing prefix, emits the target's token at the first disagreement (or the bonus token), and trims the rejected tail withKVCache::trim. Acceptance drawst ~ pand accepts ifft == d, which is modified rejection sampling for a one-hot proposal, so sampled output keeps the target's distribution.NgramIndex: an incremental index from each n-gram to its latest occurrence that has a follower. A lookup is one hash probe per n-gram length instead of a scan over the whole context, and the prompt is indexed while the GPU runs the prefill.DraftGovernor: caps the block at twice the recent accepted average plus two, and pauses proposals after three drafted rounds in a row land nothing (pause 4 rounds, doubling up to 32, reset on the next landing).CxxGeneratordoes. Pipelining starts after two plain rounds in a row, because a pipelined step cannot carry a proposal and short misses inside edits would otherwise lose one.prefill_prompt_last_logitsandresolve_kv_cache_layer_modesare extracted fromCxxGenerator(no behavior change there) so the lookup path prefills and builds caches exactly as plain decoding does.--prompt-lookup,--prompt-lookup-ngram-max(default 3),--prompt-lookup-ngram-min(default 2),--prompt-lookup-max-draft(default 7),--prompt-lookup-no-adaptive. Conflicts with--draft-model; refused for interactive chat, multimodal prompts, and pipeline parallelism.Behavior and parity
--prompt-lookup-ngram-min 50), greedy output is byte-identical to plain decoding, pipelined or withMLXCEL_FORCE_SYNC=1.--kv-cache-modeand the Boundary-V policy apply (caches come frommodel.make_caches()),--seedseeds the RNG, and the repetition-loop guard runs after every emitted token, including inside an accepted block.supports_batching() == false(recurrent or hybrid layers, model-owned sliding caches) orsupports_padded_prefill() == false. This currently excludes Qwen 3.5 (GatedDeltaNet) and Gemma 4.[Generated …]rate times the whole call, prefill included, the same asgenerate_standard;--profileprints the generator's measured prefill/decode split. The acceptance summary line goes to stderr.Benchmarks
M4 Pro (48 GB), macOS 26, release build with
metal,accelerate, greedy (--temp 0),-n 400, Qwen3/no_thinkprompts. Rate is generated tokens over the whole call including prefill, after the in-process warmup, median of 2–3 runs; run-to-run spread was within ±1%.Prompts: edit_fn renames a variable in a 20-line Python function; edit_json adds a field to each object of a 5-object JSON array; summary summarizes meeting notes reusing their wording; write is a short email with no source text.
Plain decode vs prompt lookup vs classic drafter
Classic uses Qwen3-0.6B-4bit as the drafter; the best of
--num-draft-tokens2/3/4/6 is shown per prompt. Classic was measured with a temporary warmup and whole-call timing patch (not in this PR, see "Found while measuring").The two are complementary: lookup wins wherever the reply copies its source, a drafter wins on free writing where lookup has nothing to propose. MTP and DFlash could not be compared on the same target, since their targets (Qwen 3.5, Gemma 4) are the families prompt lookup refuses today.
How each change moved the result (Qwen3-1.7B, speedup over plain)
MLXCEL_FORCE_SYNC=1) drops from 192 to 169 tok/s, matching the lookup loop with every round plain (168).Other measurements
--kv-cache-mode int8, Qwen3-1.7B, edit_json: 90 tok/s plain, 144 tok/s with lookup (1.6x).MLXCEL_QMV_WIDE=0drops edit_fn from 320 to 272 tok/s, so the wide quantized kernel is already most of what keeps the verify affordable.Greedy parity
The write divergence follows the single drafted round in that reply: the eight-token verify forward rounds differently from the one-token path and flips a near-tie. The existing draft-model path shows the same class of divergence on this prompt (even at 100% acceptance), and the module docs describe it; unlike MTP, nothing here gates on an exactness probe.
Known limitations
mlxcel generate -p). The server'sSpeculativeDispatchhas no prompt-lookup variant, and--spec-type ngram-*is still rejected at startup; wiring it (including B>1 per-row rollback) is follow-up work.rollback_speculative_cache; Gemma 4 is refused throughsupports_padded_prefill() == false, which is set for performance reasons, so it may be over-conservative, but its sliding caches have not been validated under rollback.Found while measuring (not changed here)
--draft-modelCLI measurements are not comparable to plain decoding. The classic branch ofrun_generation_modecallsSpeculativeGenerator::generateonce with no warmup, so the first-call Metal kernel compilation lands inside the timed run, and it prints the generator's own stats instead of whole-call timing likegenerate_standard. The rate it prints is therefore on a different basis from the plain path's. The Classic numbers above come from a local patch that added a 16-token warmup and wrapped the call withgeneration_stats_from_duration; this should be fixed in its own PR, alongside the MTP branch if it shares the pattern.--seeddoes not make sampled output reproducible, on plain decoding too. Two runs of plainmlxcel generate --temp 0.8 --seed 7on Qwen3-1.7B produced different outputs. Prompt lookup now callsseed_rng_if_neededlike plain decoding, so it inherits the same behavior.Test plan
cargo test -p mlxcel-core --lib -- prompt_lookup generate:: speculative::— 179 passed, including new tests for index/scan equivalence (single-token and block growth), governor shortening, pause, backoff and cap, the ngram-min default, and loop-guard stopping inside a verified block with the same tokens asCxxGenerator.cargo test -p mlxcel --bin mlxcel -- commands::generate— 89 passed.cargo clippyonmlxcel-coreand themlxcelbin — no warnings in the changed files.-n 1and-n 2,--temp 0.7,--repetition-penalty 1.1,--kv-cache-mode int8,--profile, and the refusals (conflict with--draft-model, interactive chat,--prompt-lookup-ngram-min 0, Gemma 3 through the padded-prefill gate).