Skip to content

feat(server): route production batched paged decode through the fused v2 kernel - #988

Merged
inureyes merged 10 commits into
mainfrom
feature/issue-899-production-paged-v2-dispatch
Aug 1, 2026
Merged

inureyes merged 10 commits into
mainfrom
feature/issue-899-production-paged-v2-dispatch

Conversation

@inureyes

@inureyes inureyes commented Jul 31, 2026 •

Copy link
Copy Markdown
Member

Summary

Makes the fused paged-attention decode v2 kernel from #898 the path every batched paged decode request takes, replacing the per-sequence gather_visible + dense SDPA loop that ADR 0001 measured as the paged decode hot path. The switch is gated by a measured token floor, keeps the gather path as the fallback and the parity reference, and preserves MLXCEL_PAGED_ATTENTION_NATIVE=0 as an end-to-end kill switch.

The dispatch threshold, and why it is two-regime

#898's M1 Ultra sweep (docs/benchmark_results/paged-decode-v2-m1ultra-2026-07-31.md, medians of three repetitions) has exactly one cell where v2 loses to gather: batch 1 at 1024 tokens of context, 0.91x median and below parity in all three repetitions. The plan degenerates to two pages per chunk there and the merge pass plus the workspace round trip costs more than the whole attention does. So the switch cannot be unconditional.

The first cut of this PR summed the launch's visible tokens and required 4096. The production benchmark disproved that formulation: it declined every scenario in the matrix and measured gather against gather. The loss is a property of batch 1, not of total tokens. Read the ctx-1024 column: 0.91x at batch 1, 1.41x at batch 4, 1.47x at batch 8, same per-request context and opposite outcome, because a batched launch spreads the same chunk count over more requests and amortizes the merge over more useful work. A total-token floor separates those cells only by accident and with no margin, and the benchmark's nominal 1K scenario delivered 956 tokens per request (4 * 956 = 3824), so it was declined despite being the same shape as the measured 1.41x win.

The floor is now stated the way the measurements are:

launch floor evidence
one request MIN_SINGLE_REQUEST_KV_TOKENS = 4096 visible tokens 1024 loses (0.91x), 4096 wins (1.08x median, 1.02x-1.21x across repetitions)
more than one MIN_BATCHED_KV_TOKENS_PER_REQUEST = 512 per request 1024 per request wins at batch 4 (1.41x) and batch 8 (1.47x)

512 rather than 1024 for the batched case is deliberate: 1024 is the lowest measured batched context and it wins comfortably, so a floor sitting on top of a measured point declines real workloads that land just under it for no measured reason. Batch 2 and 3 are interpolated, not measured; the trend across batch 1, 4 and 8 is monotone and the mechanism above explains why, so the interpolation is stated in the module docs rather than hidden. MLXCEL_PAGED_V2_MIN_KV_TOKENS and MLXCEL_PAGED_V2_MIN_KV_TOKENS_PER_REQUEST move the two floors; MLXCEL_PAGED_ATTENTION_NATIVE=1 bypasses them entirely, which is the supported way to benchmark the declined corner.

Both decode paths are wired, not just the batched one

dispatch_sync_decode routes a batch of one to decode_single_step, which calls the model's single-sequence Attention::forward, not forward_split_attention. Wiring only the batched path left the two single-sequence scenarios in the issue's benchmark matrix unreachable by construction, and with staggered prefills it also left most of a batched scenario on gather. Both forward_split_attention (whole batch) and Attention::forward (batch of one) now take the same entry point.

Which kernel actually ran is visible at info

The dispatch decision is a returned PagedDecodeOutcome carrying the numbers behind it, and the caller announces the first occurrence of each distinct outcome kind at info, with no RUST_LOG needed. Per kind rather than one global one-shot: a single flag reports only whatever happened first, typically a warmup request, so a later permanent decline for a different reason never surfaces. That is how the first production sweep ran entirely on the gather path without a single line saying so. scripts/benchmark_paged_decode_production.sh now exits non-zero if the after arm never logged a fused launch or the before arm did, so a null sweep fails instead of reporting.

Verified on mlxcel-server (Llama-3.2-1B-Instruct-4bit, --parallel 4 --ctx-size 131072, no RUST_LOG): 4 concurrent clients at a nominal 1K prompt log fused v2 launch (batch 4, 3828 visible KV tokens, 16 chunks, merge on); a single client at 16K logs fused v2 launch (batch 1, 13845 visible KV tokens, ...); the same binary under MLXCEL_PAGED_ATTENTION_NATIVE=0 logs gather: pinned by MLXCEL_PAGED_ATTENTION_NATIVE and never a fused launch.

The thing that was actually blocking this issue

PagedBlockPool allocated physical rows in fixed 32-block slabs (#235) and both fused decode kernels read one contiguous pool buffer per side, so any layer past 32 rows was multi-slab and declined. At block_size 32 that capped the fused path at 1024 tokens across the entire batch, which is below the token floor above: v2 would never have run in production and this issue would have been a no-op.

The slab size is now a per-pool value (PagedBlockPool::set_slab_blocks) defaulting to the historical 32, so nothing changes for any caller that does not set it. The server sizes it from the configuration it already reserved KV memory for: ceil(per_slot_ctx / block_size) * batch, floored at the old default and clamped by the per-layer share of the paged block budget (memory_estimate::resolve_paged_slab_blocks). MLXCEL_PAGED_SLAB_BLOCKS overrides it, 0 pins the old default.

The cost is that a layer's first write allocates its whole slab instead of growing into it, front-loading exactly the KV bytes the startup memory estimate already reports for that batch and context. Growth past the slab still appends without copying (#235's property is preserved), so an undersized slab degrades to the pre-#899 gather path rather than failing. This is the main thing to review; the alternative considered and rejected was geometric single-slab growth, which restores the pre-#235 realloc-and-copy behaviour and the ladder of orphaned buffer sizes in the MLX cache that #235 exists to avoid.

Operational consequence worth flagging: serving contexts longer than --ctx-size implies makes layers outgrow the slab and silently returns them to the gather path. The server logs the resolved value at startup (Paged KV slab size: N blocks per layer) and docs/CONTINUOUS_BATCHING.md says this plainly.

What changed

  • src/lib/mlxcel-core/src/cache/paged_batch_decode.rs (new): paged_batch_decode_attention, the whole-batch entry point. All-or-nothing contract: it either declines before writing anything (caller's unchanged per-sequence loop runs against an untouched pool) or it owns the step, including running the gather fallback itself when v2 is not the right choice. That is what keeps a fallback from double-writing the pool.
  • src/lib/mlxcel-core/src/paged_v2/dispatch.rs (new): the pure token-floor selector and its env override. layers::resolve_paged_v2_dispatch layers MLXCEL_PAGED_ATTENTION_NATIVE on top, checked before the selector so the kill switch cannot be out-thought by a shape the selector likes.
  • src/lib/mlxcel-core/src/paged_v2/plan_cache.rs (new): per-layer CSR page table and per-batch chunk plan caching. Building a page table costs one hash lookup per visible page (~131k per decode step for a 32-layer model at batch 4 and 32K context); refreshing the per-request scalars costs four writes per request. The expensive half is kept, the cheap half is recomputed, never incremented, because a delta would be one missed invalidation away from feeding the kernel a length that outruns its pages, which is an out-of-bounds read rather than a wrong answer.
  • src/lib/mlxcel-core/src/cache/paged.rs: PagedBlockPool::paged_decode_batched (the production entry), block_epoch (the cache invalidation counter), slab_blocks and set_slab_blocks.
  • src/models/llama3.rs, src/models/qwen3.rs: forward_split_attention calls the new entry point before its existing dense-compat paged block. Between them these two attention modules back Llama, Mistral, Qwen2 / Qwen2.5, Qwen3, Helium, and every VLM whose text backbone is one of those, which is the whole set of families that get pool-backed caches.
  • src/execution/memory_estimate.rs: resolve_paged_slab_blocks and the v2 workspace reserve charged to the KV budget.
  • src/server/model_worker.rs, src/server/batch/scheduler.rs, src/lib/mlxcel-core/src/cache.rs: plumbing the resolved slab size to the lazily created pool.
  • Docs: ADR 0001 marked superseded in part with a closing addendum, docs/CONTINUOUS_BATCHING.md gains an operator-facing section.

Plan caching and invalidation

Reuse of a cached page table is gated on a per-request page-range fingerprint (blocks, first_page, last_page, visible) plus PagedBlockPool::block_epoch, a counter bumped by every mutation that can move a block: acquire, release, row assignment, block restore, sequence restore, sequence release, and token trim. It is deliberately not bumped by a token append that fits in the current block, which is the decode hot path.

The epoch is what makes the issue's named invalidation events fall out without any scheduler plumbing, because each of them already goes through a pool method that moves blocks: admission (the sequence's first write_prefill acquires), eviction and preemption and finish (release_sequence releases), page-boundary crossing (append_tokens acquires), prompt-cache adopt and detach (restore_sequence / row forget). This is a deviation from the issue's wording, which asks for the cache to live on the active batch with scheduler-side invalidation: forwarding events the pool observes first hand would add a coupling that can be forgotten at a new call site, whereas the epoch cannot be. The epoch is global rather than per-layer, so a boundary crossing costs two fully rebuilt steps rather than one; at 32-token pages that is 2 steps in 32, about 6% of the uncached cost.

Sliding-window CSR ranges

v2 applies no windowing mask of its own, so the window has to be the CSR range. build_paged_csr_view (from #898) already emits only the pages from logical_start / page_size onward and sets first_page_offset to logical_start % page_size, so request r's token i resolves to absolute position logical_start + i and the retired prefix is never addressed. That is the same [logical_start, len) window gather_visible slices, which is why the two agree. Two properties are documented in the module because both are silent if violated: the window must be contiguous (a future paged attention-sink retention, [0, keep) ++ [start, len), is not expressible in one CSR range and must not be routed here; today it cannot be, because trim_front_keep_sink returns 0 for pool-backed caches), and RoPE is applied upstream from the scheduler's rope_offsets, so the view's own rope_offsets is not consumed here.

Workspace budgeting

The workspace is the partial kernel's (partial_v, lse) output pair, PagedDecodePlan::workspace_bytes(). resolve_paged_block_budget now subtracts a reserve before converting the KV byte budget into blocks, so admission reserves blocks it can actually back. The reserve bounds num_chunks from the plan's own search invariant (num_chunks * ctas_per_chunk lands within a factor of two of the device CTA target, and ctas_per_chunk >= Hkv), which makes the dominant term independent of the head counts: 2 * target_ctas * n_rep * (head_dim + 1) * 4, times two concurrent launches for the decode lookahead pipeline. n_rep is bounded at 16 rather than read from config, which over-reserves by a few times in the safe direction. It lands in the low tens of megabytes, well under a tenth of a percent of a serving KV budget; it is charged because the accounting should be truthful, not because it is large.

Test plan

  • cargo test --release -p mlxcel-core --lib --features metal,accelerate paged_v2:: — 55 passed (dispatch policy, outcome reporting, plan cache reuse and invalidation).
  • cargo test --release -p mlxcel-core --lib --features metal,accelerate cache::paged — 22 of them for the batched entry point, including the batch-1 fused path, the 4x956 shape the old floor declined, and the multi-slab decline naming its own fix.
  • cargo test --release -p mlxcel-core --lib --features metal,accelerate cache:: — 468 passed.
  • cargo test --release -p mlxcel-core --lib --features metal,accelerate layers:: — 64 passed; ffi_tests:: — 93 passed.
  • cargo test --release --lib --features metal,accelerate memory_estimate:: — 35 passed.
  • cargo test --release --lib --features metal,accelerate server::batch::scheduler::tests::paged_decode_storage and ::auto_decode_storage — 2 passed.
  • cargo test --release --test paged_scheduler_parity --features metal,accelerate -- --ignored — 4 passed; paged_real_model_parity 1, paged_prefix_share_parity 7, paged_kv_serialize_parity 2, paged_handoff_parity (with test-utils) 2.
  • Greedy token parity across three families: cargo test --release --test paged_decode_v2_greedy_parity --features metal,accelerate -- --ignored --nocapture at the default batch 4, 8192-token prompts with mixed lengths, 32 decode steps. Qwen3-0.6B, Llama-3.2-1B, and Qwen2.5-0.5B all produce token streams identical to the gather baseline on all four sequences. The test asserts the v2 arm actually built a decode plan, so a silent fallback cannot pass as parity.
  • cargo clippy -p mlxcel-core --lib --tests --features metal,accelerate -- -D warnings and cargo clippy --lib --tests --features metal,accelerate -- -D warnings clean.
  • cargo fmt --all -- --check clean.
  • Server-level benchmark, not run here. scripts/benchmark_paged_decode_production.sh drives the issue's mandatory matrix (4 concurrent clients at 1K / 4K / 16K, plus single-sequence 16K and 32K) against a freshly started server per arm, with MLXCEL_PAGED_ATTENTION_NATIVE=0 as the before arm. It must run serialized on a quiet host; a loaded machine produces misleading numbers for a bandwidth-bound path.

Left out, deliberately

Logit-softcap families fall back to gather, as the issue scopes them to a follow-up: neither the v2 kernel nor this function's fallback threads a soft-cap term. Speculative and MTP verify steps stay on their current path by construction rather than by a flag, because they arrive with more than one query token and the decline is on seq_len != 1. Multi-slab layers are declined rather than served; lifting that needs either a slab-aware kernel launch (one partial launch per slab, merged through the existing variable-length merge, which the merge kernel already supports) or a different pool growth strategy, and both are larger than this issue. CUDA is still unvalidated: the v2 CUDA JIT bodies from #898 have never been compiled or run, and DEFAULT_TARGET_CTAS still needs a GB10 pass.

Where I was unsure

The slab sizing policy is the judgement call with the widest blast radius. ceil(per_slot_ctx / block_size) * batch is honest (it is the KV the config already reserved) but it is eager, and I could not measure the allocation-timing effect. A reviewer who prefers laziness over fused-path eligibility should look here first; MLXCEL_PAGED_SLAB_BLOCKS=0 restores the old behaviour exactly.

The 1.08x cell at batch 1 / 4096 sits on the floor with the least margin of any win in the table. If the server benchmark shows single-sequence 4K regressing, raising MIN_TOTAL_KV_TOKENS to 8192 gives up only that cell.

resolve_paged_slab_blocks reads the per-slot context from sched_config.max_kv_size, which is None when --ctx-size is unset; it then falls back to DEFAULT_CTX_LEN (8192). I did not find a more direct per-slot context figure on the worker, so a server started without --ctx-size sizes its slab for 8192 tokens per slot and returns longer sequences to the gather path.

Closes #899

inureyes added 6 commits July 31, 2026 10:36
First half of issue #899: the machinery that lets the fused paged-attention decode v2 kernel from #898 serve a production decode batch, without yet wiring any model onto it.

`cache::paged_batch_decode::paged_batch_decode_attention` is the new entry point. It takes a whole batch of pool-backed `KVCache`s for one layer, appends this step's K/V for every sequence, and then runs attention once over the batch instead of per sequence. Its contract is all-or-nothing: it either declines before writing anything (so the caller's unchanged per-sequence loop runs against an untouched pool) or it owns the step, including running the gather-then-SDPA fallback itself when v2 is not the right choice. That is what keeps a fallback from double-writing the pool.

`paged_v2::dispatch` is the selector. #898 measured exactly one cell where v2 loses to gather (batch 1 at 1024 tokens, 0.91x median, below parity in all three repetitions), so v2 must not be dispatched unconditionally. Re-indexing the measured table by total visible KV tokens in the launch separates the loss from every win at 4096: the loss has 1024 total tokens, the weakest win (batch 1 at 4096, 1.08x) has 4096, and both cells sitting exactly on 4096 win. `MIN_TOTAL_KV_TOKENS` is that number and `MLXCEL_PAGED_V2_MIN_KV_TOKENS` moves it without a rebuild. `layers::resolve_paged_v2_dispatch` layers the existing `MLXCEL_PAGED_ATTENTION_NATIVE` override on top, so its force-off values are the kill switch the issue asks for and its force-on values bypass the floor for re-measurement.

`paged_v2::plan_cache` is the caching. Building a CSR page table costs one hash lookup per visible page (~131k per decode step for a 32-layer model at batch 4 and 32K context); refreshing the per-request scalars costs four writes per request. A decode step changes only the scalars unless a request crosses a page boundary, so the page list is kept and the scalars are recomputed (never incremented, so a missed invalidation cannot turn into an out-of-bounds read). Reuse is gated on a per-request page-range fingerprint plus `PagedBlockPool::block_epoch`, a new counter bumped by every mutation that can move a block. The epoch is what makes admission, eviction, preemption, finish, and page-boundary crossing invalidate the cache without any scheduler plumbing: each of those already goes through a pool method that moves blocks.

`PagedBlockPool::paged_decode_batched` ties the three together and is the only new pool method on the decode path.

Refs #899.
…ernel

Second half of issue #899: the models and the server now reach the whole-batch entry point, and the pool can actually serve it.

**Model wiring.** `Llama3` and `Qwen3` `forward_split_attention` call `paged_batch_decode_attention` before their existing dense-compat paged block. Between them these two attention modules back Llama, Mistral, Qwen2 / Qwen2.5, Qwen3, Helium, and every VLM whose text backbone is one of those, which is the whole set of families that get pool-backed caches. Multi-token steps never enter it (`seq_len != 1` declines), so batched prefill and speculative / MTP verify keep their current paths by construction rather than by a flag.

**The slab constraint, which was the real blocker.** Both fused decode kernels read one contiguous pool buffer per side, and `PagedBlockPool` allocated physical rows in fixed 32-block slabs (#235), so any layer past 32 rows was multi-slab and declined. At `block_size` 32 that capped the fused path at 1024 tokens across the entire batch, which is below the token floor the dispatch policy picks it for: v2 would never have run in production. The slab size is now a per-pool value (`PagedBlockPool::set_slab_blocks`) defaulting to the historical 32, so nothing changes for callers that do not set it, and the server sizes it from the configuration it already reserved KV memory for: `ceil(per_slot_ctx / block_size) * batch`, floored at the old default and clamped by the per-layer share of the paged block budget. `MLXCEL_PAGED_SLAB_BLOCKS` overrides it, with `0` pinning the old default.

The cost is that a layer's first write allocates its whole slab instead of growing into it, front-loading exactly the KV bytes the startup estimate already reports for that batch and context. Growth past the slab still appends without copying, so an under-sized slab degrades to the pre-#899 gather path rather than failing. Operators serving longer contexts than `--ctx-size` implies should raise `--ctx-size` or set the env override, otherwise the fused path silently stops being eligible.

**Memory accounting.** `resolve_paged_block_budget` now subtracts a v2 workspace reserve before converting the KV byte budget into blocks, so admission reserves blocks it can back. The reserve bounds `num_chunks` from the plan's own search invariant (`num_chunks * ctas_per_chunk` lands within a factor of two of the device CTA target), which makes the dominant term independent of the head counts. It comes to single-digit megabytes; it is charged because the accounting should be truthful, not because it is large.

Refs #899.
Fifteen tests for `paged_batch_decode_attention` (issue #899), in two halves.

The decline half is pure and pins the contract that matters most: the function must touch nothing before it has decided it can serve the batch, because a partial write followed by the caller's own loop would double-write the pool. A multi-token step (batched prefill, speculative / MTP verify) is checked to leave both the sequence lengths and the pool's block count at zero.

The parity half drives a real `PagedBlockPool` through real `KVCache::new_paged` caches, runs one whole-batch decode step, and compares against `gather_fallback` computed from the same post-write pool state, so the two see identical K/V and a mismatch can only be the kernel. Covered: four sequences at exactly the 4096-token dispatch floor, ragged lengths that are not page multiples, GQA 1, and the below-floor case where the entry point takes its own gather fallback (which must reproduce the reference exactly, not merely within tolerance). A separate test pins that a layer spread across several slabs declines the fused path and still answers correctly from gather, and another pins that four consecutive steps inside one page rebuild the CSR view at most once.

Refs #899.
…ode v2

Two harnesses issue #899 requires, one runnable now and one for the maintainer's serialized benchmark host.

`tests/paged_decode_v2_greedy_parity.rs` is the mandatory regression guard: batch 4, mixed prompt lengths, 8192-token prompts by default, greedy decode compared token for token between the fused v2 path and the gather baseline. Both arms run in one process, which `MLXCEL_PAGED_ATTENTION_NATIVE` cannot do (it is read once behind a `OnceLock`), by selecting the path structurally instead: the v2 arm sizes the pool slab so every layer is single-slab and the fused kernels are eligible, the gather arm leaves the pool on its 32-row default so they decline. Everything else, prompts included, is identical, so a token difference can only be the kernel. The test then asserts the v2 arm actually built a decode plan, so a silent fallback cannot pass as parity.

Checkpoint resolution is `MLXCEL_PARITY_MODEL`, then `models/<name>`, then `~/.cache/mlxcel/models/<repo>`, and each case soft-skips when nothing resolves. `MLXCEL_PARITY_PROMPT_LEN` and `MLXCEL_PARITY_STEPS` size the run.

Verified locally on three families at the default 8192-token, 32-step configuration: Qwen3-0.6B, Llama-3.2-1B, and Qwen2.5-0.5B all produce identical token streams across all four sequences.

`scripts/benchmark_paged_decode_production.sh` drives the server-level matrix the issue makes mandatory (4 concurrent clients at 1K / 4K / 16K, plus single-sequence 16K and 32K) against a freshly started server per arm, using the existing `bench_serving_concurrency.py` load generator. It defaults `--ctx-size` to cover the 32K case, because undersizing it makes the paged slab too small and both arms would silently measure gather; it also echoes the server's resolved slab-size line so that mistake is visible rather than silent.

Refs #899.
…erving guide

ADR 0001 chose gather-then-SDPA and named the trigger that would justify building a fused kernel: single-sequence context past ~16384 tokens, or any sustained batched decode. Since v0.4 the server defaults to `--parallel 4`, and #898's measurements met both trigger points on an M1 Ultra (1.41x to 3.08x at batch 4, 1.29x and 1.47x single-sequence at 16K and 32K), reproducing the ADR's own 2x-3x prediction almost exactly. Its status is now marked as superseded in part, with a closing addendum stating the reversal and its three qualifications: the switch is gated by a measured token floor rather than unconditional, the #235 single-slab constraint was relaxed rather than removed, and the library-only entry point #710 retired stays retired.

`docs/CONTINUOUS_BATCHING.md` gains an operator-facing section under paged decode: what the fused launch does, the measured speedups, the table of cases that still fall back, and the two environment knobs. The slab-size note is the one an operator can actually get wrong, so it says plainly that serving contexts longer than `--ctx-size` implies returns layers to the gather path silently, and points at the startup log line that shows the resolved value.

Also folds the plan cache's five batch-identity parameters into a `DecodeBatchKey` struct (clippy `too_many_arguments`, and they never travel apart anyway).

Refs #899.
`resolve_paged_block_budget` now charges the fused decode workspace to the KV byte budget before dividing it into blocks (#899), so a test asking for exactly N blocks has to ask for the workspace too. The existing explicit-bytes test is updated to do that, and keeps a case that pins the shortfall a budget which ignores the reserve comes back with, since that shortfall is the whole point: the blocks the pool hands out are blocks it can back.

Adds a test for the reserve itself, bounding it well below any serving KV budget and pinning that a model with no derivable geometry reserves nothing rather than guessing.

Refs #899.
@inureyes inureyes added status:review Under review type:performance Performance improvements priority:high High priority area:core mlxcel-core: MLX FFI, primitives, KV cache, layers area:inference Generation, sampling, decoding (incl. speculative, DRY) labels Jul 31, 2026
Which kernel a server actually ran was unobservable. The counters in `paged_batch_decode` are process-local and nothing reads them, so a launch that silently fell back to gather looked exactly like v2 running and not helping. A before/after benchmark cannot be interpreted without that distinction, and the production benchmark for this issue was in fact comparing gather against gather without any signal saying so.

Logs the first fused launch and the first gather fallback once each at info level, and gives the single-slab decline a reason line naming the actual slab counts. One-shot so a hot decode loop cannot spam the log.
The production benchmark for #899 measured gather against gather across all five scenarios. Two independent causes, both now fixed and both verified on a live `mlxcel-server` against Llama-3.2-1B.

**Cause 1: the single-sequence decode path was never wired.** `dispatch_sync_decode` routes a batch of one to `decode_single_step`, which calls the single-sequence `forward` and therefore `Attention::forward`, not `forward_split_attention`. Only the latter had the new entry point. So the two single-sequence scenarios in the issue's matrix (16K and 32K) could not reach the fused kernel by construction, and neither could any moment in a batched run where only one request was still decoding, which with staggered arrivals is most of it. `Attention::forward` in llama3 and qwen3 now takes the same whole-batch entry point with a batch of one; the `[1, H, 1, D]` decode tensors are already the shape it serves.

**Cause 2: the dispatch floor had the wrong shape.** It required 4096 total visible KV tokens summed over the launch. That fits the #898 table but misreads it: the single measured loss is a property of batch 1, not of total tokens. The ctx-1024 column runs 0.91x at batch 1, 1.41x at batch 4, 1.47x at batch 8, same per-request context and opposite outcome, because a batched launch spreads the same chunk count over more requests and amortizes the merge. A total-token floor separates those cells only by accident and with no margin, and the benchmark's nominal 1K scenario delivered 956 tokens per request, summing to 3824, so it was declined despite being the same shape as the measured 1.41x win. The floor is now stated the way the measurements are: a lone request needs 4096 visible tokens, a multi-request launch needs 512 per request. Every measured cell classifies correctly, batch 2 and 3 are interpolated and said to be.

**And the reason none of this was visible.** Declines were reported with `tracing::debug!`, which a real server does not emit; `RUST_LOG="info,mlxcel_core=debug"` produced no lines in the run that was supposed to validate the change. The decision is now a returned `PagedDecodeOutcome` carrying the numbers behind it, and the caller announces the first occurrence of each distinct outcome kind at info. Per kind, not one global one-shot: a single flag reports only whatever happened first, typically a warmup request, and a later permanent decline for another reason never surfaces, which is how a whole sweep ran on the gather path in silence.

Deciding before building the plan also stops a declined launch paying for the chunk search, and makes the plan-rebuild counter mean "the fused path ran" rather than "was considered", which is what the new tests assert on.

Verified on `mlxcel-server` with no `RUST_LOG` set: 4 concurrent clients at a nominal 1K prompt log `fused v2 launch (batch 4, 3828 visible KV tokens, 16 chunks, merge on)`, a single client at 16K logs `fused v2 launch (batch 1, 13845 visible KV tokens, ...)`, and the same binary under `MLXCEL_PAGED_ATTENTION_NATIVE=0` logs `gather: pinned by MLXCEL_PAGED_ATTENTION_NATIVE` and never a fused launch. Greedy token parity re-verified across Qwen3-0.6B, Llama-3.2-1B and Qwen2.5-0.5B at batch 4 with 8192-token prompts.

Refs #899.
… loudly

`docs/CONTINUOUS_BATCHING.md` and the ADR 0001 addendum described a single 4096-token total floor, which was the formulation the production benchmark disproved. Both now describe the two-regime floor and say why it has that shape: at 1024 tokens of context per request the kernel runs 0.91x at batch 1 but 1.41x at batch 4, so the loss is a property of the request count, not of total work.

The serving guide also gains the diagnostic an operator actually needs. The dispatch outcome is announced at info with no `RUST_LOG`, the four lines it can print are shown verbatim, and the doc states plainly that a run which never prints `fused v2 launch` did not use the fused kernel whatever its throughput looks like. The floor overrides are documented next to the kill switch.

`scripts/benchmark_paged_decode_production.sh` now checks each arm after its scenarios and exits non-zero when the after arm never logged a fused launch, or when the before arm did. The first sweep for this issue produced a complete matrix of null results because both arms were on the gather path and nothing said so; a harness that cannot tell you which kernel it measured is worse than none, so it fails instead of reporting.

Refs #899.
@inureyes

inureyes commented Aug 1, 2026

Copy link
Copy Markdown
Member Author

Root cause of the null benchmark, with evidence

Two independent causes, not one. Both are fixed and both are now verified on a live server.

1. The single-sequence decode path was never wired (the one I missed)

dispatch_sync_decode sends a batch of one to decode_single_step, which calls forward_with_sequence_id and therefore the model's single-sequence Attention::forward. My change only touched forward_split_attention, the batched path. So batch1_ctx16k and batch1_ctx32k could not reach the fused kernel by construction, and neither could any moment in a batched run where only one request was still decoding, which with staggered prefills is most of a scenario. Your server_after.log shows exactly that staggering: the batch4_ctx16k prefills complete at 4343ms, 8351ms, 12535ms and 16537ms, roughly 4 seconds apart, against a decode phase far shorter than that gap.

Fixed by calling the same whole-batch entry point with a batch of one from Attention::forward in llama3 and qwen3. The decode tensors there are already [1, H, 1, D], which is the shape it serves.

2. The floor had the wrong shape, and 4096 does not transfer

You asked me to say plainly whether the floor as specified excludes the scenarios the benchmark measures. It does, and the formulation was wrong rather than the number.

seq_lens does hold what I assumed, and the single-slab check was not firing: your log shows Paged KV slab size: 4096 blocks per layer, and the largest scenario needs about 1732 rows. What went wrong is the shape of the rule. Re-read the #898 ctx-1024 column: 0.91x at batch 1, 1.41x at batch 4, 1.47x at batch 8. Same per-request context, same plan degeneracy, opposite outcome. The loss is a property of batch 1, not of total tokens. A total-token floor separates those two cells only by accident, and with no margin: a batched launch at 1024 tokens per request sits exactly on 4096. Your run delivered 956 tokens per request, 4 x 956 = 3824, and was declined despite being the same shape as the measured 1.41x win.

The floor is now stated the way the measurements are:

launch floor evidence
one request 4096 visible tokens 1024 loses (0.91x), 4096 wins (1.08x)
more than one 512 visible tokens per request 1024 per request wins at batch 4 (1.41x) and batch 8 (1.47x)

512 rather than 1024 for the batched case is deliberate: 1024 is the lowest measured batched context and it wins comfortably, so putting the floor on top of a measured point declines real workloads landing just under it for no measured reason, which is precisely what happened. Batch 2 and 3 are interpolated, not measured; the trend across batch 1, 4 and 8 is monotone and the mechanism (a batched launch spreads the same chunk count over more requests and amortizes the merge) explains why, so the interpolation is stated in the module docs rather than hidden.

3. Why you were debugging blind

You are right that the diagnostics were unusable. They are now a returned PagedDecodeOutcome carrying the numbers behind the decision, announced by the caller at info, one line per distinct outcome kind. Per kind matters: a single global one-shot reports only whatever happened first, which in a real server is a warmup request, so a later permanent decline for a different reason never surfaces. That is how a full sweep ran on the gather path in silence.

I could not reproduce a broken RUST_LOG; EnvFilter::try_from_default_env() in initialize_server_logging should honour mlxcel_core=debug and there is no max_level feature compiling the macros out. Rather than keep chasing it, the diagnostics no longer depend on it.

Verification on a real server

mlxcel-server, Llama-3.2-1B-Instruct-4bit, --parallel 4 --ctx-size 131072, no RUST_LOG set:

# 4 concurrent clients, nominal 1K prompt
INFO decode_step{batch_size=4}: mlxcel_core::cache::paged_batch_decode:
  paged decode v2: fused v2 launch (batch 4, 3828 visible KV tokens, 16 chunks, merge on)

# 1 client, 16K prompt (fresh process)
INFO decode_step{batch_size=1}: mlxcel_core::cache::paged_batch_decode:
  paged decode v2: fused v2 launch (batch 1, 13845 visible KV tokens, 16 chunks, merge on)

# same binary, MLXCEL_PAGED_ATTENTION_NATIVE=0
INFO decode_step{batch_size=4}: mlxcel_core::cache::paged_batch_decode:
  paged decode v2: gather: pinned by MLXCEL_PAGED_ATTENTION_NATIVE

3828 tokens at batch 4 is the exact shape the old floor declined, and batch 1 at 13845 is the path that did not exist. The kill-switch arm never logs a fused launch, so the two arms are now distinguishable from the log alone.

I did not run a timing sweep. I will note one incidental number only so you know the sweep is worth your time, with the caveat that it is a single unreplicated run on a machine that had just finished a build and is not a measurement: at 4 clients on a nominal 1K prompt the fused run reported 75.4 mean decode tok/s against 41.0 for the kill-switch run. Treat that as a smoke signal, nothing more.

The harness now refuses to produce a null sweep

scripts/benchmark_paged_decode_production.sh checks each arm after its scenarios and exits non-zero if the after arm never logged a fused launch, or if the before arm did. It also prints the first few paged decode v2: lines per arm, so a decline names itself.

What to run

export DEVELOPER_DIR=/Applications/Xcode-26.6.0.app/Contents/Developer
cargo build --release --features metal,accelerate
MODEL=~/.cache/mlxcel/models/mlx-community/Llama-3.2-1B-Instruct-4bit \
  ./scripts/benchmark_paged_decode_production.sh

If it aborts, the reason is on a paged decode v2: gather: ... line in benchmarks/paged_decode_production_logs_<date>/server_after.log and names what to change.

Greedy parity re-verified after the Attention::forward change, across three families at batch 4 with 8192-token prompts and 32 steps:

cargo test --release --test paged_decode_v2_greedy_parity --features metal,accelerate -- --ignored --nocapture

Qwen3-0.6B, Llama-3.2-1B and Qwen2.5-0.5B all identical to the gather baseline on all four sequences.

… Ultra

Adds the server-level performance validation issue #899 requires: both arms across 4 concurrent clients at 1K/4K/16K and single-sequence 16K/32K, three full sweeps, fresh server per arm.

Decode throughput improves in four of five scenarios and is at parity in the fifth, with no regression anywhere. The two tightest cells are 4 clients at 16384 tokens (1.13x to 1.15x across sweeps) and a single client at 32768 (1.41x to 1.44x); both are large relative to their own spread and reproduce every time.

Records dispatch proof alongside every number, because the first attempt at this benchmark returned a completely null result that looked like a real finding and was not: both arms had silently run the gather path, so the comparison was gather against gather. The harness now fails a run where the after arm never logs a fused launch, and the report carries the log lines showing which kernel each arm executed.

States the required outcome precisely rather than rounding it up. Batched decode improves at 16384 but is at parity at 4096, and batched aggregate throughput does not improve at either, because aggregate includes time to first token and prefill dominates it at long context. ADR-0001's 2x-3x figure is an op-level claim that #898 reproduced at the op level; end-to-end serving throughput does not follow it by the same factor, since attention is one component of a decode step.

Leaves the non-monotonic 4096 cell as an open question. It gains nothing while 1024 gains 1.57x and 16384 gains 1.15x, it reproduces within 4 percent across three sweeps, so it is a real property of that shape and worth understanding before the dispatch floor is tuned further.
@inureyes inureyes added status:done Completed and removed status:review Under review labels Aug 1, 2026
@inureyes
inureyes merged commit 649f0a5 into main Aug 1, 2026
5 checks passed
@inureyes
inureyes deleted the feature/issue-899-production-paged-v2-dispatch branch August 4, 2026 12:08
@inureyes inureyes self-assigned this Aug 31, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

area:core mlxcel-core: MLX FFI, primitives, KV cache, layers area:inference Generation, sampling, decoding (incl. speculative, DRY) priority:high High priority status:done Completed type:performance Performance improvements

Projects

None yet

Development

Successfully merging this pull request may close these issues.

Route production batched paged decode through the fused v2 kernel, retiring the gather-then-SDPA hot path

1 participant