Skip to content

perf(cuda): adaptive split-KV sizing for attention_row decode (+4-5% over fixed-256; +70% deep-ctx) - #1350

Merged
justinchuby merged 1 commit into
mainfrom
squad/attention-splitkv-adaptive
Aug 18, 2026
Merged

justinchuby merged 1 commit into
mainfrom
squad/attention-splitkv-adaptive

Conversation

@justinchuby

Copy link
Copy Markdown
Owner

Summary

Follow-up to #1340 (split-KV FlashDecoding for attention_row). That PR fixed keys-per-split at chunk=256, so num_splits = cap/256. Re-profiling at deep context (~2600-tok prompt, the regime where attention_row reaches ~40% of decode) showed two things:

  1. The split-KV win is far bigger at depth than the shallow A/B in perf(cuda): split-KV FlashDecoding for attention_row decode (+10% V2-Lite wide-ctx) #1340 measured: deep-ctx E2E is 28.40 → 46.37 tok/s (+63%) with fixed-256, vs the +10% originally reported at ~500-ctx.
  2. Fixed-256 under-fills the machine at depth. At cap≈4096 the grid is (16 rows × 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 — an ncu + E2E chunk sweep found the optimum at chunk=128 (grid 512), not chunk=512 (grid 132, which is ~5% slower).

This PR sizes num_splits adaptively from the fixed cap and the fixed row count to target ~one occupancy wave, floored so no split is starved of work:

target_splits = ATTN_SPLIT_TARGET_BLOCKS.div_ceil(total_rows)   // 512 blocks ≈ one wave on 132 SMs
num_splits    = min(target_splits, cap.div_ceil(ATTN_SPLIT_MIN_CHUNK/*128*/), ATTN_SPLIT_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. The pure sizing arithmetic is factored into attention_split_geometry with a unit test.

Capture-safety (unchanged invariant)

num_splits still derives only from cap + row count, never live seqlen → eager and capture make identical launch decisions. chunk stays >= MIN_CHUNK (128), so any live context <= 128 remains on the single-split bit-exact fast path (the golden 24-tok lock sits well under this). ONNX_GENAI_ATTN_SPLIT_CHUNK still pins a fixed chunk (reproduces pre-adaptive behaviour, for A/B).

Occupancy A/B — ncu attention_split, V2-Lite deep ctx (~2600 tok)

grid waves/SM warps active DRAM throughput
monolithic (OFF) 16 0.02 9.8% 0.8%
fixed-256 (#1340) 256 0.48 11% 34%
adaptive (this PR) 512 0.97 24% 48%

Adaptive ~doubles the grid, doubles occupancy, and lifts effective DRAM throughput — filling the previously idle SMs.

Numerics gates (H200, GPU4, pinned, single-threaded)

Gate Result
V2-Lite golden 24-tok lock ✅ byte-identical (single-split path)
V2-Lite 340-tok long-ctx lock (eager==capture + golden prefix) ✅ pass (deep multi-split engaged, capture-safe)
standard_attention_capture_gpu / _fp16_gpu / _bf16_gpu ✅ pass
lib unit tests (--lib standard_attention) ✅ 15/15 (incl. new adaptive_split_geometry_*)
standard_attention_gpu ✅ 23/24 — the 1 fail ..._requires_homogeneous_floating_input_dtypes is pre-existing on clean origin/main, unrelated
Dense qwen2.5-0.5b token md5 ✅ byte-identical (split does not engage)

Perf A/B (H200 GPU4, eager, medians-of-5)

context OFF (monolithic) fixed-256 (#1340) adaptive (this PR) Δ vs fixed-256
deep (~2600-tok prompt) 28.40 46.37 48.32 +4.2%
shallow (~520-tok prompt) 51.69 58.08 60.82 +4.7%
short "Hello" (capture) 127.77 — 130.17 neutral (single-split)

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

…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
justinchuby merged commit 32ee82c into main Aug 18, 2026
6 checks passed
@justinchuby
justinchuby deleted the squad/attention-splitkv-adaptive branch August 18, 2026 23:59
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>
@github-actions

Copy link
Copy Markdown

⚠️ Benchmark Change Detected

Comparison of criterion micro-benchmarks: PR head vs merge-base, measured on the same runner in the same job (base first → PR second).

ℹ️ Absolute times are informational only — they vary with runner load. The % change column is the reliable signal because both sides ran under identical conditions.

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: ⚠️ ≥ 15% slower, 🔴 ≥ 30% slower — calibrated against measured runner noise (~27% worst-case on multi-threaded matmul)

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

codecov Bot commented Aug 19, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 80.50%. Comparing base (06c62e0) to head (108edd4).
⚠️ Report is 76 commits behind head on main.

Additional details and impacted files

Impacted file tree graph

@@            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     
Flag Coverage Δ
mlas ?
offline 80.50% <ø> (-0.30%) ⬇️

Flags with carried forward coverage won't be shown. Click here to find out more.
see 6 files with indirect coverage changes

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant