Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
59 changes: 59 additions & 0 deletions TECHNICAL_REPORTS/1820-sdpa-plan-cache-bucket-20260917.en.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,59 @@
# Bucketing the cuDNN SDPA plan-cache key, and why it should not ship

Issue #1820. Host GB10 (sm_121), MLX pin `81ba1c6a`, CUDA release build, toolchain 1.97.1, base tree `0ef0a1a4` (main, carrying PR #1817). Full measurement record: `docs/benchmark_results/sdpa-plan-cache-bucket-gb10-2026-09-17.md`.

## Background

MLX keys its cuDNN SDPA execution-plan cache on the exact shapes and strides of q, k, v and the mask. A speculative verify round appends keys every round, so its key is new every round and the cache structurally never hits. #1799 measured the consequence on the Laguna DFlash pairing: three shape classes each rebuilding a plan per round at about 22 ms, 67 to 76 ms of host time per round against 2.8 ms per classic decode token, and the LRU's lifetime-miss counter then aborting the process past about 170 rounds.

PR #1817 fixed that by routing these calls off cuDNN entirely, onto MLX's ops fallback, which has no per-shape build. It works, and it pays for it: the fallback materializes a `[B, heads, q_len, k_len]` score matrix, about 10 ms more GPU time per round at block 2. #1820 asked whether making the key reusable would recover that 10 ms while keeping cuDNN's flash kernel.

## What was built

MLX's own one-row decode canonicalization, extended to a small array-masked multi-row call. k and v are widened to a bucket, the additive 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, and by a zero-padded copy when they are not.

It works mechanically. The same generation builds 12 plans instead of 82, one per shape class rather than three per round, and greedy token ids are byte-identical to the exact-shape cuDNN path on both binaries.

## The recommendation: do not ship it

Against the standing bar that a default-on change must not lose on any measured workload, bucketing fails on the workload it was designed for.

| width | bucketed | #1817 fallback | ratio | ranges |
|---|---|---|---|---|
| 2 | 33.10 (32.43 to 33.77) | 35.12 (34.65 to 35.79) | **0.943x** | disjoint |
| 4 | 37.82 (37.40 to 38.29) | 36.64 (32.88 to 38.68) | 1.032x | overlapping |
| 6 | 35.76 (34.48 to 36.80) | 37.54 (36.70 to 38.21) | 0.953x | overlapping |
| 8 | 31.15 (30.64 to 31.51) | 32.68 (32.36 to 33.04) | **0.953x** | disjoint |
| 16 | 22.21 (21.89 to 22.43) | 24.26 (23.91 to 24.45) | **0.915x** | disjoint |

152-token prompt, n = 3, same binary on every arm, widths interleaved. 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 10 ms #1799 attributed to the fallback is not recovered anywhere it was measured.

A measured loss disqualifies it as a default regardless of what the unmeasured rungs would have shown. The recommendation does not rest on the rung that could not be run.

`MLXCEL_SDPA_FALLBACK_MAX_QUERIES` therefore stays as #1817 shipped it. The implementation is left on the branch behind `MLXCEL_SDPA_PLAN_BUCKET_MAX_QUERIES`, not removed, because the open question below is answerable and the implementation is the expensive part of answering it.

## What remains open

The crossover is **unmeasured on this host in its current state, not absent**, and the mechanism that motivates it is real. Only 10 of the pairing's 45 attention calls scale with prompt length: the target's full-attention layers, which take the free unslice arm. Its 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, and a crossover should exist somewhere above 2634 tokens.

Finding it would not produce an unconditional ship. It would produce a context-gated one, which then needs a short-context retry to confirm the loss is real before a gate could be designed, so the remaining question is larger than one rung.

## Why the deciding rung could not be run

A 16k-token run costs about 38 `NVRM: NV_ERR_NO_MEMORY` allocation failures on this host, against near zero at 2634. The 13-run rung projected to roughly 678 cumulative failures against a 400 budget, and was stopped after two runs with the count at 222 from a pre-session baseline of 2.

The runs complete successfully with valid throughput. That is not evidence the failures are benign: cumulative accumulation under spiky delivery, with individual runs looking fine, is the shape the 2026-07-06 hard freezes took on this host. A follow-up needs a fresh boot, whose clean driver count is the only condition under which the rung fits inside the budget.

## Defects found and fixed along the way

An audit of the overlay against its Rust mirror found a live regression this branch introduced. `cuda_sdpa_plan_bucket_eligible` took a single `masked` flag, but the C++ gate treats an array mask and `do_causal` differently: bucketing needs a mask to widen, so it rejects a causal block, while #1799's fallback fires on either. The maskless causal call site passes `true` for that flag, so the mirror claimed a causal block stayed on cuDNN while C++ actually routed it to the materializing fallback, silently dropping score-matrix query chunking from exactly those calls. Both predicates now take `arr_masked` and `do_causal` separately.

Four overlay fixes: `kv_cache_slice_extent` now requires `offset() == 0`, since its element-count identity proves the widened view is the same size as the allocation but not that it starts at it; the bucket `try`/`catch` no longer wraps `sdpa_cache().emplace`, whose LRU thrashing throw is the exact abort this work exists to prevent and would have been misread as a cuDNN refusal; a shape-eligible call whose layout declines bucketing now warns once, because such a call gets neither fix and rebuilds a plan every round; and the unslice arm takes `>=` rather than `>` so an exactly-full cache does not copy onto a buffer of its own size.

## Verification

Byte identity, compared as token ids: `7d044d8534201cb4` across base-binary cuDNN, branch-binary cuDNN and branch-binary bucketed, three repeats. `baa2c6b55d7f874d` on both binaries' ops-fallback arms, a dispatch difference #1799 already recorded, not a padding error. Non-speculative control `qwen3-1.7b-4bit`: 200 ids identical base against branch, and a trace showing zero bucketed calls in 1204 cuDNN SDPA calls.

Data layout, because two files there must not be read as results. `laguna_ladder_excluded.jsonl` holds the rows that were measured and then disqualified, and is cited for nothing: the halted 8k rung (seven rows, sparse at n = 1 to 2 and contaminated by a driver allocation storm, one row reporting a 24.8 ms per round drafter host build against about 12 ms in its siblings) and the two 16k rows from the rung that was stopped when each run proved to cost about 38 NVRM allocation failures. It is a separate file because no automatic filter can distinguish those rows: their foreign-process fields are clean, so every contamination check passes them. Each row carries an `excluded_reason`, and `summarize_ladder.py` drops such rows ahead of every other signal. The traces are frequency extracts rather than raw, one row per distinct line with its occurrence count, which preserves every count, distinct-value count and maximum the record cites while discarding per-call ordering that nothing cites.

Not verified: the `metal,accelerate` gate is not runnable on this Linux/CUDA host and was not run. Attention sinks are unexercised, since this checkpoint carries none. `peak_rss_kib` in the harness reports 1.9 GiB for a run whose weights alone are 20.97 GiB, because CUDA unified allocations on GB10 do not appear in the resident set; the field is recorded but is not a valid footprint proxy, and the harness's memory floor stays derived rather than measured.
59 changes: 59 additions & 0 deletions TECHNICAL_REPORTS/1820-sdpa-plan-cache-bucket-20260917.ko.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,59 @@
# cuDNN SDPA 플랜 캐시 키 버킷화, 그리고 적용하지 않기로 한 이유

이슈 #1820. 호스트 GB10 (sm_121), MLX 핀 `81ba1c6a`, CUDA 릴리스 빌드, 툴체인 1.97.1, 기준 트리 `0ef0a1a4` (PR #1817을 포함한 main). 전체 측정 기록: `docs/benchmark_results/sdpa-plan-cache-bucket-gb10-2026-09-17.md`.

## 배경

MLX는 cuDNN SDPA 실행 플랜 캐시를 q, k, v와 마스크의 정확한 shape 및 stride로 키잉한다. 추측 실행(speculative)의 verify 라운드는 매 라운드 키를 덧붙이므로 키가 매번 새로 생기고, 캐시는 구조적으로 절대 적중하지 않는다. #1799는 Laguna DFlash 조합에서 그 결과를 측정했다. 세 개의 shape 클래스가 각각 매 라운드 약 22 ms짜리 플랜을 다시 만들었고, 라운드당 호스트 시간이 67~76 ms로 클래식 디코드 토큰당 2.8 ms와 대비되었으며, LRU의 수명 미스 카운터가 약 170 라운드를 넘기면 프로세스를 중단시켰다.

PR #1817은 이 호출들을 cuDNN에서 완전히 우회시켜 shape별 빌드가 없는 MLX 자체 ops 폴백으로 보내 문제를 해결했다. 동작하지만 대가가 있다. 폴백은 `[B, heads, q_len, k_len]` 점수 행렬을 실체화하므로 블록 2에서 라운드당 GPU 시간이 약 10 ms 더 든다. #1820은 키를 재사용 가능하게 만들면 cuDNN의 flash 커널을 유지한 채 그 10 ms를 회수할 수 있는지 물었다.

## 구현한 것

MLX 자체의 1행 디코드 정규화를 배열 마스크가 붙은 소규모 다중 행 호출로 확장했다. k와 v를 버킷 크기로 넓히고, 덧셈 마스크도 같은 너비로 넓히되 새로 생긴 열은 `-inf`로 채우며, 실제 길이는 `set_padding_mask`와 `set_seq_len_q` / `set_seq_len_kv`를 통해 cuDNN에 전달한다. 넓힌 영역이 결과에 반영되지 않게 하는 것이 바로 이 부분이다. k와 v는 하나의 캐시 버퍼의 선행 슬라이스이고 여유 공간이 있으면 unslice로, 그렇지 않으면 0으로 채운 복사로 버킷에 도달한다.

기계적으로는 동작한다. 동일한 생성에서 플랜 빌드가 82회에서 12회로 줄어 라운드당 3회가 아니라 shape 클래스당 1회가 되며, greedy 토큰 id는 두 바이너리 모두에서 정확 shape cuDNN 경로와 바이트 단위로 동일하다.

## 결론: 적용하지 않는다

기본 활성화 변경은 측정된 어떤 워크로드에서도 손해를 보면 안 된다는 기준에 대해, 버킷화는 바로 그 변경이 겨냥한 워크로드에서 실패한다.

| 폭 | 버킷화 | #1817 폴백 | 비율 | 범위 |
|---|---|---|---|---|
| 2 | 33.10 (32.43 ~ 33.77) | 35.12 (34.65 ~ 35.79) | **0.943x** | 분리됨 |
| 4 | 37.82 (37.40 ~ 38.29) | 36.64 (32.88 ~ 38.68) | 1.032x | 겹침 |
| 6 | 35.76 (34.48 ~ 36.80) | 37.54 (36.70 ~ 38.21) | 0.953x | 겹침 |
| 8 | 31.15 (30.64 ~ 31.51) | 32.68 (32.36 ~ 33.04) | **0.953x** | 분리됨 |
| 16 | 22.21 (21.89 ~ 22.43) | 24.26 (23.91 ~ 24.45) | **0.915x** | 분리됨 |

152 토큰 프롬프트, n = 3, 모든 arm이 동일 바이너리, 폭은 교차 실행. 다섯 폭 중 세 곳이 범위가 겹치지 않는 손해다. 2634 토큰에서는 양쪽이 무승부다 (1.017x, 1.011x이며 처리량과 라운드당 verify 시간 모두 범위가 겹친다). 버킷화가 이기는 컨텍스트 길이는 측정 범위 안에 없으므로, #1799가 폴백에 귀속시킨 10 ms는 측정한 어느 지점에서도 회수되지 않는다.

측정된 손해가 있으면 측정하지 못한 구간에서 무엇이 나왔을지와 무관하게 기본값으로는 부적격이다. 이 결론은 실행하지 못한 구간에 기대지 않는다.

따라서 `MLXCEL_SDPA_FALLBACK_MAX_QUERIES`는 #1817이 내보낸 상태 그대로 둔다. 구현은 제거하지 않고 `MLXCEL_SDPA_PLAN_BUCKET_MAX_QUERIES` 뒤에 브랜치에 남겨둔다. 아래의 미해결 질문은 답할 수 있는 질문이고, 구현이 그 답을 구하는 데 드는 비용의 대부분이기 때문이다.

## 남은 미해결 질문

교차점은 **부재한 것이 아니라 현재 상태의 이 호스트에서 측정되지 않은 것**이며, 그것을 뒷받침하는 메커니즘은 실재한다. 이 조합의 어텐션 호출 45개 중 프롬프트 길이에 따라 커지는 것은 10개뿐이다. 타깃의 full attention 레이어이고, 이들은 비용이 없는 unslice 경로를 탄다. 나머지 sliding 레이어 30개와 드래프터 레이어 5개 전부는 512 윈도로 상한이 걸려 있다. 즉 버킷화의 비용은 컨텍스트에 대해 평평한 반면 폴백의 점수 행렬은 선형으로 커지므로, 2634 토큰 위 어딘가에 교차점이 존재해야 한다.

교차점을 찾는다고 해서 무조건 적용으로 이어지지는 않는다. 컨텍스트 기반 게이팅이 필요해지고, 게이트를 설계하려면 먼저 짧은 컨텍스트에서의 손해가 실재함을 재측정으로 확인해야 한다. 남은 질문은 한 구간보다 크다.

## 결정적 구간을 실행하지 못한 이유

이 호스트에서 16k 토큰 실행 1회는 `NVRM: NV_ERR_NO_MEMORY` 할당 실패를 약 38회 발생시킨다. 2634 토큰에서는 거의 0이다. 13회 실행 구간은 누적 약 678회로 예측되어 400 예산을 초과했고, 세션 전 기준값 2에서 시작해 222에 도달한 시점에 2회 실행 후 중단했다.

실행 자체는 정상 종료되고 처리량도 유효하다. 그것이 이 실패가 무해하다는 근거는 아니다. 개별 실행은 멀쩡해 보이면서 누적이 스파이크 형태로 쌓이는 양상은 2026-07-06에 이 호스트에서 발생한 하드 프리즈와 같은 모양이다. 후속 작업에는 재부팅이 필요하다. 드라이버 카운트가 깨끗한 상태여야만 이 구간이 예산 안에 들어간다.

## 과정에서 발견해 수정한 결함

오버레이를 Rust 미러와 대조 감사하는 과정에서 이 브랜치가 유입시킨 실제 회귀를 발견했다. `cuda_sdpa_plan_bucket_eligible`이 단일 `masked` 플래그를 받았는데, C++ 게이트는 배열 마스크와 `do_causal`을 다르게 취급한다. 버킷화는 넓힐 마스크가 필요하므로 causal 블록을 거부하는 반면, #1799의 폴백은 둘 중 어느 쪽이든 발동한다. 마스크 없는 causal 호출 지점은 그 플래그에 `true`를 넘기므로, 미러는 causal 블록이 cuDNN에 남는다고 판단했지만 C++는 실제로는 실체화 폴백으로 보내고 있었고, 결과적으로 바로 그 호출들에서 점수 행렬 쿼리 청킹이 조용히 꺼졌다. 이제 두 술어 모두 `arr_masked`와 `do_causal`을 따로 받는다.

오버레이 수정 네 건. `kv_cache_slice_extent`는 이제 `offset() == 0`을 요구한다. 원소 개수 항등식은 넓힌 뷰가 할당과 같은 크기임은 증명하지만 할당의 시작점에서 출발한다는 것은 증명하지 못하기 때문이다. 버킷 `try`/`catch`는 더 이상 `sdpa_cache().emplace`를 감싸지 않는다. 그 LRU 스래싱 예외는 이 작업이 막으려는 바로 그 중단이며 cuDNN 거부로 오독되었을 것이다. shape 조건은 만족하지만 레이아웃 때문에 버킷화가 거부된 호출은 이제 한 번 경고한다. 그런 호출은 두 수정 중 어느 것도 적용받지 못한 채 매 라운드 플랜을 다시 만들기 때문이다. unslice 경로는 `>` 대신 `>=`를 쓰게 하여 정확히 가득 찬 캐시가 같은 크기의 버퍼로 복사되지 않게 했다.

## 검증

바이트 동일성, 토큰 id 기준 비교: 기준 바이너리 cuDNN, 브랜치 바이너리 cuDNN, 브랜치 바이너리 버킷화 모두 `7d044d8534201cb4`이며 3회 반복. 두 바이너리의 ops 폴백 arm은 모두 `baa2c6b55d7f874d`로, 이는 #1799가 이미 기록한 디스패치 차이이지 패딩 오류가 아니다. 비추측 대조군 `qwen3-1.7b-4bit`: 기준 대비 브랜치에서 200개 id 동일, 그리고 cuDNN SDPA 호출 1204건 중 버킷화된 호출이 0건임을 트레이스로 확인.

데이터 배치에 관한 주의. 결과로 읽으면 안 되는 파일이 두 묶음 있다. `laguna_ladder_excluded.jsonl`은 측정했으나 자격을 상실한 행들을 보관하며 어디에서도 인용하지 않는다. 드라이버 할당 폭주 중 중단된 8k 구간 7행 (구성당 n = 1~2로 희소하고, 그와 별개로 오염되어 있다. 한 행이 라운드당 드래프터 호스트 빌드를 24.8 ms로 보고하는데 형제 행들은 약 12 ms다) 과, 실행 1회당 NVRM 할당 실패가 약 38회임이 드러나 중단한 16k 구간 2행이다. 별도 파일로 둔 이유는 자동 필터로는 이 행들을 구분할 수 없기 때문이다. 외부 프로세스 필드가 깨끗해서 모든 오염 검사를 통과한다. 각 행은 `excluded_reason`을 담고 있고 `summarize_ladder.py`는 그 필드를 다른 어떤 신호보다 먼저 보고 해당 행을 제외한다. 트레이스는 원본이 아니라 빈도 추출본이다. 고유한 줄 하나당 출현 횟수를 붙인 형식으로, 기록이 인용하는 모든 합계와 고유값 개수, 최댓값은 보존하고 아무도 인용하지 않는 호출 순서만 버린다.

검증하지 못한 것: `metal,accelerate` 게이트는 이 Linux/CUDA 호스트에서 실행할 수 없어 실행하지 않았다. 이 체크포인트에는 어텐션 싱크가 없으므로 싱크 경로는 검증되지 않았다. 하네스의 `peak_rss_kib`는 가중치만 20.97 GiB인 실행에 대해 1.9 GiB를 보고하는데, GB10의 CUDA 통합 메모리 할당이 프로세스 resident set에 나타나지 않기 때문이다. 이 필드는 기록하되 유효한 메모리 사용량 지표가 아니며, 하네스의 메모리 하한값은 측정값이 아니라 유도값으로 남는다.
Loading
Loading