Repository navigation
perf(cuda): prefetch-pipeline the block-32 scales-fp16 split-K int4 GEMV (+4.4% q38 decode) - #1584
Conversation
The block-32 asymmetric (zero-point) scales-fp16 split-K int4 decode GEMV (`matmul_nbits_gemv_f16_scales_f16_zp_splitk`) fills idle SMs and adds per-column memory-level parallelism (K_SPLIT warps/column), but each lane still walks its weight-word chain with a single 32-bit global load in flight -> a load->accumulate->load dependency that stalls on the Long-Scoreboard weight-load latency the single-warp `_pipe` kernel already diagnosed (DRAM ~24% of an H200's 4.8 TB/s; the math is not the bottleneck). Add `_pf` prefetch-pipelined siblings of the split-K entries (symmetric + asymmetric) that keep PF weight-word loads in flight per lane via a register-resident shift register (manual rotation, no dynamic-indexed array -> no local spill), reusing the proven `_pipe` transform on the split-K loop. The lane->nibble mapping, `accumulate_int4x8_f16_zp` calls, per-split accumulation order and K_SPLIT shared-memory reduction are UNCHANGED, so the output is bit-identical to the plain split-K entry (verified by a new byte-identity test and by identical greedy token IDs on both models). Default-on; `ONNX_GENAI_ZP_SPLITK_PREFETCH=0` restores the plain entry for A/B. Capture-safe (registers + launch-time shared memory only). Generality: the swap keys off the already-shape-gated split-K entry, no hardcoded dims. A/B (H200 ordinal 2, idle, batch=1, 128-tok steady, 1 warmup, median-of-4 steady_median runs, qwen38-27b-int4-cuda asymmetric int4 block-32 bf16 io, features bench-native,cuda, ORT_ROOT=.ort-cuda-1.28/root, base 70db811): prefetch OFF 59.62 tok/s -> ON 62.51 tok/s = +4.85%. captures>0, fallbacks=0. mary control (qwen3.6-27b-int4-cuda symmetric f16) stays coherent + capture -safe with byte-identical greedy tokens ON vs OFF. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Codecov Report✅ All modified and coverable lines are covered by tests. Additional details and impacted files@@ Coverage Diff @@
## main #1584 +/- ##
==========================================
- Coverage 80.77% 80.70% -0.07%
==========================================
Files 381 379 -2
Lines 173645 170738 -2907
Branches 173645 170738 -2907
==========================================
- Hits 140262 137799 -2463
+ Misses 28492 28060 -432
+ Partials 4891 4879 -12
Flags with carried forward coverage won't be shown. Click here to find out more. 🚀 New features to boost your workflow:
|
…-gemv-m1 #1585 refactored the split-K template to <bool HasZp, int K_SPLIT> and added a K_SPLIT=8 deep-split entry (matmul_nbits_gemv_f16_scales_f16_zp_splitk8) that takes the GQA k/v projection when K_SPLIT=2 still leaves the grid under one wave. Reconcile the prefetch (_pf) transform with it: - Resolve the bf16_direct_capable conflict as a union (both my _pf entries and #1585's _splitk8 entry narrow through matmul_nbits_store_narrowed). - Parametrize the prefetch template on K_SPLIT too (matmul_nbits_gemv_f16_scales_f16_splitk_pf_tpl<bool HasZp, int K_SPLIT>) so the deep-split path can share the identical PF register-shift-register transform, and add a matmul_nbits_gemv_f16_scales_f16_zp_splitk8_pf entry (<true, 8>). The deep-split path is still single-load-in-flight per lane after the grid is deepened, so prefetch is an additional latency-hiding win there. - Extend the dispatch swap and bf16_direct_capable to route the _splitk8 entry to its _pf sibling under the same ONNX_GENAI_ZP_SPLITK_PREFETCH gate. Bit-identity preserved (K_SPLIT-generic transform, unchanged accumulation order): 43 matmul_nbits tests pass incl. the prefetch parity test and #1585's deepen test; q38 + mary greedy token IDs identical ON vs OFF. Re-measured A/B on the post-#1585 base (H200 ordinal 2, idle, batch=1, 128-tok steady, 1 warmup, median-of-4 steady_median runs, qwen38-27b-int4-cuda asymmetric int4 block-32 bf16, features bench-native,cuda, ORT_ROOT=.ort-cuda-1.28/root, base 0faa635): prefetch OFF 59.76 -> ON 62.49 tok/s = +4.57%. captures=4, fallbacks=0 both arms. mary control coherent + byte-identical ON vs OFF. Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Rebased onto #1585 (deep-split GQA k/v) + re-measuredorigin/main advanced to Reconciliation (the deep-split path also wants prefetch):
Re-measured A/B on the post-#1585 base (so #1585's deep-split is in both arms; this isolates the prefetch):
Env: batch=1, 128-tok steady, 1 warmup, median-of-4 Tests: all 43 |
main advanced past this branch's base (#1578, #1580, #1583..#1585, #1588, #1491, #1571, #1592). 23 files changed on main, 3 of them also touched here. Conflict resolution, and the audit behind it: kernels/matmul_nbits.rs -- taken from main wholesale. This branch's blob is byte-identical to main's blob at 3542ae7 (#1580): the only "change" on our side was restoring main's probe after an earlier merge dropped it (the Tier 2 audit). main has since carried that same file forward through #1584/#1585, so our content is a strict ancestor of main's and taking main loses nothing. Verified by blob hash, not by reading the diff. .github/workflows/ci.yml -- auto-merged. Result equals main plus the four `-p onnx-runtime-memory-api` package selections this stack adds. main's two new `cargo fetch --locked` steps are both present (counted). pipeline/decoder_component.rs -- auto-merged. Result equals main plus the two `ProcessMemoryManager` lines this stack adds. main's #1592 additions (supports_argmax, step_argmax, decode_argmax_with_step_inputs, captured_step_input_greedy_supported) all survive at main's reference counts. The other 20 files main changed are byte-identical to origin/main in the merge result, checked by blob hash for every one rather than by eye. Verified: cargo check --workspace --all-targets and cargo check -p onnx-genai-engine --features cuda,native-backend clean; cargo fmt --all --check clean; onnx-genai-engine native-backend suite 559 passed / 2 failed / 1 ignored, the two failures being the same macOS statvfs pair that fails on main and is fixed by #1586. Co-authored-by: Copilot App <223556219+Copilot@users.noreply.github.com> Copilot-Session: c80f8522-983c-47f7-8241-2155a823aabe
What
Adds prefetch-pipelined
_pfsiblings of the block-32 scales-fp16 split-K int4 decode GEMV (both the asymmetric zero-point entrymatmul_nbits_gemv_f16_scales_f16_zp_splitk— q38's dominant decode kernel — and the symmetric entry) and routes to them by default.ONNX_GENAI_ZP_SPLITK_PREFETCH=0restores the plain split-K entry for A/B.Why
The int4 M=1 GEMV (
MatMulNBits) is the confirmed #1 q38 decode cost (~25% of captured decode). The split-K kernel already fills idle SMs and adds per-column memory-level parallelism (K_SPLIT warps/column), but each lane still walks its weight-word chain with a single 32-bit global load in flight — aload -> accumulate -> loaddependency that stalls on the Long-Scoreboard weight-load latency the single-warp_pipekernel already diagnosed (ncu on the sibling path: ~75% of warp cycles stalled on the one in-flight load; DRAM ~24% of an H200's 4.8 TB/s; the dequant math is not the bottleneck).The
_pfvariant reuses the proven_pipetransform on the split-K loop: each lane keepsPF=3weight-word loads in flight via a register-resident shift register (manual rotation, no dynamic-indexed array -> no local spill), adding per-lane MLP on top of split-K's per-column MLP.Bit-identity
The lane->nibble mapping,
accumulate_int4x8_f16_zpcalls, per-split accumulation order and the K_SPLIT shared-memory reduction are unchanged — only the timing of the weight loads changes. Output is bit-identical to the plain split-K entry:scales_f16_zp_splitk_prefetch_is_bit_identical_to_splitk(asserts 0/N fp16 bit mismatches on grid-starved split-K shapes; also asserts the dims actually select split-K so the check isn't vacuous). All 41matmul_nbitstests pass.Capture-safety
Registers + launch-time shared memory only; no alloc/sync. Verified
captures>0, fallbacks=0on q38 and mary.Generality
The swap keys off the already-shape-gated split-K entry (
use_scales_f16_zp_splitk/use_f16_symmetric_splitk), which derive fromk/n/block_size/live SM count. No hardcoded head counts, dims, or block sizes.A/B (mandatory env block)
CUDA_VISIBLE_DEVICES)steady_medianruns (each itself median-of-3)--features bench-native,cuda,ORT_ROOT=.ort-cuda-1.28/root70db81192(merged up tob77e3d3fc); branch:squad/int4-gemv-m1(Pre-merge A/B on the same box measured +4.85%; both are well above the +2% gate.)
Profiling disclosure
ncuis blocked on this box (RmProfilingAdminOnly=1for non-admin), so attribution uses op-count + roofline + captured A/B timing rather than per-kernel counters. The bottleneck diagnosis (Long-Scoreboard weight-load latency) is quoted from the pre-existing_pipekernel's committed ncu analysis on the sibling single-warp path.Files
crates/onnx-runtime-ep-cuda/src/kernels/matmul_nbits.rs— new_pfdevice template + two extern entries,use_scales_f16_zp_splitk_pfgate, dispatch swap,bf16_direct_capableupdate,run_symmetric_block_rawprefetch toggle, new parity test.Co-authored-by: Copilot 223556219+Copilot@users.noreply.github.com