Skip to content

perf(cuda): let the split-K decode GEMV store bf16 directly - #1550

Merged
justinchuby merged 1 commit into
mainfrom
perf/gemv-bf16-direct-store-v2
Aug 20, 2026
Merged

justinchuby merged 1 commit into
mainfrom
perf/gemv-bf16-direct-store-v2

Conversation

@justinchuby

Copy link
Copy Markdown
Owner

Relands #1549 with the bug that forced its revert (7633aa2) fixed, and with
honest numbers.

On a bf16 model, MatMulNBitsKernel::run_bf16 runs the tuned fp16 path into
an fp16 staging buffer and then issues a cast_half launch per node to
narrow the result back to bf16. In a 64-token decode of Muse Glimmer 30B
that is 3024 cast_half launches. The block-32 split-K GEMV can just write
bf16 itself: its epilogue already rounds the fp32 accumulator to fp16 on a
single lane, so narrowing further is one instruction executed once per
output column.

matmul_nbits_store_narrowed takes a launch-uniform out_bf16 flag. The
bf16 arm rounds fp32 -> fp16 -> bf16 rather than fp32 -> bf16: the staged
route it replaces rounds to fp16 in the GEMV and then casts that fp16 to
bf16, and reproducing the double rounding is what makes this bit-identical
rather than merely close.

run_bf16 publishes the real bf16 destination in a thread-local for the
duration of the synchronous fp16 call it makes, and skips its cast if a GEMV
accepted it. Entries without the narrowing store never accept, and keep the
staged route.

What broke in #1549: for 1 < m <= decode_gemv_loop_max_m(), run_f16
launches one single-row GEMV per row, and row 0 starts exactly at the
staging pointer. Keyed on the pointer alone, row 0 claimed the whole bf16
output, wrote only its own n columns, and suppressed the cast -- leaving
rows 1.. stale. The offer now also carries the element count it covers and
is only accepted by a launch that writes all of it, so per-row launches
decline. bf16_direct_store_declines_the_small_batch_row_loop covers this
and fails if the guard is removed.

Measured on A100, Muse Glimmer 30B INT4, native backend, GPU 6:

  • nsys, 64-token decode: cast_half 3024 -> 2398 launches, 22.5 -> 21.3 ms
    of kernel time. The 626 removed casts match the accepted direct stores.
  • Wall clock, 128-token decode, levers alternated: 30.11 -> 30.52 tok/s
    (+1.4%), but the direct store led in only 2 of 3 reps and the spread
    within a single configuration is larger than the gap. Treat the speedup as
    unproven; what is established is that the work is strictly smaller.

Prefill is unaffected: the offer is published there too, but the M>1 path
has no narrowing store and never accepts.

ONNX_GENAI_GEMV_BF16_DIRECT_OUT=0 restores the staged route.

Verification, redone properly. The end-to-end check in #1549 ran
onnx-genai run --backend native without --device cuda, which silently
decodes on the CPU (zero NVRTC compilations), so it exercised none of this
and reported a vacuous match. Repeated with --device cuda, greedy, two
prompts, 60 and 80 new tokens: byte-identical both to the lever-off route
and to a build of the pre-change commit.

Tests: 439 passed / 0 failed (437 on main + two new). Both new tests assert
the direct-store counter, so a routing change that silently reverts to
staging fails instead of comparing the reference to itself.

Refs #1528

Co-authored-by: Copilot 223556219+Copilot@users.noreply.github.com
Copilot-Session: 0190e2eb-abe4-451f-b36d-44a035a99b7e

Relands #1549 with the bug that forced its revert (7633aa2) fixed, and with
honest numbers.

On a bf16 model, `MatMulNBitsKernel::run_bf16` runs the tuned fp16 path into
an fp16 staging buffer and then issues a `cast_half` launch per node to
narrow the result back to bf16. In a 64-token decode of Muse Glimmer 30B
that is 3024 `cast_half` launches. The block-32 split-K GEMV can just write
bf16 itself: its epilogue already rounds the fp32 accumulator to fp16 on a
single lane, so narrowing further is one instruction executed once per
output column.

`matmul_nbits_store_narrowed` takes a launch-uniform `out_bf16` flag. The
bf16 arm rounds fp32 -> fp16 -> bf16 rather than fp32 -> bf16: the staged
route it replaces rounds to fp16 in the GEMV and then casts that fp16 to
bf16, and reproducing the double rounding is what makes this bit-identical
rather than merely close.

`run_bf16` publishes the real bf16 destination in a thread-local for the
duration of the synchronous fp16 call it makes, and skips its cast if a GEMV
accepted it. Entries without the narrowing store never accept, and keep the
staged route.

What broke in #1549: for `1 < m <= decode_gemv_loop_max_m()`, `run_f16`
launches one single-row GEMV per row, and row 0 starts exactly at the
staging pointer. Keyed on the pointer alone, row 0 claimed the whole bf16
output, wrote only its own n columns, and suppressed the cast -- leaving
rows 1.. stale. The offer now also carries the element count it covers and
is only accepted by a launch that writes all of it, so per-row launches
decline. `bf16_direct_store_declines_the_small_batch_row_loop` covers this
and fails if the guard is removed.

Measured on A100, Muse Glimmer 30B INT4, native backend, GPU 6:

- nsys, 64-token decode: `cast_half` 3024 -> 2398 launches, 22.5 -> 21.3 ms
  of kernel time. The 626 removed casts match the accepted direct stores.
- Wall clock, 128-token decode, levers alternated: 30.11 -> 30.52 tok/s
  (+1.4%), but the direct store led in only 2 of 3 reps and the spread
  within a single configuration is larger than the gap. Treat the speedup as
  unproven; what is established is that the work is strictly smaller.

Prefill is unaffected: the offer is published there too, but the M>1 path
has no narrowing store and never accepts.

`ONNX_GENAI_GEMV_BF16_DIRECT_OUT=0` restores the staged route.

Verification, redone properly. The end-to-end check in #1549 ran
`onnx-genai run --backend native` without `--device cuda`, which silently
decodes on the CPU (zero NVRTC compilations), so it exercised none of this
and reported a vacuous match. Repeated with `--device cuda`, greedy, two
prompts, 60 and 80 new tokens: byte-identical both to the lever-off route
and to a build of the pre-change commit.

Tests: 439 passed / 0 failed (437 on main + two new). Both new tests assert
the direct-store counter, so a routing change that silently reverts to
staging fails instead of comparing the reference to itself.

Refs #1528

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
Copilot-Session: 0190e2eb-abe4-451f-b36d-44a035a99b7e
@justinchuby
justinchuby merged commit c03d843 into main Aug 20, 2026
3 checks passed
@justinchuby
justinchuby deleted the perf/gemv-bf16-direct-store-v2 branch August 20, 2026 06:39
justinchuby added a commit that referenced this pull request Aug 20, 2026
`cargo fmt --all --check` — one of this repo's two required checks — has
been failing on unmodified `main` since #1550 merged. A test's
`eprintln!` exceeds the line width:

```
Diff in crates/onnx-runtime-ep-cuda/src/kernels/matmul_nbits.rs:18919:
-            eprintln!("skipping MatMulNBits bf16 small-batch direct-store test: CUDA runtime unavailable");
+            eprintln!(
+                "skipping MatMulNBits bf16 small-batch direct-store test: CUDA runtime unavailable"
+            );
```

Reproduced on a clean checkout of `origin/main` with no local edits, so
it is not an artifact of any open branch.

This PR is **pure `cargo fmt --all` output** over that one file — 3
insertions, 1 deletion, no other change.

This is the third gate found red on `main` in as many days (see #1536,
#1546). The mechanism is the same each time: with only two required
checks and `strict_required_status_checks_policy=false`, a PR can merge
on a green run that predates the commit that breaks the gate.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
@codecov

codecov Bot commented Aug 20, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 82.64%. Comparing base (06b372c) to head (149c40f).
⚠️ Report is 41 commits behind head on main.

Additional details and impacted files

Impacted file tree graph

@@             Coverage Diff             @@
##             main    #1550       +/-   ##
===========================================
+ Coverage   80.36%   82.64%    +2.28%     
===========================================
  Files         378       12      -366     
  Lines      167774     5475   -162299     
  Branches   167774     5475   -162299     
===========================================
- Hits       134836     4525   -130311     
+ Misses      28088      757    -27331     
+ Partials     4850      193     -4657     
Flag Coverage Δ
cli-ort-linux 82.60% <ø> (?)
cli-ort-windows 82.10% <ø> (-0.10%) ⬇️
offline ?

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

🚀 New features to boost your workflow:
  • ❄️ Test Analytics: Detect flaky tests, report on failures, and find test suite problems.
  • 📦 JS Bundle Analysis: Save yourself from yourself by tracking and limiting bundle sizes in JS merges.

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