Skip to content

feat(core): fused paged-attention decode v2 with CSR page table and cross-CTA split-KV - #984

Merged
inureyes merged 7 commits into
mainfrom
feature/issue-898-paged-decode-v2
Jul 31, 2026
Merged

inureyes merged 7 commits into
mainfrom
feature/issue-898-paged-decode-v2

Conversation

@inureyes

Copy link
Copy Markdown
Member

Summary

Builds paged-attention decode v2 as a library capability: a CSR page table, a cross-CTA split-KV partial kernel, a generic variable-length merge kernel, a host-side plan wired into the autotuner seam issue #906 reserved, and an env-gated entry point. v1 is untouched and remains the only path any caller reaches unless MLXCEL_PAGED_ATTENTION_V2=1 is set.

The structural problem v1 has is that its KV split happens inside one threadgroup, so one CTA serves one (batch, query head) pair no matter how long the context is. A small batch with a long context leaves most of the GPU idle, and adding context adds no parallelism. v2 moves the split across CTAs: parallelism becomes num_chunks * kv_heads * q_groups and grows with context.

What changed

1. CSR view builder. src/lib/mlxcel-core/src/cache/paged_csr.rs plus PagedBlockPool::paged_csr_view. The layout is the standard indices / indptr / last_page_len triple plus a first_page_offset extension: the canonical form assumes every request starts at entry 0 of its first page, which a sliding-window sequence with a non-zero logical_start violates. Carrying the offset lets the view drop fully retired pages instead of emitting them and skipping them inside the kernel, so the plan's page counts are the real work counts. validate asserts the geometry identity (pages - 1) * page_size + last_page_len - first_page_offset == seq_len on every build, because each way it can fail is an out-of-bounds read in the kernel rather than an error. rope_offsets records the absolute next-token position (the written length), which differs from the visible length exactly when a front trim has happened. Shared refcounted blocks need no special case: two requests that share a prefix resolve the same block id to the same physical row, so the row appears in both slices of indices. 17 unit tests, including a 500-case randomized geometry sweep and prefix sharing through a real pool with refcount assertions.

2. Kernel v2. src/lib/mlx-cpp/turbo/paged_attention_v2.cpp, Metal plus CUDA JIT strings in one translation unit with runtime backend selection, following the dual-source pattern of paged_attention.cpp. One CTA owns one (chunk, kv head, q-head group) triple: grid z is the flat chunk index, grid y is kv_head * QGroups + q_group, 32 lanes partition the head dimension, and NumWarps warps stripe the chunk's tokens. The CTA reads each KV element once and reuses it across every query head of its group. Scores run in base 2 (log2(e) folded into the attention scale) so the online softmax is exp2, and the emitted LSE is in log2 units. Two helpers, paged_attention_v2_q_heads_per_cta and paged_attention_v2_num_warps, are exported to Rust so the plan's CTA count and the launcher's grid are derived from one source and cannot drift.

3. Merge kernel. src/lib/mlx-cpp/turbo/paged_attention_v2_merge.cpp, split into its own translation unit because it is the reusable half. It takes v_in [N, H, D] (each partial already normalized by its own softmax denominator), lse_in [N, H] in log2 units, and an o_indptr [M + 1], and returns (v_out [M, H, D], lse_out [M, H]) using the closed-form exp2/log2 rescale generalized to a variable-length group. One thread owns one (output row, head, dim) element and walks its partial list once, so there is no threadgroup memory and no barrier, which is what keeps it usable for an arbitrary grouping. A partial with lse = -inf (an empty chunk) contributes nothing, and an output row with no finite partial yields zeros and lse_out = -inf rather than NaN.

4. Plan. src/lib/mlxcel-core/src/paged_v2/plan.rs. PagedDecodePlan binary-searches pages_per_chunk: the chunk count is non-increasing in the chunk size, so the search takes the largest chunk size whose CTA count still reaches the device target, because among the sizes that saturate the device that one does the least merge work. A floor derived from the 65535 grid-z bound keeps every emitted plan launchable on both backends. It is plain data with a matches predicate, so a caller can hold it across decode steps and only rebuild when a request crosses a page boundary. Every request gets at least one chunk, including a fully trimmed one, so the output always has exactly one row per request. autotune::ops::paged_decode_v2_chunk occupies the OP_PAGED_DECODE_V2_KV_CHUNK seam: it enumerates the feasible chunk sizes (powers of two inside the plan's own bounds, plus the heuristic so the default is measurable), returns the binary-search value as its default tactic, and its run performs a real launch against the borrowed context so profiling measures the kernels rather than array construction. Nothing in the autotuner core changed to accept it.

5. Entry point and flag. paged_v2::run_decode_v2 plus PagedBlockPool::paged_decode_fused_v2, reached from paged_decode_fused behind MLXCEL_PAGED_ATTENTION_V2=1. The merge kernel runs only when some request has more than one chunk; when every request fits in one chunk the plan emits chunks in request order and the partial output is already the answer, so it is reshaped rather than merged. That "write O directly" case is decided on the host, which is what keeps it safe: an MLX custom kernel's output array is uninitialized, so a kernel that skipped elements would leave garbage rather than produce a fast path.

6. Correctness harness. examples/paged_decode_v2_correctness.rs, and a three-way extension of examples/page_gather_microbench.rs (gather-then-SDPA vs fused v1 vs fused v2) that reports v2/v1 and v2/gatherA next to the v2 plan that produced them.

Deliberate scope decisions

CUDA is unvalidated. The implementing host is an Apple M1 Ultra with no CUDA hardware and no nvcc. Both CUDA JIT bodies are structural transliterations of the Metal ones (same thread mapping, same chunk arithmetic, same accumulation order; simd_sum becomes a __shfl_xor_sync butterfly, threadgroup becomes __shared__, exp2/log2/fmax become the f variants) and are marked unvalidated in the source. They have never been compiled or run. The same applies to DEFAULT_TARGET_CTAS: a CUDA target should come from the SM count, which needs a CUDA host to derive.

Default off, v1 intact. With the env unset the entry point is one OnceLock read away from the exact pre-#898 code, no v2 array is built, and neither v2 kernel is JIT-compiled. Verified by running the full ffi_tests module (93 tests, which includes the v1 fused-vs-gather parity tests) both with the variable unset and with it set to 1: 93 pass either way, so v2 also satisfies the parity assertions v1 does when it takes over the path.

No benchmarks were run. The three-way microbench is written but not measured here; these paths are bandwidth-bound and a loaded host produces misleading numbers. It was smoke-tested for launch correctness only (one configuration, two iterations), and no timing from that run is recorded anywhere. docs/benchmark_results/paged-decode-v2-<hw>-<date>.md is therefore not in this PR and should be written from a serialized run on a quiet machine.

The kernel lives in two new files, not in paged_attention.cpp. The issue names that file, but it is already 534 lines and the project's 500-line ceiling would be badly broken by adding ~880 lines to it. The split also puts the merge kernel, which issue #903 reuses unchanged, in its own translation unit.

The harness builds its page table by hand rather than driving PagedBlockPool. That is forced, not a shortcut: the pool allocates physical rows in 32-block slabs and both fused kernels read one contiguous buffer per side, so any context past ~1024 tokens at page size 32 is multi-slab and declined by v1 and v2 alike. A pool-driven matrix could only cover the short-context corner. The pool path is covered instead by the cache::paged_csr and paged_v2::launch unit tests.

Correctness results (Apple M1 Ultra, Metal)

The committed default sweep is a 24-configuration subset (head_dim {64,128} x GQA {1,4,8} x block {32} x ctx {512,4096} x batch {1,4}), not the issue's full 216. The full matrix is available behind --full, and the axis flags select anything in between. What was actually run here is 68 configurations across three sweeps, all passing against the 2e-2 tolerance:

  • The 24-configuration default: worst max relative deviation 5.24e-4, worst relative RMS 3.05e-4.
  • 12 long-context configurations (head_dim 128, GQA 4, block {16,32,64}, ctx {16384,32768}, batch {1,4}): worst max relative deviation 4.43e-4.
  • 32 force-merge configurations with pages_per_chunk = 1 (head_dim {64,128}, GQA {1,8}, block {16,64}, ctx {512,4096}, batch {1,8}), reaching 1408 chunks and 45056 CTAs: worst max relative deviation 5.24e-4.

Block sizes 16 and 64 and batch 8 are covered by the second and third sweeps rather than by the committed default.

Test plan

  • cargo test --release -p mlxcel-core --lib --features metal,accelerate paged_v2:: passes 30/30 (19 plan tests including an exhaustive brute-force check that the binary search really is maximal, 11 GPU tests against a host-side f64 reference computed from the same values written into the pool).
  • cargo test --release -p mlxcel-core --lib --features metal,accelerate cache::paged_csr passes 17/17.
  • cargo test --release -p mlxcel-core --lib --features metal,accelerate cache:: passes 433/433.
  • cargo test --release -p mlxcel-core --lib --features metal,accelerate autotune:: passes 55/55.
  • cargo test --release -p mlxcel-core --lib --features metal,accelerate ffi_tests:: passes 93/93 with MLXCEL_PAGED_ATTENTION_V2 unset and 93/93 with it set to 1.
  • cargo test --release --test autotuner_outcome --features metal,accelerate passes 2/2.
  • cargo run --release --features metal,accelerate --example paged_decode_v2_correctness and two additional sweeps: 68 configurations, 0 failing.
  • cargo clippy -p mlxcel-core --lib --tests --features metal,accelerate -- -D warnings clean, and the same for both touched examples.
  • cargo fmt --all -- --check clean.

Not run: the full cargo test -p mlxcel-core suite, which exceeds this environment's watchdog. The pre-existing fused_moe_parity_tests::fused_moe_geglu_kernel_matches_references_gemma4_shape failure tracked as #964 is untouched by this PR.

Notes for the follow-ups

  • Route production batched paged decode through the fused v2 kernel, retiring the gather-then-SDPA hot path #899 (production wiring): the plan is cacheable, and PagedDecodePlan::matches is the cheap predicate for reuse. workspace_bytes is the budgeting figure; the workspace is MLX-allocated (an MLX custom kernel cannot write through a caller-owned buffer, so "caller-provided workspace" became "outputs whose exact size the plan states up front"). The single-slab constraint from perf(core): chunked slab storage for the paged pool to eliminate growth reallocation entirely #235 still applies and is the first thing to lift for real contexts. device_target_ctas derives its Apple figure from hardware::gpu_core_count, which is the performance-core proxy the tree already uses as a device-scale signal rather than a true GPU core count; it scales with the part but is not calibrated, so treat it as a starting point that MLXCEL_PAGED_DECODE_V2_TARGET_CTAS or the autotuner supersedes.
  • Cascade attention: compute shared prompt prefixes once per decode batch #903 (cascade): the merge kernel is reused unchanged. Its contract is (v_in [N, H, D] normalized partials, lse_in [N, H] in log2 units, o_indptr [M + 1]) returning (v_out, lse_out), with lse_out also in log2 units so a cascade can feed one merge's output into the next. merge_is_associative_across_regroupings in paged_v2::launch_tests pins exactly that property. Note the log2 units: a consumer wanting the natural-log LSE multiplies by ln(2).

Closes #898

inureyes added 6 commits July 31, 2026 02:27
Adds the two GPU kernels behind paged-attention decode v2 (issue #898). v1 in `paged_attention.cpp` is untouched and stays the only path any caller reaches until the Rust plan and entry point land.

`paged_attention_v2.cpp` holds the partial kernel: KV is split across CTAs instead of inside one threadgroup, so parallelism becomes `num_chunks * Hkv * q_groups` and grows with context instead of being pinned at one CTA per `(batch, query head)` pair. One CTA owns one `(chunk, kv head, q-head group)` triple, reads each KV element once and reuses it across every query head of the group, and emits a normalized partial plus its LSE. The page table is a CSR view (`indices` / `indptr` / `last_page_len` plus the `first_page_offset` mlxcel extension that expresses a sliding-window sequence starting mid-page), so no gather pre-pass is needed and one launch covers the whole batch.

`paged_attention_v2_merge.cpp` holds the variable-length merge kernel, split into its own translation unit because it is the reusable half: it takes `(V, LSE)` arrays plus an `o_indptr` and knows nothing about paging, so the cascade-attention issue #903 can drive it with a different grouping and nothing else. Scores run in base 2 (`log2(e)` folded into the attention scale) and the LSE is emitted in log2 units, which is what the merge rescale consumes.

The CUDA bodies are structural transliterations of the Metal ones and are marked unvalidated in the source: the implementing host was Apple Silicon with no CUDA hardware and no `nvcc`, so they have never been compiled or run.

Validated with `cargo check -p mlxcel-core --lib --features metal,accelerate`, which compiles both new translation units through the cxx bridge.

Refs #898
Adds `cache::paged_csr`, the flat page table the v2 decode kernels consume in place of v1's `rows` / `row_offsets` / `logical_starts` / `visible_lens` quadruple (issue #898). `PagedBlockPool::paged_csr_view` builds it for one layer of an active batch; the builder itself is a pure function over `PagedLayerState` plus a row-resolving closure, so it is testable without a pool.

The layout is the standard `indices` / `indptr` / `last_page_len` triple plus a `first_page_offset` extension. The canonical form assumes every request starts at entry 0 of its first page, which a sliding-window sequence with a non-zero `logical_start` violates; carrying the offset lets the view drop fully retired pages instead of emitting them and skipping them inside the kernel, so the plan's page counts are the real work counts. `validate` asserts the geometry identity `(pages - 1) * page_size + last_page_len - first_page_offset == seq_len` on every build, because each way it can fail is an out-of-bounds read in the kernel rather than an error.

`rope_offsets` records the absolute next-token position (the written length), not the visible length; the two differ exactly when a front trim has happened, and RoPE needs the former.

Shared blocks need no special case: two requests that share a prefix resolve the same block id to the same physical row, so the row appears in both slices of `indices`. That property is what lets issue #903 build a cascade decomposition on the same view.

17 unit tests: full and partial final pages, the exactly-full final page that a plain `len % page_size` would report as zero, front-trimmed windows, a 500-case randomized geometry sweep, prefix sharing through a real pool with refcount assertions, empty and fully-trimmed requests, and every rejection path.

`cargo test --release -p mlxcel-core --lib --features metal,accelerate cache::paged_csr` passes 17/17.

Refs #898
Completes the library-level capability behind issue #898: the host-side plan that sizes the cross-CTA split, the run step that turns a CSR view plus a plan into one or two kernel launches, and the `MLXCEL_PAGED_ATTENTION_V2=1` entry point.

`paged_v2::plan` binary-searches `pages_per_chunk`. The chunk count is non-increasing in the chunk size, so the search takes the largest chunk size whose CTA count still reaches the device target: among the sizes that saturate the device, that is the one that does the least merge work. A floor derived from the 65535 grid-z bound keeps every emitted plan launchable on both backends, and the plan is plain data with a `matches` predicate, so a caller can hold it across decode steps and only rebuild when a request crosses a page boundary. Every request gets at least one chunk, including a fully trimmed one, so the output always has exactly one row per request.

`paged_v2::launch` runs the merge kernel only when some request has more than one chunk. When every request fits in one chunk, the plan emits chunks in request order and the partial output is already the answer, so it is reshaped rather than merged. That is the issue's write-O-directly case, decided on the host, which is what keeps it safe: an MLX custom kernel's output array is uninitialized, so a kernel that skipped elements would leave garbage rather than a fast path.

The autotuner seam reserved by issue #906 is now occupied. `autotune::ops::paged_decode_v2_chunk` registers under `OP_PAGED_DECODE_V2_KV_CHUNK`, enumerates the feasible chunk sizes (powers of two inside the plan's own bounds, plus the heuristic so the default is measurable), and returns the binary-search value as its default tactic. Its `run` performs a real launch against the borrowed context, so profiling measures the kernels rather than array construction. Nothing in the autotuner core changed to accept it.

The entry point is one `OnceLock` read inside `PagedBlockPool::paged_decode_fused`: with the env unset no v2 array is built and neither v2 kernel is JIT-compiled, so the default path is unchanged. When v2 declines a shape (multi-slab layer, no visible tokens, a geometry the kernel cannot serve) it returns `None` and v1 runs exactly as before. `single_slab_tensors` is factored out of the v1 path and made public so the correctness harness can drive the kernels directly.

30 tests: 19 plan tests including an exhaustive brute-force check that the binary search really is maximal, and 11 GPU tests that compare the kernels against a host-side f64 attention computed from the same values written into the pool. The GPU set covers single-chunk and merged paths agreeing with the reference and with each other, chunk-size invariance across 7 sizes, GQA head-group mapping, front-trimmed windows, an empty request merging to zeros without poisoning its neighbours, f16 pools against the gather-then-SDPA reference, and the merge kernel on its own (closed form, all-empty rows, and the regrouping associativity that issue #903 depends on).

`cargo test --release -p mlxcel-core --lib --features metal,accelerate paged_v2::` passes 30/30 on an M1 Ultra.

Refs #898
Adds `examples/paged_decode_v2_correctness.rs`, the elementwise comparison of the v2 kernels against the gather-then-SDPA reference (ADR 0001 strategy A) over the issue #898 sweep: head dimension, GQA ratio, page size, context, and batch, with a randomized partial last page and half the requests starting mid-page so `first_page_offset` is exercised.

The harness builds the pool tensors and the CSR page table directly rather than going through `PagedBlockPool`. That is forced, not a shortcut: the pool allocates physical rows in 32-block slabs and both fused kernels read one contiguous buffer per side, so any context past ~1024 tokens at page size 32 is multi-slab and declined by v1 and v2 alike. A pool-driven matrix could only cover the short-context corner. The pool path is instead covered by the `cache::paged_csr` and `paged_v2::launch` unit tests. Physical pages are assigned in reverse order out of a pool with 2x slack, so a kernel that ignored `indices` would fail outright rather than pass on an accidentally contiguous layout.

The committed default is a 24-configuration subset (head_dim {64,128} x gqa {1,4,8} x block {32} x ctx {512,4096} x batch {1,4}) that runs in seconds; `--full` expands to the issue's 216-configuration matrix, several entries of which allocate gigabytes, and the axis flags select anything in between. `--force-merge` pins `pages_per_chunk = 1` so every request goes through the merge kernel with a maximally ragged grouping. The output states the plan (chunk size, chunk count, CTA count, whether a merge ran) next to the deviations, so a recorded result says which regime it measured.

Measured on an M1 Ultra, 68 configurations across three sweeps, all passing against the 2e-2 tolerance: the 24-config default (worst max_rel 5.24e-4), 12 long-context configs at ctx 16384 and 32768 over block sizes 16/32/64 (worst 4.43e-4), and 32 force-merge configs across block sizes 16/64 and batch 8, up to 1408 chunks and 45056 CTAs (worst 5.24e-4). Worst relative RMS across all three was 3.05e-4.

Refs #898
Extends `examples/page_gather_microbench.rs` with the comparison issue #898 asks for: the existing gather-then-SDPA paths now run alongside the fused v1 kernel (#123, split-K inside one threadgroup) and the fused v2 kernels (#898, cross-CTA split-KV plus merge), over the same ADR 0001 sweep of context, batch, and block size.

All three read the same layout-A pool buffers, so they see an identical scatter pattern, and both fused paths consume an f32 query cast once outside the timed region rather than having the cast charged to one of them. The v2 plan is likewise built outside the timed region, because that is the production model: the plan is plain data that stays valid across decode steps until a request crosses a page boundary.

Each row reports `v2/v1` and `v2/gatherA` next to the v2 plan that produced them (chunk size, chunk count, whether a merge launch was needed), so a recorded result states which regime it measured rather than leaving the reader to infer it. Those two ratios are the issue's acceptance bars: `v2/v1 >= 1.0` across the sweep, and `v2/gatherA > 1.0` at the ADR 0001 trigger points (batch 4 at ctx >= 1024, single-sequence at ctx >= 16384).

CSV columns are appended, never reordered, so readers of the pre-#898 schema keep working.

Smoke-tested for launch correctness only (one configuration, two iterations). No measurement is recorded here: these paths are bandwidth-bound and a loaded host produces misleading numbers, so the sweep is meant to run serialized on a quiet machine under `caffeinate -i`.

Refs #898
Adds a #898 addendum to ADR 0001 and names the v2 kernels in the architecture doc's fused-kernel list.

The addendum ties v2 back to the ADR's own "what reopens this" bullets: both of them trace to the same root cause, that v1 splits KV inside one threadgroup so parallelism does not grow with context. v2 attacks that directly by splitting across CTAs. It also states plainly what v2 does not change: the #235 single-slab constraint still holds, so v2 by itself does not reopen the scheduler wire-in question, and the correctness harness drives a hand-built page table for exactly that reason.

Also folds in the lint fixes clippy flagged on the new test code: an argument-count allow on the synthetic-batch builder, an iterator rewrite of the reference accumulator loop, and a range-contains form in the warp-budget assertion.

Verified with `cargo fmt --all -- --check`, `cargo clippy -p mlxcel-core --lib --tests --features metal,accelerate -- -D warnings`, and `cargo clippy --example paged_decode_v2_correctness --example page_gather_microbench --features metal,accelerate -- -D warnings`, all clean.

Refs #898
@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 labels Jul 30, 2026
…n M1 Ultra

Adds the performance-validation report issue #898 requires: gather-then-SDPA against fused v1 against fused v2, batch {1,4,8} by context {1024,4096,16384,32768}, three repetitions with medians and per-cell spread.

Both required outcomes hold. v2 is at or above v1 in every cell, and it beats gather-then-SDPA at both ADR-0001 trigger points: 1.41x to 3.08x for batch 4 at every context measured, and 1.29x and 1.47x for single-sequence decode at 16384 and 32768. The 2x-3x that ADR-0001 predicted for batch 4 past 4096 tokens reproduces closely, at 2.04x, 2.78x and 3.08x. Batched cells at 4096 and above varied by at most 0.07 across repetitions, and batch 4 at 16384 returned 2.78x all three times.

The batch-1 column against v1 is the structural result. v1 splits KV inside a single threadgroup, so context length adds no parallelism, and at batch 1 it is slower than the path it was meant to replace: 3023us against gather's 1063us at 32768. v2's cross-CTA split is 4.11x faster than v1 there.

One cell goes the other way and constrains the follow-up. At batch 1 with 1024 tokens of context, v2 is slower than gather-then-SDPA at 0.91x median, below parity in all three repetitions, because the plan degenerates to two pages per chunk and the merge pass costs more than the attention. That is outside the ADR trigger points and violates no requirement here, but issue #899 must not dispatch v2 unconditionally; the production selector needs a floor below which the gather path stays.
@inureyes inureyes added status:done Completed and removed status:review Under review labels Jul 31, 2026
@inureyes
inureyes merged commit 30f1ca8 into main Jul 31, 2026
5 checks passed
@inureyes
inureyes deleted the feature/issue-898-paged-decode-v2 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 priority:high High priority status:done Completed type:performance Performance improvements

Projects

None yet

Development

Successfully merging this pull request may close these issues.

Fused paged-attention decode v2: CSR page table, cross-CTA split-KV, and variable-length merge kernels

1 participant