Repository navigation
perf(core): compute shared prompt prefixes once per decode batch - #1004
Merged
Merged
Conversation
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
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.
6 of 7 tasks
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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_prefixhands 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 underdocs/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
hto KV headh / 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_headis 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 them + 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_pathrestates #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 propertypaged_v2::launch_tests::merge_is_associative_across_regroupingspins is exactly what makes a prefix-plus-suffix decomposition legal at all. Nothing inpaged_attention_v2.cpp,paged_attention_v2_merge.cppor 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_prefixcompares the CSR view'sindicespositionally across requests. That is not an approximation of the refcount check, it is the refcount check: within onePagedCsrViewevery 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 insidemlxcel-corewith 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_poolassertsrefcount == 3on a fixture built throughretain_block, the same fork an APC prefix clone performs.Task items
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.paged_v2::cascade::build_cascade_plansplits 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_decoderuns both and merges. Non-members keep their full range at level 1 and a merge group of one.V2Context::launch_with_lse, described above.MLXCEL_CASCADE_ATTENTION(0/false/off/nodisables,1/true/on/yesenables, 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.mlxcel-corenext to the decision, copied intoBatchObservabilitybypublish_metrics, exposed asmlxcel_paged_decode_launches_total{path="fused_v2"|"gather"|"cascade"}plusmlxcel_cascade_shared_tokens_total,mlxcel_cascade_member_sequences_totalandmlxcel_cascade_failures_total.Proving which path a benchmark arm took
Three independent ways, because this is the failure mode the epic keeps paying for.
PagedDecodeOutcomegainsFusedCascade { batch, members, shared_pages, shared_tokens, prefix_chunks, suffix_chunks }andCascadeFailed(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.rsprints, 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 reportingmembers=1, or missing, is visibly not measuring what its name says.scripts/bench_serving_concurrency.py --metricsscrapes/metricsaround each concurrency level, prints the per-path counter deltas, and prints an explicitNOTE: no cascade launch in this levelwhen 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):
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).
--repsmakes 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=1on the same build: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. ProvesPagedBlockPool::paged_decode_batchedreally 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.rsexecuted 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
MLXCEL_FUSED_QK_NORM).PagedDecodeV2Cache, but the scan itself re-runs per layer per step. It is a host-side pass overindices(order 1e3 comparisons per layer at 8K context and batch 8), and the harness's overhead gate is what bounds it; aCascadePlanslot in the existing plan cache is the obvious follow-up if the gate says it matters.NRepisM * 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.scheduler.rscallsclone_detached_paged_prefix, which pins each block twice and hands the same ids to the adopting sequence, andappend_tokenscopy-on-writes a shared last block) and the test fixtures reproduce it, but I did not observe a live server doing it. If thecascadecounter stays at zero under the serving scenario above whileprompt_cache_hitsclimbs, that is the thing to look at first, not the kernel.Closes #903