perf(cuda): measure SDPA plan-cache bucketing, recommend against it - #1906
Merged
Merged
Conversation
MLX keys its cuDNN SDPA execution-plan cache on the exact shapes and strides of q, k, v and the mask, and a speculative verify round appends keys every round, so every round of every attention-layer class missed that cache and rebuilt a plan on the host: about 22 ms per build on GB10, 67 to 76 ms per round on the Laguna DFlash pairing, and a fatal `Cache thrashing` abort once lifetime misses passed twice the capacity. #1799 routed those calls off cuDNN entirely; this makes the key reusable so they can stay on it. Measured on the pairing #1799 attributed the defect on (GB10, block 4, 25 verify rounds, `MLXCEL_SDPA_PLAN_DEBUG=1`): three shape classes, three plan builds per round, 82 in the process. Only three key fields move round over round, the k/v sequence length and the mask's column count and row stride; the k/v strides and the mask's two leading strides (0, from the broadcast `fast::scaled_dot_product_attention` builds) do not. With bucketing each class builds one plan for the whole run, 11 in the process, and the round's accept counts are unchanged. The canonicalization is MLX's own one-row decode path extended to a small array-masked multi-row call: k and v are widened to a bucket, the mask is widened to the same width with the new columns set to `-inf`, and the true lengths reach cuDNN through `set_padding_mask` with `set_seq_len_q` / `set_seq_len_kv`, which is what keeps the widened region out of the result. k and v reach the bucket by unslicing when they are a leading slice of one cache buffer with room to spare (the target's dense and speculative-buffered caches), and by a zero-padded copy when they are not, which is the third class: the drafter concatenates its proposal keys onto the cache window, so its k and v are freshly built arrays. The copy is bounded by `MLXCEL_SDPA_PLAN_BUCKET_MAX_MB` (64 per tensor); above it the call keeps its exact shape. `MLXCEL_SDPA_PLAN_BUCKET_MAX_QUERIES` (default 32, 0 disables) is the kill switch, and it narrows #1799's gate rather than leaving both mechanisms in place: `MLXCEL_SDPA_FALLBACK_MAX_QUERIES` no longer claims a call whose key can be bucketed, so what it still covers is a causal-mode block with no array mask, where there is no mask to widen. `MLXCEL_SDPA_PLAN_DEBUG=1` traces the key fields and plan builds per call. A cuDNN that cannot plan the widened shape degrades to per-length plans instead of aborting, per sinks setting so one refusal cannot disable the other shape. Refs #1820
Measured on GB10 against main (`60341873`, which carries #1817). What the record establishes: the three shape classes and exactly which key fields move round over round, measured per call rather than inferred; three plan builds per verify round becoming zero after warm-up (82 builds to 11 in a 25-round generation); byte-identical token ids between bucketed and exact-shape cuDNN on both binaries; a non-speculative head_dim-128 model that does not enter the new path at all; and a plan-cache growth bound of shape classes times query widths times buckets crossed rather than rounds. What it does not establish: that bucketing beats #1817's ops fallback. On the 152-token prompt they are a wash at widths 2 and 4 and bucketing is 3 to 6% behind at 6, 8 and 16, so the roughly 10 ms per round #1799 attributed to the fallback is not recovered at those key lengths. The long-context sweep that would have tested #1799's prediction (the fallback's score matrix grows with the key length, cuDNN's flash does not) was stopped partway when a second session loaded a model on the same GPU and the host's NVRM failure count went from 0 to 356 in two minutes. Those 10 rows are kept in `laguna_long_cli.jsonl` and marked unusable rather than quietly reported; the clean sweep finished before the foreign process existed. Refs #1820
The long-context ladder sweep stopped at 31 rows on 2026-09-12 14:35 when its session ended, and the host was later rebooted, so these results were at risk of being lost uncommitted. This commit preserves them verbatim along with the harness changes that produced them: the 8k, 16k and 32k code prompts, the ladder driver, and the host-gate and CLI harness edits. The data is committed as raw evidence only. Completeness per configuration and whether each row ran on a clean host have not been reviewed yet, and nothing here is cited as a result until that review is done. Refs #1820
An audit of the #1820 overlay against its Rust mirror found that the mirror had collapsed two axes the C++ gate keeps apart. `cuda_sdpa_plan_bucket_eligible` took a single `masked` flag, and the maskless causal call site in `fast_scaled_dot_product_attention_causal_windowed` passes `true` for it. That was correct before #1820, when the #1799 gate fired for `has_arr_mask || do_causal` alike, and wrong after: bucketing needs an array mask to widen, so C++ rejects a causal block at `do_causal` and still routes it to MLX's materializing ops fallback, while the mirror now claimed it stayed on cuDNN. The score-matrix query chunking was therefore switched off for exactly the calls that still materialize a `[B, heads, q_len, k_len]` matrix, and `cuda_sdpa_small_query_fallback` had become dead code on CUDA builds. Both predicates now take `arr_masked` and `do_causal` separately, both call sites pass what they mean, and the unit tests assert the causal axis rather than re-asserting the array-mask one. Four fixes in the overlay itself, all found by the same audit: `kv_cache_slice_extent` now requires `offset() == 0`. Its element-count identity proves the widened view is the same size as the allocation, not that it starts at it, so a mid-buffer window slice passed every test and `unslice_kv` would then subtract a non-zero byte offset from the base pointer. No mlxcel producer hands SDPA such a slice today, so this guards a future one. The bucket try/catch no longer wraps `sdpa_cache().emplace`. The LRU's lifetime-miss "Cache thrashing" throw is the exact abort this bucketing exists to prevent, and catching it as a cuDNN refusal would have disabled bucketing permanently and then hit the abort again, uncaught, on the exact-shape retry. A shape-eligible call whose layout declines bucketing now warns once. `supports_sdpa_cudnn` narrows #1799's fallback on the shape alone because it must answer identically during graph building and at eval, so a call that then declines on layout (the by-copy size cap, a mask that is not the broadcast plane) gets neither fix and rebuilds a plan every round. That hole is real and was silent. The unslice arm now takes `k_extent >= k_len` rather than `>`, so an exactly-full cache unslices as a no-op view instead of copying onto a buffer of identical size. `docs/environment-variables.md` said a `RotatingKVCache` multi-row append is not bucketed. The by-copy arm does bucket it; corrected, and `MLXCEL_SDPA_PLAN_BUCKET_MAX_MB` now has a row of its own. Refs #1820
…rebase The branch was rebased onto `0ef0a1a4`, 30 commits of main, which in that window gained f32 attention work touching the attention path, so the short-context table the branch already carried had to be re-measured rather than assumed to survive. It does not fully survive, and this records the new run rather than editing the old one. Every arm is 2 to 7% below the committed table, on both sides, and the fallback arm moves as much as the bucketed one, so most of that is a session offset: load1 ran 1.00 to 3.14 against 0.72 to 2.29 before, another session held `rustc` and `cargo-clippy` on the box throughout, and six runs had a CI job alive. The comparison itself also moved. Bucketing was a wash at widths 2 and 4 and 3 to 6% behind at 6, 8 and 16; it is now 0.94x at width 2 with disjoint ranges, 1.03x at 4, and 0.92x to 0.95x at 6, 8 and 16. So the conclusion the committed table drew is unchanged and slightly stronger: at a 152-token prompt bucketing does not beat #1799's ops fallback. One term did not move with the session. The fallback arm's drafter host build is flat against the committed run (7.4 to 7.6, 9.1 to 9.2, 10.1 to 10.3 ms per round at widths 2, 4 and 16), while the bucketed arm's rose about 3 ms at every width (7.3 to 10.5, 9.4 to 13.4, 13.0 to 16.2). That term is host kernel enqueues, which is what the widening adds and what CPU contention hits hardest, so contention and a real change are not separable from this run alone. It is recorded as unattributed rather than explained away. The 10 rows of a first attempt at this sweep are not included: the harness was killed partway when the host ran low on memory, and they were taken while that pressure was building. Refs #1820
First rung of the rebased context ladder, committed as it completes rather than at the end of the sweep, so a driver spike on a later rung cannot take the earlier evidence with it. At 2634 prompt tokens the two fixes are a wash, which is what the pre-rebase run found: bucketed against #1799's ops fallback is 1.017x at width 4 and 1.011x at width 8 on throughput, with per-round verify time overlapping in both. n = 3 per arm, no foreign model process on the host, no CI job, load1 1.24 to 3.81. The table carries two metrics because only one of them is a kernel measurement. The arms run different kernels, and #1799 already recorded that the ops fallback shifts the greedy path at ties, so they generate different text and need a different number of verify rounds to reach the same 150 tokens (64 against 66 at width 4). Throughput therefore mixes kernel speed with acceptance luck; verify ms per round divides that out and is the number to read for a crossover. Refs #1820
The context ladder was halted partway through its 8k rung when the driver trip wire fired: 107 `NVRM: NV_ERR_NO_MEMORY` errors in a five-minute window against a pre-run baseline of 2 that had been stable for three days. Only this unit's own processes were on the GPU, checked by process cwd rather than ancestry, so the errors are this workload's own, not a foreign session's. These six non-warmup rows are committed as raw evidence and are cited for nothing. They are incomplete (n = 1 to 2 per configuration against the n = 3 the record requires) and partly contaminated: the second `b8` row reports a drafter host build of 24.8 ms per round against about 12 ms in every other row here, which is the signature of the allocation storm rather than of the arm under test. Both `b8` rows sit well under the 27.63 tok/s the pre-rebase run measured for that configuration. What the kernel log shows about the burst is worth keeping, because it changes how the halt should be read. The 107 errors are not a climb: 96 of them land inside one second at 16:12:02, on top of 6 at 16:11:56 and 4 at 16:11:25, and the count returns to zero for the minute after. The earlier burst has the same shape, 27 errors inside one second at 15:38:21. So this is a discrete allocation storm repeating at intervals, not the compounding accumulation that preceded the 2026-07-06 hard freezes, and it is an order of magnitude below the 2026-09-12 halt (234 errors in two minutes). The 2634-token rung completed at n = 3 and is already committed. 16k and 32k are unrun. Refs #1820
Three preconditions added to the measurement gate, each because a run was actually lost to the thing it now refuses, on a shared self-hosted host where peer sessions and CI compete for the same box. Driver. `NV_ERR_NO_MEMORY` is now checked before every run rather than by eye between rungs, on two independent limits: at least 50 in the trailing five minutes, or 400 cumulative since boot. Cumulative is the primary gauge and is deliberately not waitable, because it never decays, so a run that trips it stops rather than waits. That asymmetry is the lesson from the 2026-09-17 halt, where a per-second burst shape was read as reassurance while the total went 2 to 77 to 184; the instantaneous rate says nothing about the documented freeze precursor, which is cumulative over hours. Each row now carries `nvrm_window_before`, `nvrm_total_before` and `nvrm_total_after`, so a row's driver conditions are auditable from the file instead of reconstructed from whoever was watching. Memory. Three concurrent foreign builds drove `MemAvailable` to zero and the supervisor killed background tasks three times; a model run started in that state dies mid-row, which is worse than a bad number because a bad number at least shows up in the spread. The floor is 32 GiB. It is derived rather than measured and is commented as such: 20.97 GiB of weights measured on disk (GB10 memory is unified, so weights count against system RAM) plus 0.69 GiB of KV at 16k computed from `config.json`, about 21.7 GiB resident, plus headroom for prefill-chunk activations and cuDNN workspace. Rows now record `peak_rss_kib` sampled from `VmHWM` so the floor can be replaced with the measured figure. Sustained quiet, and CI. Two samples five seconds apart was not enough: a lull between two queued CI runs clears a ten-second check, and the timed run that starts in it is overrun when the next job ramps. Quiet is now a minute of continuous quiet, and an active `Runner.Worker` gates rather than merely being recorded, because a CI job between compiler invocations is not an idle host, it is a host about to compile again. No measurement policy changed: same binary per arm, same warmup, same round-robin, same repeat count. These only decide when a run is allowed to start. Refs #1820
The gate watched for compilers and read a frontend build as silence. A Tauri build bundles a node toolchain on top of its Rust link step, and a CI job's own artifact step can be pure node, and both saturate this box exactly as `rustc` does. On 2026-09-17 a compiler-only predicate reported the host clear while a `WebUI installed artifact` step was still running and the `cargo-clippy` job it was gating on had not even started. This is the same failure the sustained-quiet change fixed one commit earlier, in a different costume: there the check was too short, here it was too narrow. Both let a timed run start next to work that invalidates it. `node`, `npm`, `npx`, `yarn`, `pnpm`, `bun`, `tsc`, `esbuild`, `vite`, `rollup` and `webpack` now gate alongside the compilers. A node process was live on the host when this was written, so the gap was load-bearing rather than theoretical. Refs #1820
The 8k and 16k rows sat in the same file as the clean 2634 rung with clean `foreign_models_*` fields, so every filter the summarizer had would have passed them through and averaged them in. They are now in `laguna_ladder_excluded.jsonl`, each carrying an `excluded_reason` naming what is wrong with it, and `summarize_ladder.py` drops any row carrying that field ahead of every other signal. Structural rather than annotative: the clean file yields only the 2634 rung, and the excluded file refuses its own contents. The 8k rung was halted mid-run under a driver allocation storm, and is contaminated as well as sparse: one b8 row reports a 24.8 ms per round drafter host build against about 12 ms in its siblings, which is the storm and not the arm under test. The 16k rung was stopped after two runs. Each 16k run cost about 38 `NVRM: NV_ERR_NO_MEMORY` allocation failures against near-zero at 2634, so the 13-run rung projected to roughly 678 cumulative against a 400 budget and could not fit. That step change with prompt length is a measured property of this host and is reported as a result rather than as an incident. Two more guards while the summarizer was open. A row whose run moved the cumulative driver count by more than 20 is surfaced as suspect but kept, since that threshold is judgement rather than fact. Any configuration under n = 3 is labelled UNDERPOWERED in the table itself, so a sparse arm cannot read as a measured one at a glance. Refs #1820
Bucketing the cuDNN SDPA plan-cache key works and should not ship. Against the bar that a default-on change must not lose on any measured workload, it fails on the workload it was built for: at a 152-token prompt it is behind #1799's ops fallback at three of five block widths with non-overlapping ranges (0.943x, 0.953x, 0.915x at widths 2, 8 and 16), and at 2634 tokens it is a wash on both throughput and per-round verify time. There is no measured context length at which it wins, so the roughly 10 ms per round #1799 attributed to the fallback is not recovered anywhere it was measured. `MLXCEL_SDPA_FALLBACK_MAX_QUERIES` therefore stays exactly as #1817 shipped it, and the implementation stays on the branch behind its own kill switch rather than being removed. What bucketing does achieve is recorded rather than dismissed: 12 plan builds in a generation instead of 82, zero per verify round after warm-up, and greedy token ids byte-identical to the exact-shape cuDNN path on both binaries across three repeats. The mechanism is sound; it is the throughput that does not justify a default. The crossover question is recorded as unmeasured on this host rather than absent, with the mechanism that motivates it: only 10 of the pairing's 45 attention calls scale with prompt length, since the target's 30 sliding layers and all 5 drafter layers are capped at a 512 window, so bucketing's cost is flat in context while the fallback's score matrix grows linearly. Answering it would produce a context-gated ship, not an unconditional one, and would then need the short-context loss re-confirmed before a gate could be designed. Two host findings are recorded as first-class results. A 16k-token run costs about 38 `NVRM: NV_ERR_NO_MEMORY` allocation failures against near-zero at 2634, which projected the deciding 13-run rung to roughly 678 cumulative against a 400 budget and is why it could not be run; a follow-up needs a fresh boot. And `peak_rss_kib` reports 1.9 GiB for a run whose weights alone are 20.97 GiB, because CUDA unified allocations on GB10 do not enter the resident set, so that field is marked as not a valid footprint proxy rather than left as a plausible number someone trusts later. The 2026-09-12 record is marked superseded in place: its mechanism findings hold, its throughput table does not. Refs #1820
#1820 built this path and then measured it losing. At a 152-token prompt it is behind #1799's ops fallback at three of five block widths with non-overlapping ranges (0.943x, 0.953x and 0.915x at widths 2, 8 and 16), and at 2634 tokens it is a wash on both throughput and per-round verify time. Against the bar that a default-on change must not lose on any measured workload, that disqualifies it, so `MLXCEL_SDPA_PLAN_BUCKET_MAX_QUERIES` now defaults to `0` on both sides of the bridge and the shipped dispatch is #1799's alone. With the bound at 0 nothing is bucket-eligible, so the #1799 gate reclaims every array-masked verify block it claimed before this branch existed: the default reproduces main's dispatch exactly. `plan_bucketing_is_off_by_default_so_1799_dispatch_ships` pins that, so a future default change cannot alter shipped dispatch silently. The code stays rather than being deleted, and the C++ comment says why so it is read as neither dead code nor a soft endorsement: the mechanism is sound (12 plan builds per generation instead of 82, greedy token ids byte-identical to exact-shape cuDNN across three repeats), and a crossover above 2634 keys is plausible and unmeasured, because this path's cost is flat in context while the fallback materializes a score matrix that grows with the key length. A long-context re-measurement on a freshly booted host should be a flag flip, not a reimplementation. The eligibility predicate is split into `plan_bucket_eligible_with`, which takes the bound as an argument. The old test asserted bucketing was eligible, which held only while the default was on, and a `OnceLock` read of the environment cannot be varied per test. `docs/environment-variables.md` is corrected on both rows. The `MLXCEL_SDPA_FALLBACK_MAX_QUERIES` row claimed #1820 had narrowed it; that is true only when bucketing is enabled, and the shipped default leaves that gate's behaviour unchanged. Refs #1820
Two cleanups on the data directory, which was 1.7 MB against 336 KB for the merged #1798 precedent at the same file count. The three raw traces were 1.2 MB of it and 88 to 97 percent redundant (the non-speculative control is 3444 lines with 125 distinct). They are replaced by frequency extracts: one row per distinct line as the count, a tab, then the line. `sort -u` would have been wrong here, because the call totals the records cite are exactly those occurrence counts. What survives is every count, distinct-value count and maximum; what does not is per-call ordering, and no figure in either record depends on it. `harness/extract_trace.py` generates them and `summarize_trace.py` now reads either format and reports which it was given, using `max(plans)` rather than the last record for the resident plan count, which is equivalent because that counter is monotonic within a run. Every number the records cite was re-derived from the extracts rather than assumed to survive: 82 plan builds upstream against 12 rebased (and 11 pre-rebase); the three shape classes at 25, 25 and 24 builds over 25 rounds, which is the three-per-round figure, against 12 in the whole rebased run; 1204 non-speculative calls with zero bucketed and three builds; cache extents 256 and 768 with `copy=0` on both target arms and `copy=1` on the drafter; and prefill chunks of 2048 and 586 query rows. Doing that surfaced a gap this commit also closes. The 2026-09-17 record cited 12 plan builds and 1204 non-speculative calls from rebased runs that had never been committed; the committed traces were the pre-rebase ones, at 11 and 3444. Extracts of both rebased runs are added, so that record's own numbers now have evidence in the tree. Second, `laguna_ladder_excluded.jsonl` was referenced by nothing. That is the defect rather than an oversight: the quarantined rows have clean foreign-process fields, so every automatic filter passes them, and an uncited file beside the clean data cannot be told from a stray one. Both records and both reports now name it, say it holds the halted 8k rung and the two 16k rows, say why each is disqualified, and say it is retained as raw evidence and cited for nothing. Directory is 1.7 MB to 692 KB. Refs #1820
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.
What this asks for
A decision, not a merge-and-forget. This branch implements #1820's plan-cache bucketing, measures it, and recommends against enabling it. It lands the implementation default-off behind
MLXCEL_SDPA_PLAN_BUCKET_MAX_QUERIES=0so the shipped dispatch is unchanged from main, and lands the evidence that produced that recommendation.The result: bucketing loses to #1799's ops fallback where it was measured
Laguna DFlash on GB10, 152-token prompt, 200 tokens, n = 3, same binary on every arm, widths interleaved round-robin.
fb-is main's dispatch.Three of five widths are losses with non-overlapping ranges. At 2634 tokens the two are a wash (1.017x and 1.011x, ranges overlapping on both throughput and per-round verify time). There is no measured context length at which bucketing wins, so the roughly 10 ms per round #1799 attributed to the fallback is not recovered anywhere it was measured. Against the bar that a default-on change must not lose on any measured workload, that disqualifies it.
MLXCEL_SDPA_FALLBACK_MAX_QUERIEStherefore stays exactly as #1817 shipped it.Why the code lands anyway instead of being deleted
Not an oversight, and not a soft endorsement. The mechanism is sound, and the question it was built to answer is still open and still answerable:
[B, heads, q_len, k_len]score matrix that grows with the key length.A long-context re-measurement should be a flag flip, not a reimplementation. The motivation is in the C++ comment at the default so a later reader does not delete it as dead code.
Why the deciding measurement is missing
A 16k-token run costs about 38
NVRM: NV_ERR_NO_MEMORYallocation failures on this host, against near-zero at 2634. The 13-run rung projected to roughly 678 cumulative against a 400 budget and was stopped after two runs at 222, from a pre-session baseline of 2 that had been stable for three days. The runs complete successfully with valid throughput, which is not evidence the failures are benign: that is the shape the 2026-07-06 hard freezes took here. The follow-up needs a fresh boot, whose clean driver count is the only condition under which the rung fits the budget. Not filed; raised in the report for a decision.Defects fixed
A live regression this branch introduced:
cuda_sdpa_plan_bucket_eligiblecollapsed the array-mask anddo_causalaxes, which the C++ gate treats differently once bucketing exists. The maskless causal call site passestrue, so the mirror claimed a causal block stayed on cuDNN while C++ routed it to the materializing fallback, silently dropping score-matrix query chunking. Checked against main: not a defect there, because main's gate is(has_arr_mask || do_causal)and both axes behave identically, so it is branch-local and needs no separate PR.Four overlay fixes:
kv_cache_slice_extentrequiresoffset() == 0; the buckettry/catchno longer wrapssdpa_cache().emplace, whose LRU thrashing throw is the exact abort this work exists to prevent; a shape-eligible call whose layout declines bucketing warns once, since it gets neither fix; and the unslice arm takes>=so an exactly-full cache does not copy onto a buffer of its own size.Test plan
cargo fmt --all -- --checkcargo clippy --profile test-fast --features cuda --lib --bins --tests -- -D warningscargo test --profile test-fast --features cuda -p mlxcel-core layers:: -- --test-threads=1: 92 passed, 0 failed7d044d8534201cb4, three repeats)qwen3-1.7b-4bit: 200 ids identical base against branch, and zero bucketed calls in 1204 traced cuDNN SDPA callsmetal,accelerategate: not runnable on this Linux/CUDA host, unrunRecord:
docs/benchmark_results/sdpa-plan-cache-bucket-gb10-2026-09-17.md. Report:TECHNICAL_REPORTS/1820-sdpa-plan-cache-bucket-20260917.{en,ko}.md.Closes #1820