Skip to content

feat(speculative): add prompt-lookup decoding to mlxcel generate - #2074

Merged
inureyes merged 3 commits into
lablup:mainfrom
daemyung-lablup:feature/prompt-lookup-decoding
Oct 2, 2026
Merged

inureyes merged 3 commits into
lablup:mainfrom
daemyung-lablup:feature/prompt-lookup-decoding

Conversation

@daemyung-lablup

Copy link
Copy Markdown
Contributor

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 with KVCache::trim. Acceptance draws t ~ p and accepts iff t == 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).
  • Pipelining: rounds without a proposal submit the next one-token step from the still-lazy sample before the host reads it, as CxxGenerator does. 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_logits and resolve_kv_cache_layer_modes are extracted from CxxGenerator (no behavior change there) so the lookup path prefills and builds caches exactly as plain decoding does.
  • CLI: --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

  • The prompt goes through the plain-decode prefill routine, so caches and the first token match plain decoding, and a round without a proposal is the same one-token forward. With proposals disabled (--prompt-lookup-ngram-min 50), greedy output is byte-identical to plain decoding, pipelined or with MLXCEL_FORCE_SYNC=1.
  • --kv-cache-mode and the Boundary-V policy apply (caches come from model.make_caches()), --seed seeds the RNG, and the repetition-loop guard runs after every emitted token, including inside an accepted block.
  • Models whose state a cache trim cannot roll back are refused with the reason: supports_batching() == false (recurrent or hybrid layers, model-owned sliding caches) or supports_padded_prefill() == false. This currently excludes Qwen 3.5 (GatedDeltaNet) and Gemma 4.
  • The [Generated …] rate times the whole call, prefill included, the same as generate_standard; --profile prints 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_think prompts. 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-tokens 2/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").

Target Prompt Plain (tok/s) Prompt lookup Classic (best k)
Qwen3-1.7B edit_fn 172 331 (1.93x) 171 (0.99x, k=3)
Qwen3-1.7B edit_json 173 298 (1.73x) 170 (0.98x, k=3)
Qwen3-1.7B summary 172 195 (1.13x) 151 (0.88x, k=2)
Qwen3-1.7B write 184 185 (1.01x) 142 (0.77x, k=2)
Qwen3-4B edit_fn 79 136 (1.73x) 111 (1.40x, k=4)
Qwen3-4B edit_json 81 137 (1.70x) 112 (1.39x, k=3)
Qwen3-4B summary 77 116 (1.52x) 101 (1.32x, k=4)
Qwen3-4B write 88 83 (0.95x) 90 (1.03x, k=2)

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)

Configuration edit_fn edit_json summary write
Initial: ngram-min 1, full block every round, synchronous rounds 1.84x 1.82x – 0.74x
+ ngram-min 2 and the adaptive governor 1.80x 1.73x 1.12x 0.91x
+ pipelining plain rounds 1.89x 1.51x 1.22x 1.04x
+ tuned governor and pipelining threshold (this PR) 1.93x 1.73x 1.13x 1.01x
  • One-token n-gram matches were the prose regression: at ngram-min 1, free writing proposed 187 tokens and 8 landed.
  • The remaining 8% on free writing was host synchronization, not verify width: plain decoding forced synchronous (MLXCEL_FORCE_SYNC=1) drops from 192 to 169 tok/s, matching the lookup loop with every round plain (168).
  • Pipelining every plain round cost edit_json 0.2x, because a pipelined step cannot carry a proposal and JSON edits miss for a token or two at every object. The governor multiplier, miss streak and pipelining threshold were chosen from a 36-point grid over the four prompts; the chosen point is the simplest of the top cluster (geometric-mean speedups 1.36–1.38x).

Other measurements

  • --kv-cache-mode int8, Qwen3-1.7B, edit_json: 90 tok/s plain, 144 tok/s with lookup (1.6x).
  • Verify cost: an eight-token verify round costs about 2.6 one-token steps (13.6 ms vs 5.3 ms) on this machine, so a round must land about two proposals to break even. MLXCEL_QMV_WIDE=0 drops edit_fn from 320 to 272 tok/s, so the wide quantized kernel is already most of what keeps the verify affordable.

Greedy parity

Prompt Output vs plain decode (greedy)
edit_fn byte-identical
edit_json byte-identical
summary byte-identical
write diverges at char 309 of 566 ("avoid" vs "reduce")

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

  • CLI only (mlxcel generate -p). The server's SpeculativeDispatch has 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.
  • Hybrid and recurrent targets are refused. Qwen 3.5 could reuse the MTP path's rollback_speculative_cache; Gemma 4 is refused through supports_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.
  • Rounds that carry a proposal are still synchronous: the host reads the verify argmax before building the next round.
  • Measured on one M4 Pro with the Qwen3 family only. M5-class hardware selects different verify kernels and was not measured.

Found while measuring (not changed here)

  • Classic --draft-model CLI measurements are not comparable to plain decoding. The classic branch of run_generation_mode calls SpeculativeGenerator::generate once 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 like generate_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 with generation_stats_from_duration; this should be fixed in its own PR, alongside the MTP branch if it shares the pattern.
  • --seed does not make sampled output reproducible, on plain decoding too. Two runs of plain mlxcel generate --temp 0.8 --seed 7 on Qwen3-1.7B produced different outputs. Prompt lookup now calls seed_rng_if_needed like 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 as CxxGenerator.
  • cargo test -p mlxcel --bin mlxcel -- commands::generate — 89 passed.
  • cargo clippy on mlxcel-core and the mlxcel bin — no warnings in the changed files.
  • Live CLI on Qwen3-1.7B-4bit and Qwen3-4B-4bit: the benchmarks and parity checks above, -n 1 and -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).

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.
@inureyes inureyes self-assigned this Oct 2, 2026
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
@inureyes

inureyes commented Oct 2, 2026

Copy link
Copy Markdown
Member

Hardening review

Reviewed against the description (no linked issue). Pushed a1fa9e70 on a merge of current main (007f8981), which clears the two red checks inherited from main (fixed there by #2080). The algorithm checks out: the k - a trim, NgramIndex against find_draft, the in-flight invariants, and sampler-match being acceptance-optimal for a one-hot proposal.

Fixed

  • Gate: a ModelOwned sequence-state layout is refused directly. Gemma 4 and Llama 4 batch but own their caches, and were refused only through supports_padded_prefill, a different hazard (feat(core): shared snapshot serialization for KVCache / RotatingKVCache / ChunkedKVCache so Gemma 3, AFMoE and Llama 4 join exact-prefix prompt-cache reuse #1335); trimming their empty placeholder caches rolls nothing back.
  • Warmup: one forward per verify width before the timed run. The 16-token warmup compiled only widths its own proposals used, none on prose, so the timed run paid first-use kernel costs (Qwen3-1.7B email on GB10: 0.65x to 1.06x).
  • ngram-max at most 16 and max-draft at most 64 (index memory grows with the square of the range); incremental SamplerState on the penalty path; mirostat and adaptive-p refused.
  • CLI: tuning flags require --prompt-lookup; bad configuration or multimodal input is refused before the checkpoint is resolved.
  • Tests: a cache-dependent target must match plain decoding through accepted, rejected and pipelined rounds (six deliberate mutations of the trim and in-flight logic each fail them), plus gate, bounds, warmup and CLI tests.

Found, not changed here

GB10 (CUDA): release build of a1fa9e70, greedy, whole-call tok/s including prefill, median of 3 (spread mostly under 2%), -n 400 (story -n 800), speedup over plain:

Model edit_fn edit_json summary write story
Qwen3-1.7B-4bit 1.10x 1.37x 0.79x 1.06x 0.84x
Qwen3-4B-4bit 1.21x 1.55x 1.32x 0.93x 0.87x
Qwen3-8B-4bit 1.36x 1.81x 1.50x 0.97x 0.90x
Llama-3.1-8B-Instruct-4bit 1.51x 1.94x 0.92x 0.96x 0.92x

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

  • speculative::prompt_lookup 29 passed; CLI tests in the mlxcel bin 95 passed; clippy -D warnings and fmt clean.
  • Full CUDA gate (cargo test --workspace --profile test-fast --features cuda --no-fail-fast -- --test-threads=1): 11800 passed, 5 failed, the five tracked in fix: six lib tests fail under --features cuda on main #2087 and failing identically on main at 78af3e2d.
  • Real checkpoints: refusals on gemma-3-1b, gemma-4-e2b and qwen3.6-35b-a3b; sampled, --repetition-penalty, --kv-cache-mode int8 and --profile runs on Qwen3-1.7B.
  • Security: the only new input surface is CLI flags, now bounded; no unsafe code, file or network access.

Not run: Metal (no Apple machine on this host).

@inureyes
inureyes merged commit 9914b1a into lablup:main Oct 2, 2026
25 checks passed
@inureyes inureyes added status:done Completed type:enhancement New features, capabilities, or significant additions priority:medium Medium priority area:inference Generation, sampling, decoding (incl. speculative, DRY) area:cli Command-line interface / CLI flags labels Oct 2, 2026
inureyes added a commit that referenced this pull request Oct 2, 2026
#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
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

area:cli Command-line interface / CLI flags area:inference Generation, sampling, decoding (incl. speculative, DRY) priority:medium Medium priority status:done Completed type:enhancement New features, capabilities, or significant additions

Projects

None yet

Development

Successfully merging this pull request may close these issues.

2 participants