Repository navigation
perf(cuda): adaptive split-KV sizing for attention_row decode (+4-5% over fixed-256; +70% deep-ctx) - #1350
Merged
Conversation
…over fixed-256; +70% deep-ctx vs monolithic) Follow-up to #1340. The split-KV path fixed keys-per-split at chunk=256, so num_splits = cap/256. An ncu + E2E chunk sweep on V2-Lite (H200) at deep context showed that under-fills the machine: at cap~4096 the grid is (16 rows x 16 splits) = 256 blocks, waves/SM 0.48, warps active 11%. The decode split kernel is memory-latency bound, so it wants ~a full occupancy wave of resident CTAs, not one block per SM. Size num_splits adaptively from the fixed cap AND fixed row count to target ~ATTN_SPLIT_TARGET_BLOCKS (512 = ~one wave on 132 SMs), floored so no split covers fewer than ATTN_SPLIT_MIN_CHUNK (128) keys: target_splits = TARGET_BLOCKS.div_ceil(total_rows) num_splits = min(target_splits, cap.div_ceil(MIN_CHUNK), MAX_SPLITS) chunk = cap.div_ceil(num_splits) For V2-Lite's 16 decode rows this lands num_splits=32 (grid 512) at deep context and scales down cleanly at shallow context (cap alone caps the useful splits). The sizing arithmetic is factored into a pure `attention_split_geometry` helper with a unit test. Capture-safe: num_splits still derives ONLY from cap + row count, never live seqlen, so eager and capture make identical launch decisions. chunk stays >= MIN_CHUNK (128), so any live context <= 128 stays on the single-split bit-exact fast path (the golden 24-tok lock sits well under this). Set ONNX_GENAI_ATTN_SPLIT_CHUNK to pin a fixed chunk (reproduces pre-adaptive behaviour, for A/B). ncu deep-ctx (V2-Lite, ~2600-tok prompt) attention_split fill: monolithic grid 16 waves 0.02 warps 9.8% DRAM 0.8% fixed-256 grid 256 waves 0.48 warps 11% DRAM 34% adaptive grid 512 waves 0.97 warps 24% DRAM 48% Gates (H200, GPU4, pinned single-threaded): - V2-Lite golden 24-tok lock: byte-identical (single-split path). - V2-Lite 340-tok long-ctx lock: eager==capture holds (deep multi-split). - standard_attention capture/fp16/bf16 gpu-tests + 15 lib unit tests pass. - Dense qwen2.5-0.5b: byte-identical, split does not engage. - E2E medians-of-5: deep 46.37 -> 48.32 tok/s (+4.2% vs fixed-256, +70% vs monolithic 28.40); shallow 58.08 -> 60.82 (+4.7% vs fixed-256, +17.7% vs monolithic 51.69); short "Hello" capture neutral (127.77 -> 130.17). Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby
added a commit
that referenced
this pull request
Aug 19, 2026
Records the merged #1350 adaptive split-KV attention win (grid≈512 full-wave targeting, +63% deep-ctx over monolithic + adaptive +4.2%/+4.7%) in the campaign brain. now.md only. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
justinchuby
added a commit
that referenced
this pull request
Aug 19, 2026
`main` is red on the required `Rust quality` check again. #1350 (`32ee82cb5`) landed a test block in `crates/onnx-runtime-ep-cuda/src/kernels/standard_attention.rs` that rustfmt rewraps, so `cargo fmt --all -- --check` fails at line 3569 on `main` — and therefore on every open PR, none of which can go green until this lands. This is the rustfmt output for that block and nothing else: two `assert_eq!` calls in `attention_split_geometry`'s unit test get split across lines. No behaviour changes; the crate does not even build on the Linux lanes. Same root cause as #1346, which fixed the previous instance of this in the same file: the Actions queue is saturated (every recent run is `queued`, nothing `in_progress`), so post-merge CI never reports and formatting drift reaches `main` unobserved. Verified locally on rustc 1.97.1 — the toolchain CI resolves to: | step | on `main` (`530d9c3a3`) | on this branch | |---|---|---| | `cargo fmt --all -- --check` | **FAIL** (`standard_attention.rs:3569`) | PASS | | `verify_documented_env_vars.py` | PASS | PASS | Working as sebastian (CPU perf). Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
|
| Status | Scenario | Base | PR | Change |
|---|---|---|---|---|
gather/medium_f32_threads=1-internal/32768 |
5.85 µs | 7.25 µs | +23.8% | |
block_quantized_matmul_cached_dense/mxfp4_preexpanded_dense_oncelock_like_proxy/1x1024x1024 |
44.99 µs | 55.61 µs | +23.6% | |
tokenization/encode_tokens_per_second |
371.38 µs | 451.46 µs | +21.6% | |
tokenization/decode_tokens_per_second |
6.71 ms | 7.85 ms | +16.9% | |
| ✅ | add/medium_f16_threads=1-internal/262144 |
133.35 µs | 146.26 µs | +9.7% |
| ✅ | sampling_latency/top_p_per_token |
374.06 µs | 410.13 µs | +9.6% |
| ✅ | gather/medium_bf16_threads=1-internal/32768 |
2.43 µs | 2.59 µs | +6.4% |
| ✅ | sampling_latency/greedy_per_token |
3.40 µs | 3.51 µs | +3.0% |
| ✅ | add/large_f32_threads=1-internal/4194304 |
892.80 µs | 908.44 µs | +1.8% |
| ✅ | matmul/large_generic_f16_threads=8/32x1024x1024 |
90.73 µs | 91.60 µs | +1.0% |
| ✅ | matmul/large_generic_bf16_threads=1/32x1024x1024 |
2.26 ms | 2.26 ms | +0.0% |
| ✅ | qwen3_sampling_processors/top_k_top_p_full_sort_baseline |
5.76 ms | 5.77 ms | +0.0% |
| ✅ | matmul/medium_generic_f32_threads=1/32x512x512 |
2.36 ms | 2.36 ms | -0.0% |
| ✅ | qwen3_sampling_processors/top_k_top_p_fast |
667.57 µs | 651.92 µs | -2.3% |
| ✅ | qwen3_sampling_processors/top_p_full_sort_after_top_k_baseline |
3.76 ms | 3.66 ms | -2.4% |
| ✅ | qwen3_sampling_processors/top_p_fast_after_top_k |
557.72 µs | 541.94 µs | -2.8% |
| ✅ | matmul/large_generic_f32_threads=1/32x1024x1024 |
9.18 ms | 8.85 ms | -3.7% |
| ✅ | gather/small_bf16_threads=1-internal/4096 |
616.2 ns | 582.7 ns | -5.4% |
| ✅ | add/medium_f32_threads=1-internal/262144 |
31.03 µs | 29.25 µs | -5.8% |
| ✅ | matmul/large_generic_f16_threads=1/32x1024x1024 |
81.13 µs | 76.05 µs | -6.3% |
| ✅ | sampling_latency/top_k_per_token |
56.41 µs | 52.26 µs | -7.3% |
| ✅ | add/large_f16_threads=1-internal/4194304 |
2.33 ms | 2.15 ms | -8.1% |
| ✅ | matmul/large_generic_bf16_threads=8/32x1024x1024 |
1.81 ms | 1.67 ms | -8.2% |
| ✅ | sampling_latency/min_p_per_token |
229.58 µs | 210.51 µs | -8.3% |
| ✅ | add/medium_bf16_threads=1-internal/262144 |
140.64 µs | 128.24 µs | -8.8% |
| ✅ | qwen3_sampling_processors/top_k_full_sort_baseline |
2.46 ms | 2.25 ms | -8.9% |
| ✅ | gather/medium_f16_threads=1-internal/32768 |
2.77 µs | 2.52 µs | -9.0% |
| ✅ | matmul/small_generic_bf16_threads=1/1x256x256 |
37.12 µs | 33.67 µs | -9.3% |
| ✅ | matmul/large_generic_f32_threads=8/32x1024x1024 |
4.04 ms | 3.66 ms | -9.5% |
| ✅ | kv_cache/alloc_dealloc_pages |
42.24 µs | 38.14 µs | -9.7% |
| ✅ | logit_processing/seven_processor_chain_per_step |
353.64 µs | 316.47 µs | -10.5% |
| ✅ | matmul/small_generic_f32_threads=1/1x256x256 |
49.94 µs | 44.49 µs | -10.9% |
| ✅ | add/small_f16_threads=1-internal/1024 |
547.4 ns | 477.4 ns | -12.8% |
| ✅ | matmul/medium_generic_bf16_threads=8/32x512x512 |
413.07 µs | 359.84 µs | -12.9% |
| ✅ | matmul/small_generic_bf16_threads=8/1x256x256 |
37.08 µs | 32.15 µs | -13.3% |
| 🟢 | matmul/medium_generic_f16_threads=1/32x512x512 |
32.99 µs | 27.92 µs | -15.4% |
| 🟢 | add/small_bf16_threads=1-internal/1024 |
564.2 ns | 477.2 ns | -15.4% |
| 🟢 | matmul/medium_generic_f32_threads=8/32x512x512 |
1.09 ms | 919.31 µs | -15.9% |
| 🟢 | matmul/medium_generic_bf16_threads=1/32x512x512 |
609.99 µs | 498.78 µs | -18.2% |
| 🟢 | gather/small_f16_threads=1-internal/4096 |
656.5 ns | 536.6 ns | -18.3% |
| 🟢 | matmul/medium_generic_f16_threads=8/32x512x512 |
34.32 µs | 27.88 µs | -18.8% |
| 🟢 | reduce_mean/large_f32_threads=1-internal/262144 |
1.26 ms | 1.02 ms | -19.6% |
| 🟢 | qwen3_sampling_processors/top_k_partial_selection |
179.21 µs | 141.54 µs | -21.0% |
| 🟢 | grammar_masking/llguidance_compute_mask/32 |
88.20 µs | 69.31 µs | -21.4% |
| 🟢 | gather/small_f32_threads=1-internal/4096 |
890.1 ns | 675.2 ns | -24.1% |
| 🟢 | reduce_mean/small_f32_threads=1-internal/4096 |
20.12 µs | 14.97 µs | -25.6% |
| 🟢 | reduce_mean/medium_f32_threads=1-internal/65536 |
331.44 µs | 244.87 µs | -26.1% |
| 🟢 | add/large_bf16_threads=1-internal/4194304 |
2.29 ms | 1.69 ms | -26.1% |
| 🟢 | add/small_f32_threads=1-internal/1024 |
286.5 ns | 210.7 ns | -26.4% |
| 🟢 | gather/large_f16_threads=1-internal/131072 |
22.66 µs | 16.26 µs | -28.3% |
| 🟢 | gather/large_bf16_threads=1-internal/131072 |
20.08 µs | 13.80 µs | -31.3% |
| 🟢 | block_quantized_matmul_cached_dense/mxfp4_cached_dense_repeated_call/1x1024x1024 |
74.96 µs | 51.40 µs | -31.4% |
| 🟢 | block_quantized_moe_cached_dense/mxfp4_uncached_expert_dequant_each_call/rows=1,H=256,I=256,E=4,top_k=1 |
604.44 µs | 372.47 µs | -38.4% |
| 🟢 | matmul/small_generic_f16_threads=1/1x256x256 |
56.73 µs | 32.25 µs | -43.2% |
| 🟢 | matmul/small_generic_f16_threads=8/1x256x256 |
61.69 µs | 34.55 µs | -44.0% |
| 🟢 | gather/large_f32_threads=1-internal/131072 |
51.02 µs | 27.71 µs | -45.7% |
| 🟢 | block_quantized_moe_cached_dense/mxfp4_cached_dense_expert_repeated_call/rows=1,H=256,I=256,E=4,top_k=1 |
131.43 µs | 68.03 µs | -48.2% |
| 🟢 | block_quantized_matmul_cached_dense/mxfp4_uncached_dequant_each_call/1x1024x1024 |
1.74 ms | 603.25 µs | -65.3% |
| 🟢 | matmul/small_generic_f32_threads=8/1x256x256 |
122.50 µs | 39.14 µs | -68.0% |
Visual flags:
Host info
CPU: Apple M1 (Virtual)
Cores: 3
OS: Darwin 25.5.0 arm64
Rust: rustc 1.97.1 (8bab26f4f 2026-07-14)
Load avg: { 3.25 3.49 5.94 }
What this cannot catch
- Regressions in code paths not covered by these benchmarks (e.g., end-to-end decode with a real model)
- Sub-threshold regressions that compound over multiple PRs
- Performance changes that only manifest under GPU execution
- Latency changes in the ORT integration path (these benchmarks exercise the native Rust kernels)
Codecov Report✅ All modified and coverable lines are covered by tests. Additional details and impacted files@@ Coverage Diff @@
## main #1350 +/- ##
==========================================
- Coverage 80.88% 80.50% -0.39%
==========================================
Files 364 362 -2
Lines 160729 157677 -3052
Branches 160729 157677 -3052
==========================================
- Hits 130005 126937 -3068
- Misses 26069 26098 +29
+ Partials 4655 4642 -13
Flags with carried forward coverage won't be shown. Click here to find out more. 🚀 New features to boost your workflow:
|
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
Follow-up to #1340 (split-KV FlashDecoding for
attention_row). That PR fixed keys-per-split atchunk=256, sonum_splits = cap/256. Re-profiling at deep context (~2600-tok prompt, the regime whereattention_rowreaches ~40% of decode) showed two things:cap≈4096the grid is(16 rows × 16 splits) = 256blocks, waves/SM 0.48, warps active 11%. The decode split kernel is memory-latency bound, so it wants ~a full occupancy wave of resident CTAs, not one block per SM — an ncu + E2E chunk sweep found the optimum atchunk=128(grid 512), notchunk=512(grid 132, which is ~5% slower).This PR sizes
num_splitsadaptively from the fixedcapand the fixed row count to target ~one occupancy wave, floored so no split is starved of work:For V2-Lite's 16 decode rows this lands
num_splits=32(grid 512) at deep context and scales down cleanly at shallow context. The pure sizing arithmetic is factored intoattention_split_geometrywith a unit test.Capture-safety (unchanged invariant)
num_splitsstill derives only fromcap+ row count, never live seqlen → eager and capture make identical launch decisions.chunkstays>= MIN_CHUNK (128), so any live context<= 128remains on the single-split bit-exact fast path (the golden 24-tok lock sits well under this).ONNX_GENAI_ATTN_SPLIT_CHUNKstill pins a fixed chunk (reproduces pre-adaptive behaviour, for A/B).Occupancy A/B — ncu
attention_split, V2-Lite deep ctx (~2600 tok)Adaptive ~doubles the grid, doubles occupancy, and lifts effective DRAM throughput — filling the previously idle SMs.
Numerics gates (H200, GPU4, pinned, single-threaded)
eager==capture+ golden prefix)standard_attention_capture_gpu/_fp16_gpu/_bf16_gpu--lib standard_attention)adaptive_split_geometry_*)standard_attention_gpu..._requires_homogeneous_floating_input_dtypesis pre-existing on clean origin/main, unrelatedPerf A/B (H200 GPU4, eager, medians-of-5)
Verdict: GO, default-ON (adaptive)
Adaptive beats the fixed-256 default at both depths (+4-5%), is neutral at short context, capture-safe, byte-identical short / f64-tol wide, no dense regression. It also documents that split-KV's true deep-context win is +70% vs monolithic, far larger than the +10% shallow number in #1340.
Do not self-merge — reporting to coordinator for admin-merge after review.
Co-authored-by: Copilot 223556219+Copilot@users.noreply.github.com