Skip to content

perf(core): compute shared prompt prefixes once per decode batch - #1004

Merged
inureyes merged 4 commits into
mainfrom
feature/issue-903-cascade-attention
Aug 2, 2026
Merged

inureyes merged 4 commits into
mainfrom
feature/issue-903-cascade-attention

Conversation

@inureyes

@inureyes inureyes commented Aug 2, 2026 •

Copy link
Copy Markdown
Member

Summary

Cascade (shared-prefix) decode: a whole-page prompt prefix shared by several sequences of a decode batch is attended once for the subgroup and merged into each member's per-request suffix state, instead of being re-read once per sequence. mlxcel already deduplicates the storage of such a prefix (the paged pool refcounts blocks, and CachePool::clone_detached_paged_prefix hands the same block ids to every sequence that adopts the same APC entry); this deduplicates the compute.

No kernel was written or modified. Both levels are ordinary paged-decode v2 launches and the merge is issue #898's paged_attention_merge_states, used exactly as its contract states. Two issues have now consumed that contract without editing it.

Default off (MLXCEL_CASCADE_ATTENTION). The thresholds are derived from what the decomposition costs, not from a measurement, and I did not run benchmarks. Correctness is verified; throughput is not. Flipping the default is one constant, paged_v2::cascade::DEFAULT_CASCADE_ENABLED, once a number exists under docs/benchmark_results/.

The decomposition, and the one part that is not obvious

Level 0 is the shared span, level 1 is each request's private suffix, and the merge pairs each member's two (V, LSE) states. The two key ranges are disjoint, so the merge is exact up to f32 rounding.

The part that makes it a win rather than two launches doing the same work is that level 0 is a one-request launch with M times as many query heads. The v2 partial kernel gives one threadgroup all query heads of one KV head for one (request, page tile) pair, and that threadgroup loads each K and V element once and reuses it across every query head it owns. Handing the same page list to M separate requests would read the span M times; handing it to one request with M times as many query heads reads it once.

The kernel maps query head h to KV head h / NRep, so the stacking order is load-bearing and must be KV-head major: h0 = kv_head * (M * G) + m * G + g, which is a [M, Hkv, G, D] -> [Hkv, M, G, D] transpose in and its inverse out. Stack member-major instead and every query head silently reads the wrong KV head, with correct shapes and a successful launch. member_major_head_stacking_reads_the_wrong_kv_head is the negative control that pins it.

How the #898 merge kernel is reused, and whether its contract held

Unchanged, and yes. v_in [N, H, D] with each partial already normalized by its own denominator, lse_in [N, H] in log2 units, o_indptr [M + 1] grouping contiguous rows into output rows. Both cascade levels are v2 launches, so their LSE is already the m + log2(l) the partial kernel emits and no unit conversion happens anywhere on this path, which is the safest possible way to honour the clause that fails silently. merge_rejects_natural_log_lse_units_on_the_cascade_path restates #907's negative control on cascade-shaped partials: the log2 merge reproduces the host reference and the natural-log merge does not.

Two further contract clauses are load-bearing here and are relied on rather than worked around. A merge group of one row resolves to the identity (l = 1, out = v), which is what lets a sequence outside the sharing subgroup ride along in the same launch with its whole range at level 1. And the regrouping property paged_v2::launch_tests::merge_is_associative_across_regroupings pins is exactly what makes a prefix-plus-suffix decomposition legal at all. Nothing in paged_attention_v2.cpp, paged_attention_v2_merge.cpp or the FFI signature was touched, so #907's MLA split-KV is unaffected.

The one additive change is V2Context::launch_with_lse, which returns the (V, LSE) pair both kernels already produce instead of discarding the LSE; launch() is now that call plus the reshape it always performed. That is issue task 3 ("make it a returnable output") landing as plumbing at the Rust boundary rather than as a kernel change.

Detection reads the page table, not the refcounts

detect_shared_prefix compares the CSR view's indices positionally across requests. That is not an approximation of the refcount check, it is the refcount check: within one PagedCsrView every live block id resolves to exactly one physical pool row and two live ids never share a row, so equal rows at the same position means one refcounted block and therefore identical bytes. It keeps the decision inside mlxcel-core with no scheduler plumbing, is testable without a server, and catches sharing that arose any way at all rather than only sharing the server knows it created. the_shared_blocks_really_are_shared_in_the_pool asserts refcount == 3 on a fixture built through retain_block, the same fork an APC prefix clone performs.

Task items

  1. Shared-span identification: paged_v2::cascade::detect_shared_prefix. Groups candidates by first emitted page, computes the longest common page prefix per group, and picks the subgroup that removes the most page reads (shared_pages * (members - 1)). One shared level only, as scoped.
  2. Two-level decode step: paged_v2::cascade::build_cascade_plan splits the batch page table into a one-request level-0 view over the shared span and a whole-batch level-1 view with the members' shared pages removed; paged_v2::cascade_launch::run_cascade_decode runs both and merges. Non-members keep their full range at level 1 and a merge group of one.
  3. Both kernels emit LSE: V2Context::launch_with_lse, described above.
  4. Fallback and kill switch: MLXCEL_CASCADE_ATTENTION (0/false/off/no disables, 1/true/on/yes enables, anything else takes the shipped default, which is off). Mid-page window starts, spans that would reach a member's last page, sub-threshold spans and sub-threshold member counts all fall back to flat. Soft-cap families, multi-token steps, non-pool-backed caches and multi-slab layers are declined upstream by the Route production batched paged decode through the fused v2 kernel, retiring the gather-then-SDPA hot path #899 path before cascade is consulted.
  5. Scheduler observability: counters in mlxcel-core next to the decision, copied into BatchObservability by publish_metrics, exposed as mlxcel_paged_decode_launches_total{path="fused_v2"|"gather"|"cascade"} plus mlxcel_cascade_shared_tokens_total, mlxcel_cascade_member_sequences_total and mlxcel_cascade_failures_total.

Proving which path a benchmark arm took

Three independent ways, because this is the failure mode the epic keeps paying for.

  • PagedDecodeOutcome gains FusedCascade { batch, members, shared_pages, shared_tokens, prefix_chunks, suffix_chunks } and CascadeFailed(reason), both announced once per kind at info by the existing one-shot-per-outcome-kind table, so a cascade that degrades to flat says so instead of degrading in silence.
  • examples/cascade_attention_bench.rs prints, on stdout, the launch statistics of the arm it just timed and asserts that level 0 really folded the member queries onto the head axis. An arm reporting members=1, or missing, is visibly not measuring what its name says.
  • scripts/bench_serving_concurrency.py --metrics scrapes /metrics around each concurrency level, prints the per-path counter deltas, and prints an explicit NOTE: no cascade launch in this level when the cascade counter did not move.

How to run the benchmarks

Layer level, which is where the whole effect lives (everything outside attention is identical between arms):

cargo run --release --features metal,accelerate --example cascade_attention_bench -- \
    --shared 2048,8192 --batches 4,8 --tail 256 --steps 200 --warmup 40 --reps 5

It sweeps the issue's grid, prints median/min/max ms per step per arm, and finishes with the issue's no-sharing overhead gate (the same shapes with nothing shared, timed with and without the detection scan in front of the flat launch). --reps makes dispersion visible from a single invocation; repeat the whole command for run-to-run dispersion.

End to end against a running server with the prompt cache on, once with and once without MLXCEL_CASCADE_ATTENTION=1 on the same build:

python3 scripts/bench_serving_concurrency.py \
    --shared-prefix-tokens 2048 --prompt-tokens 32 \
    --concurrency 4,8 --max-tokens 256 --metrics

Results belong in docs/benchmark_results/cascade-attention-<hw>-<date>.md.

Test plan

  • cargo test --release -p mlxcel-core --lib --features metal,accelerate paged_v2::: 75 passed, 0 failed (19 of them new: 13 host-side planning, 6 GPU).
  • cargo test --release -p mlxcel-core --lib --features metal,accelerate cache::paged: 126 passed, 0 failed.
  • cargo test --release -p mlxcel-core --lib --features metal,accelerate mla::: 33 passed, 0 failed, so MLA matrix-absorbed decode path with compressed-latent KV cache for DeepSeek-family models #907's reuse of the same merge kernel is unaffected.
  • cargo test --release --test cascade_decode_dispatch --features metal,accelerate: 1 passed. Proves PagedBlockPool::paged_decode_batched really dispatches cascade and agrees with the flat Fused paged-attention decode v2: CSR page table, cross-CTA split-KV, and variable-length merge kernels #898 launch.
  • ./target/release/deps/mlxcel-<hash> server::batch::observability: 13 passed; server::routes::metrics: 1 passed.
  • cargo clippy -p mlxcel-core --lib --tests --features metal,accelerate -- -D warnings: clean.
  • cargo clippy --lib --example cascade_attention_bench --test cascade_decode_dispatch --features metal,accelerate -- -D warnings: clean.
  • cargo fmt --all -- --check: clean.
  • cargo test --release --lib --features metal,accelerate --no-run: the root lib test target compiles.
  • examples/cascade_attention_bench.rs executed once at tiny settings to confirm the harness runs end to end and prints its path attribution. Those numbers are not measurements and are not reported.

Correctness coverage

Cascade output is checked against a host f64 reference and against the flat v2 launch, on an exact f32 pool (tolerance 2e-5) and on a realistic f16 pool (5e-3), for a whole-batch subgroup and for a mixed batch with two non-sharing sequences. Planning is checked exhaustively on the host: threshold enforcement, the cap that keeps every shared page full, exclusion of mid-page window starts, subgroup selection by saved page reads, page-for-page conservation between the two levels and the flat view, and every rejection the builder makes.

What I left out, and where I was unsure

  • No benchmark numbers. As instructed. The default therefore stays off, with several precedents in this epic (Fused residual-add RMSNorm and fused RoPE + KV-append decode kernels #905, MLA matrix-absorbed decode path with compressed-latent KV cache for DeepSeek-family models #907, Distribution-preserving speculative-decoding acceptance (chain speculative sampling) #902, MLXCEL_FUSED_QK_NORM).
  • No CUDA code. Nothing here is backend-specific; both levels reuse the existing v2 kernels, whose CUDA bodies Fused paged-attention decode v2: CSR page table, cross-CTA split-KV, and variable-length merge kernels #898 already shipped unvalidated. This adds no new CUDA and validates none.
  • Detection is not memoized across layers or steps. The page table it reads is already cached by PagedDecodeV2Cache, but the scan itself re-runs per layer per step. It is a host-side pass over indices (order 1e3 comparisons per layer at 8K context and batch 8), and the harness's overhead gate is what bounds it; a CascadePlan slot in the existing plan cache is the obvious follow-up if the gate says it matters.
  • The Shape-bucketed kernel autotuner and cold-L2 benchmark methodology #906 autotuner is not wired to either cascade geometry. Both plans use the heuristic chunk size, which matches what the Route production batched paged decode through the fused v2 kernel, retiring the gather-then-SDPA hot path #899 production path already does.
  • Subgroup search is deliberately simple. Every member of a candidate must share the whole span; a member that diverges early shortens the span for the group rather than being dropped. The shape this exists for is N clients behind one system prompt, where all members diverge at the same page.
  • Unsure: whether the level-0 JIT specialization on member count matters in practice. Level 0's NRep is M * G, so a new batch size compiles a new kernel variant once. Bounded and one-time, but a workload whose subgroup size oscillates would pay it repeatedly. I did not measure it.
  • Unsure: whether the threshold defaults are near right. 16 pages (512 tokens at block size 32) and 2 members are reasoned, not measured. Both move by environment variable so the benchmark can find the real crossover without a rebuild.
  • Unsure: how often production actually produces shared blocks. I traced the path (scheduler.rs calls clone_detached_paged_prefix, which pins each block twice and hands the same ids to the adopting sequence, and append_tokens copy-on-writes a shared last block) and the test fixtures reproduce it, but I did not observe a live server doing it. If the cascade counter stays at zero under the serving scenario above while prompt_cache_hits climbs, that is the thing to look at first, not the kernel.
  • Not attempted: nested sharing. A common system prompt plus a per-tenant sub-prefix would need a third level, which the issue scopes out.

Closes #903

mlxcel already deduplicates the storage of a shared prompt prefix: the paged pool refcounts blocks and CachePool::clone_detached_paged_prefix hands the same block ids to every sequence that hits the same APC entry. The compute stays duplicated, so attention bandwidth over the shared span scales with the batch even though every byte read is identical.

paged_v2::cascade detects a whole-page prefix shared by a subgroup of the decode batch and paged_v2::cascade_launch decomposes the step into a shared-span launch, a per-request suffix launch, and one merge. No kernel changed: level 0 is a one-request v2 launch whose query is every member's query stacked KV-head-major onto the query-head axis, which is what makes the shared pages load once per chunk instead of once per member, and the merge is issue #898's paged_attention_merge_states over a grouping that pairs each member's two states and leaves non-members with a group of one (the kernel's identity case).

Detection reads the CSR view's indices positionally rather than the pool refcounts. Within one PagedCsrView every live block id resolves to exactly one physical row and two live ids never share a row, so equal rows at the same position means one refcounted block and identical bytes. That keeps the whole decision inside mlxcel-core and catches sharing that arose any way at all.

V2Context::launch_with_lse returns the (V, LSE) state both v2 kernels already produce; launch() is now that call plus the reshape it always did. Both cascade levels are v2 launches, so their LSE is already in the merge kernel's log2 units and no conversion happens anywhere on this path.

Default off (MLXCEL_CASCADE_ATTENTION), thresholds MLXCEL_CASCADE_MIN_SHARED_PAGES and MLXCEL_CASCADE_MIN_MEMBERS, because the thresholds are derived from what the decomposition costs rather than from a measurement. PagedDecodeOutcome gains FusedCascade and CascadeFailed so a benchmark can prove which path an arm ran and a cascade that degrades to flat says so.

Tests: 13 host-side planning tests plus 6 GPU tests, including two negative controls for the silent failure modes. Member-major head stacking reads the wrong KV head with correct shapes and a successful launch; a natural-log LSE merges into a plausible wrong weighted average, the same clause mla::split_kv_tests::merge_rejects_natural_log_lse_units pins for the MLA caller.

Refs #903
…ch gate

Counters for the cascade path (activations, cumulative shared-span tokens, cumulative member count, fallbacks) live in mlxcel-core next to the decision that produces them, and BatchScheduler::publish_metrics copies them into the batch observability alongside the existing paged gauges. They surface as mlxcel_paged_decode_launches_total{path="fused_v2"|"gather"|"cascade"} plus mlxcel_cascade_shared_tokens_total, mlxcel_cascade_member_sequences_total and mlxcel_cascade_failures_total, so "which attention path is this server running, and how much prefix is it hoisting" is answerable from /metrics without a profiler.

examples/cascade_attention_bench.rs times flat against cascade over a pool whose shared pages are genuinely one refcounted block per page. Every cell prints the launch stats that attribute it (members, shared pages, both chunk counts, the folded level-0 head count) to stdout rather than through tracing, because the binaries install no subscriber. It also carries the issue's no-sharing overhead gate: the same shapes with nothing shared, timed with and without the detection scan in front of the flat launch.

tests/cascade_decode_dispatch.rs is a separate integration binary because the gate and both thresholds are OnceLocks over environment variables, which no unit test inside a shared process can flip deterministically. It sets them, calls PagedBlockPool::paged_decode_batched, and asserts the outcome is FusedCascade with the expected member count and span, then that the result matches the flat #898 library launch. Issue #899 shipped a fused path that never activated and whose benchmark compared the fallback against itself; a cascade that is correct in isolation but unreachable from the production entry point would reproduce that exactly, and nothing in paged_v2 would notice.

The whole-batch case now interleaves the two levels with a stack on a new axis instead of a concatenate plus a gather, which removes two launches from the shape the feature exists for.

Refs #903
…ng scenario

docs/cascade-attention.md covers the decomposition, why the level-0 launch folds the member queries onto the query-head axis KV-head-major (and what silently breaks if that order is wrong), how the shared span is detected off the page table rather than off the refcounts, every refusal and the reason it is load-bearing, the three flags, how to read the active path off a log line or /metrics, and the current limitations. Indexed from docs/README.md, cross-linked from CONTINUOUS_BATCHING.md next to the existing #899 dispatch section, and the three variables are in the environment-variables table.

scripts/bench_serving_concurrency.py gains --shared-prefix-tokens, which sends a byte-identical system prompt from every client of a level and gives each a distinct short question, so the requests share a KV prefix and diverge after it. That is the precondition for the prompt cache to hand out the same refcounted blocks, which is the precondition for cascade to have anything to hoist. --metrics scrapes /metrics around each level and prints the per-path decode counter deltas, and says so explicitly when a level's cascade counter did not move, so a run cannot report a shared-prefix number for an arm that ran the flat path.

Also applies rustfmt to the files added in the previous two commits.

Refs #903
@inureyes inureyes added type:performance Performance improvements priority:medium Medium priority area:core mlxcel-core: MLX FFI, primitives, KV cache, layers status:review Under review labels Aug 2, 2026
Adds the performance-validation report issue #903 requires. Cascade decode is slower than the flat #898 launch in all eight measurements across two independent invocations, ranging 0.356x to 0.784x, so the path ships available but unwired and `DEFAULT_CASCADE_ENABLED` stays false.

Every cascade row carries the launch stats it was measured on, and the harness asserts that member queries were folded, so an arm that silently fell back panics rather than reporting a number. Observed prefix_q_heads of 128 at four members and 256 at eight, exactly q_heads times members. Three earlier measurements in this epic turned out to be comparing a path against itself, which is why this is stated rather than assumed.

The decomposition adds roughly six extra small launches per layer-step, two permutations in, two out, and the merge. At a 2048-token shared span the flat launch reads little enough that this fixed cost dominates. Raising the span to 8192 improves the ratio without reaching parity, so any crossover lies beyond the span sizes the issue asks about.

Records both invocations of the no-sharing overhead gate rather than only the favourable one. The first reported plus 62 percent at batch 4 with a flat arm spanning a 64 percent range within a single invocation, wider than the effect being measured; the second reported minus 0.4 percent with an 11 percent range. The second is the credible number and the gate is roughly met, but the pair is a useful record that one invocation of this harness can produce a large artifact at an apparently quiet load average.
@inureyes inureyes added status:done Completed and removed status:review Under review labels Aug 2, 2026
@inureyes
inureyes merged commit 78c3d6e into main Aug 2, 2026
5 checks passed
@inureyes
inureyes deleted the feature/issue-903-cascade-attention 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:medium Medium priority status:done Completed type:performance Performance improvements

Projects

None yet

Development

Successfully merging this pull request may close these issues.

Cascade attention: compute shared prompt prefixes once per decode batch

1 participant