feat(server): route production batched paged decode through the fused v2 kernel - #988
Conversation
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.
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.
Root cause of the null benchmark, with evidenceTwo 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)
Fixed by calling the same whole-batch entry point with a batch of one from 2. The floor had the wrong shape, and 4096 does not transferYou 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.
The floor is now stated the way the measurements are:
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 blindYou are right that the diagnostics were unusable. They are now a returned I could not reproduce a broken Verification on a real server
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
What to runexport 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.shIf it aborts, the reason is on a Greedy parity re-verified after the cargo test --release --test paged_decode_v2_greedy_parity --features metal,accelerate -- --ignored --nocaptureQwen3-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.
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 preservesMLXCEL_PAGED_ATTENTION_NATIVE=0as 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:
MIN_SINGLE_REQUEST_KV_TOKENS= 4096 visible tokensMIN_BATCHED_KV_TOKENS_PER_REQUEST= 512 per request512 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_TOKENSandMLXCEL_PAGED_V2_MIN_KV_TOKENS_PER_REQUESTmove the two floors;MLXCEL_PAGED_ATTENTION_NATIVE=1bypasses them entirely, which is the supported way to benchmark the declined corner.Both decode paths are wired, not just the batched one
dispatch_sync_decoderoutes a batch of one todecode_single_step, which calls the model's single-sequenceAttention::forward, notforward_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. Bothforward_split_attention(whole batch) andAttention::forward(batch of one) now take the same entry point.Which kernel actually ran is visible at info
The dispatch decision is a returned
PagedDecodeOutcomecarrying the numbers behind it, and the caller announces the first occurrence of each distinct outcome kind at info, with noRUST_LOGneeded. 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.shnow 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, noRUST_LOG): 4 concurrent clients at a nominal 1K prompt logfused v2 launch (batch 4, 3828 visible KV tokens, 16 chunks, merge on); a single client at 16K logsfused v2 launch (batch 1, 13845 visible KV tokens, ...); the same binary underMLXCEL_PAGED_ATTENTION_NATIVE=0logsgather: pinned by MLXCEL_PAGED_ATTENTION_NATIVEand never a fused launch.The thing that was actually blocking this issue
PagedBlockPoolallocated 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. Atblock_size32 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_BLOCKSoverrides it,0pins 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-sizeimplies 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) anddocs/CONTINUOUS_BATCHING.mdsays 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_dispatchlayersMLXCEL_PAGED_ATTENTION_NATIVEon 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_blocksandset_slab_blocks.src/models/llama3.rs,src/models/qwen3.rs:forward_split_attentioncalls 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_blocksand 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/CONTINUOUS_BATCHING.mdgains 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) plusPagedBlockPool::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_prefillacquires), eviction and preemption and finish (release_sequencereleases), page-boundary crossing (append_tokensacquires), 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 fromlogical_start / page_sizeonward and setsfirst_page_offsettological_start % page_size, so requestr's tokeniresolves to absolute positionlogical_start + iand the retired prefix is never addressed. That is the same[logical_start, len)windowgather_visibleslices, 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, becausetrim_front_keep_sinkreturns0for pool-backed caches), and RoPE is applied upstream from the scheduler'srope_offsets, so the view's ownrope_offsetsis not consumed here.Workspace budgeting
The workspace is the partial kernel's
(partial_v, lse)output pair,PagedDecodePlan::workspace_bytes().resolve_paged_block_budgetnow subtracts a reserve before converting the KV byte budget into blocks, so admission reserves blocks it can actually back. The reserve boundsnum_chunksfrom the plan's own search invariant (num_chunks * ctas_per_chunklands within a factor of two of the device CTA target, andctas_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_repis 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_storageand::auto_decode_storage— 2 passed.cargo test --release --test paged_scheduler_parity --features metal,accelerate -- --ignored— 4 passed;paged_real_model_parity1,paged_prefix_share_parity7,paged_kv_serialize_parity2,paged_handoff_parity(withtest-utils) 2.cargo test --release --test paged_decode_v2_greedy_parity --features metal,accelerate -- --ignored --nocaptureat 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 warningsandcargo clippy --lib --tests --features metal,accelerate -- -D warningsclean.cargo fmt --all -- --checkclean.scripts/benchmark_paged_decode_production.shdrives 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, withMLXCEL_PAGED_ATTENTION_NATIVE=0as 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, andDEFAULT_TARGET_CTASstill 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) * batchis 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=0restores 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_TOKENSto 8192 gives up only that cell.resolve_paged_slab_blocksreads the per-slot context fromsched_config.max_kv_size, which isNonewhen--ctx-sizeis unset; it then falls back toDEFAULT_CTX_LEN(8192). I did not find a more direct per-slot context figure on the worker, so a server started without--ctx-sizesizes its slab for 8192 tokens per slot and returns longer sequences to the gather path.Closes #899