Skip to content

feat(ep-cuda): fp16 TopK kernel (unblocks dense_fallback MoE routers on CUDA) - #612

Merged
justinchuby merged 1 commit into
mainfrom
squad/cuda-fp16-topk
Aug 3, 2026
Merged

justinchuby merged 1 commit into
mainfrom
squad/cuda-fp16-topk

Conversation

@justinchuby

Copy link
Copy Markdown
Owner

Summary

The Qwen3.6-35B-A3B (and any dense_fallback MoE decoder) 256-expert / top-8
router runs its gate TopK in fp16. The CUDA EP used to reject fp16 TopK for
all 40 router nodes with input 0 ('X') dtype Float16 unsupported; expected Float32, forcing a whole-session CPU fallback.

The fp16/bf16 TopK kernel and claim gate already landed in #445
(feat(cuda): support fp16 router TopK). This PR closes the remaining
conformance-coverage gap: #445 only tested a small final-axis ties case, but the
router unblock needs parity proven at the real router shape and axis, plus an
explicit claim-regression guard.

Correctness argument (upcast-for-compare is exact)

The kernel widens each compared value to f32 (static_cast<float>) and reuses the
existing total-order before() compare (equivalent to f32::total_cmp, matching the
CPU EP). fp16→f32 widening is lossless, so the ordering is identical to comparing
the fp16 values directly; ties break on ascending index, matching ONNX/ORT/CPU. The
kernel writes back the original raw fp16/bf16 element, so values are byte-identical
to the CPU oracle (widen→narrow round-trips exactly, preserving sign-of-zero).
Indices are Int64; values output dtype matches the input dtype, per the ONNX TopK spec.

What this PR adds (crates/onnx-runtime-ep-cuda/tests/indexing_gpu.rs)

  • topk_fp16_router_scale_and_non_final_axis_match_cpu — fp16 [2,256] top-8
    (the exact 35B router shape) with a %37 tie-heavy pattern, asserting GPU == CPU
    byte-for-byte and that each token's 8 selected experts are distinct/in-range.
    Also covers a non-final axis (axis 0), oracled against the proven f32 GPU kernel.
  • cuda_ep_claims_fp16_topk_router_nodes — regression guard asserting the EP now
    CLAIMS fp16 and bf16 TopK via supports_op, and still claims f32.

35B router unblock evidence

  • qwen36-35b-a3b-artifacts/decoder/model.onnx: 40 TopK nodes, all input0
    elem_type=10 (Float16)
    — confirms the exact op gap.
  • Before (pre-feat(cuda): support fp16 router TopK #445 claim gate require_dtype(.., Float32, "X")): 40× dtype Float16 unsupported; expected Float32 → whole-session CPU fallback.
  • After (require_one_of(.., CUDA_FLOAT_DTYPES, "X")): fp16 TopK is claimed;
    the new claim-proof test exercises this through the public supports_op path.

Validation

  • cargo test -p onnx-runtime-ep-cuda --features cuda --test indexing_gpu — 15/15
    pass
    on H200 (incl. all 3 fp16 TopK cases).
  • every_covered_op_has_a_conformance_entry (coverage-of-coverage) — pass.
  • cargo clippy -p onnx-runtime-ep-cuda --features cuda -- -D warnings — clean.
  • cargo fmt --all --check — clean.

Note: pre-existing CPU/GPU non-final-axis layout divergence (out of scope)

While adding axis coverage I found the CPU EP writes non-final-axis TopK in
push order (outer, inner, k) while the CUDA kernel and ONNX use k-major
(outer, k, inner, the layout asserted by topk_non_final_axes_*). They disagree for
a non-final axis with inner>1 and k>1. This is independent of dtype (affects
f32 identically) and predates this work — flagged for a separate CPU-EP fix; this PR
sidesteps it by oracling the fp16 non-final-axis case against the f32 GPU path.

Scope

Native full 35B decode is still blocked by native pipeline decode (GAP 3) and rank-3
mRoPE positions — out of scope. This PR only proves the TopK op gap is closed.

The fp16/bf16 TopK kernel and claim gate landed in #445. Close the
remaining conformance-coverage gap for the Qwen3.6-35B-A3B MoE router
unblock: add byte-exact fp16 GPU/CPU parity at the real 256-expert /
top-8 router shape (tie-heavy), a non-final-axis fp16 case oracled
against the proven f32 GPU kernel, and a supports_op regression guard
asserting the EP now claims fp16/bf16 TopK (and still claims f32).

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

codecov Bot commented Aug 3, 2026 •

Copy link
Copy Markdown

Codecov Report

✅ All modified and coverable lines are covered by tests.
✅ Project coverage is 81.29%. Comparing base (a447bbc) to head (450be58).
⚠️ Report is 1 commits behind head on main.

Additional details and impacted files

Impacted file tree graph

@@            Coverage Diff             @@
##             main     #612      +/-   ##
==========================================
+ Coverage   80.71%   81.29%   +0.57%     
==========================================
  Files         316      318       +2     
  Lines      125320   127564    +2244     
  Branches   125320   127564    +2244     
==========================================
+ Hits       101158   103706    +2548     
+ Misses      20019    19696     -323     
- Partials     4143     4162      +19     
Flag Coverage Δ
cli-ort-linux 86.69% <ø> (ø)
cli-ort-windows 83.84% <ø> (+0.10%) ⬆️
mlas 81.29% <ø> (?)
offline 81.08% <ø> (+0.60%) ⬆️

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

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

@justinchuby

Copy link
Copy Markdown
Owner Author

VERDICT: APPROVE

Independent review by Harry (reviewer, opus-4.8) — reviewed a detached worktree off origin/main + FETCH_HEAD of squad/cuda-fp16-topk. I did not author this. Diff is +137/−1 in crates/onnx-runtime-ep-cuda/tests/indexing_gpu.rs plus one decisions-inbox note — test-only, no kernel/product change on this branch (git diff --name-only origin/main...HEAD = the test file + .squad/decisions/inbox/...).

2 — fp16 write-back & compare (byte-exactness): CONFIRMED

src/kernels/topk.rs topk_float<T>:

  • Reads the ORIGINAL element T raw_value = input[...] (L44), keeps best_value = raw_value (L48), and writes values[offset] = best_value (L53). It writes back the original raw fp16/bf16 bits — never static_cast<T>(static_cast<float>(x)). So sign-of-zero and subnormals cannot be perturbed by an upcast/downcast round-trip; the selected value is bit-identical to the input element (and hence to the CPU oracle).
  • Compare only upcasts for ordering: float value = static_cast<float>(raw_value) feeding before(), which builds the standard total-order key (__float_as_int + sign-flip) ≡ f32::total_cmp, tie-broken by ascending index (ia < ib, L23). fp16→f32 widening is exact, injective and order-preserving for all finite values and both zeros (−0→−0, +0→+0), so the upcast total order is identical to a native fp16 total_cmp; equal fp16 bits ⇒ equal keys ⇒ ascending-index tie-break, matching CPU. No write-back path can differ from CPU for finite/zero fp16 inputs. No correctness hole.

4 — Non-final-axis oracle honesty: author is CORRECT (ruling against CPU)

ONNX TopK spec: outputs have shape [a_0,…,a_{axis-1}, k, a_{axis+1},…,a_n] — the K dimension sits at the axis position, so output[…, j, …] along axis is the j-th ranked element (k-major).

  • CUDA writes offset = (outer * k + out) * inner + i → rank at the axis position, inner columns contiguous = k-major = spec-correct. Verified against the existing topk_non_final_axes_* expectations (axis-0 [3,4] k=2 → values [5,4,8,9,5,4,7,9], indices [0,1,2,1,1,2,0,2]), which I recomputed by hand — matches.
  • CPU EP (onnx-runtime-ep-cpu/src/kernels/selection.rs L589-607) pushes in loop order for outer { for i(inner) { for d in top-k } } = (outer, inner, k) = push order, K innermost — wrong per ONNX whenever inner>1 && k>1. For a final axis inner==1, both layouts coincide, so CPU is correct there.
  • The author therefore has it the right way round: the new non-final-axis case oracles fp16 against the proven-correct f32 GPU result (both k-major), NOT against the buggy CPU path — it is not masking a CUDA bug. The final-axis router case does assert GPU==CPU, which is legitimate because CPU is correct on the final axis.
  • The flagged CPU bug is real and genuinely pre-existing: selection.rs is untouched on this branch (only the test file + inbox note changed), so it exists identically on main.

3 — Tests non-vacuous / tight: CONFIRMED via mutation

run()/run_cpu() return raw output bytes per tensor, so assert_eq!(router, run_cpu(...)) is a byte-exact GPU==CPU comparison of BOTH fp16 values and Int64 indices at the real router shape [2,256], k=8, with a %37 tie-heavy pattern (exercises the ascending-index tie-break). Mutation test: I flipped the kernel tie-break ia < ib → ia > ib and reran — both topk_fp16_router_scale_and_non_final_axis_match_cpu and topk_fp16_and_bf16_router_values_match_cpu_order FAIL on index bytes (values stay equal since ties are equal-valued, indices diverge), proving the assertions are tight and would catch a selection/order regression. Reverted the mutation after.

5 — Scope + hygiene

  • cargo fmt -p onnx-runtime-ep-cuda -- --check: clean.
  • Prescribed CUDA gate cargo clippy -p onnx-runtime-ep-cuda --features cuda -- -D warnings: clean. The changed test target also clippy-clean under -D warnings (--test indexing_gpu). (A --tests-wide clippy run surfaces lints in other files — group_query_attention_gpu.rs type_complexity, and #[cfg(test)] blocks in normalization.rs/standard_attention.rs for is_multiple_of/repeat_n — all pre-existing on main and pure toolchain drift from local clippy 1.97; none in this PR's file.)
  • cargo test -p onnx-runtime-ep-cuda --features cuda --test indexing_gpu on H200: 15/15 pass, including both new tests and the coverage-of-coverage guard.

CPU non-final-axis bug — follow-up?

Yes, worth a separate, low-priority follow-up PR against onnx-runtime-ep-cpu/src/kernels/selection.rs to write TopK outputs by tensor strides (k-major) instead of push order. It is latent: real models put TopK on the final axis (MoE routers, sampling/top-k), so the non-final-axis + inner>1 + k>1 path is rarely if ever hit in production — hence not a blocker for this PR, but it is a genuine ONNX-conformance defect that should be tracked so it doesn't bite an exotic graph later.

Not merging (review-only).

@justinchuby
justinchuby merged commit be6d4e3 into main Aug 3, 2026
14 checks passed
@justinchuby
justinchuby deleted the squad/cuda-fp16-topk branch August 3, 2026 07:43
@github-actions

github-actions Bot commented Aug 3, 2026

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/large_bf16_threads=1-internal/131072 14.37 µs 17.89 µs +24.5%
⚠️ gather/large_f16_threads=1-internal/131072 14.14 µs 16.87 µs +19.3%
✅ sampling_latency/top_k_per_token 60.75 µs 64.52 µs +6.2%
✅ matmul/large_generic_f32_threads=8/32x1024x1024 6.46 ms 6.83 ms +5.6%
✅ sampling_latency/top_p_per_token 434.07 µs 455.81 µs +5.0%
✅ matmul/medium_generic_bf16_threads=8/32x512x512 657.22 µs 689.00 µs +4.8%
✅ sampling_latency/greedy_per_token 3.54 µs 3.70 µs +4.4%
✅ matmul/large_generic_f16_threads=1/32x1024x1024 92.96 µs 94.88 µs +2.1%
✅ qwen3_sampling_processors/top_k_full_sort_baseline 2.45 ms 2.48 ms +1.5%
✅ matmul/small_generic_bf16_threads=8/1x256x256 44.18 µs 44.77 µs +1.3%
✅ sampling_latency/min_p_per_token 236.13 µs 239.17 µs +1.3%
✅ qwen3_sampling_processors/top_k_top_p_fast 730.36 µs 739.74 µs +1.3%
✅ matmul/medium_generic_f16_threads=1/32x512x512 40.56 µs 40.87 µs +0.7%
✅ matmul/small_generic_bf16_threads=1/1x256x256 40.80 µs 41.10 µs +0.7%
✅ qwen3_sampling_processors/top_k_partial_selection 159.88 µs 160.66 µs +0.5%
✅ matmul/medium_generic_f16_threads=8/32x512x512 49.37 µs 49.31 µs -0.1%
✅ matmul/large_generic_bf16_threads=8/32x1024x1024 2.37 ms 2.37 ms -0.2%
✅ grammar_masking/llguidance_compute_mask/32 90.04 µs 89.54 µs -0.6%
✅ matmul/medium_generic_bf16_threads=1/32x512x512 627.21 µs 622.55 µs -0.7%
✅ logit_processing/seven_processor_chain_per_step 365.29 µs 361.77 µs -1.0%
✅ matmul/large_generic_f16_threads=8/32x1024x1024 115.66 µs 114.49 µs -1.0%
✅ matmul/large_generic_f32_threads=1/32x1024x1024 11.69 ms 11.54 ms -1.2%
✅ matmul/small_generic_f32_threads=1/1x256x256 48.51 µs 47.83 µs -1.4%
✅ qwen3_sampling_processors/top_p_fast_after_top_k 583.45 µs 574.98 µs -1.5%
✅ matmul/medium_generic_f32_threads=1/32x512x512 2.84 ms 2.76 ms -2.7%
✅ qwen3_sampling_processors/top_k_top_p_full_sort_baseline 6.67 ms 6.49 ms -2.8%
✅ matmul/small_generic_f16_threads=8/1x256x256 43.08 µs 41.07 µs -4.7%
✅ qwen3_sampling_processors/top_p_full_sort_after_top_k_baseline 4.30 ms 4.08 ms -5.2%
✅ matmul/small_generic_f16_threads=1/1x256x256 42.00 µs 39.62 µs -5.7%
✅ matmul/medium_generic_f32_threads=8/32x512x512 1.92 ms 1.80 ms -5.9%
✅ tokenization/encode_tokens_per_second 471.58 µs 443.40 µs -6.0%
✅ matmul/large_generic_bf16_threads=1/32x1024x1024 2.57 ms 2.39 ms -6.9%
✅ tokenization/decode_tokens_per_second 7.84 ms 7.30 ms -6.9%
✅ add/medium_f32_threads=1-internal/262144 3.10 ms 2.86 ms -7.7%
✅ kv_cache/alloc_dealloc_pages 45.01 µs 41.44 µs -7.9%
✅ add/medium_f16_threads=1-internal/262144 3.25 ms 2.87 ms -11.7%
✅ add/medium_bf16_threads=1-internal/262144 3.44 ms 3.01 ms -12.5%
✅ add/small_bf16_threads=1-internal/1024 16.35 µs 14.21 µs -13.1%
✅ add/large_f32_threads=1-internal/4194304 52.25 ms 45.28 ms -13.3%
✅ reduce_mean/small_f32_threads=1-internal/4096 19.43 µs 16.74 µs -13.9%
✅ add/small_f16_threads=1-internal/1024 16.92 µs 14.51 µs -14.2%
✅ add/large_f16_threads=1-internal/4194304 54.37 ms 46.55 ms -14.4%
✅ add/large_bf16_threads=1-internal/4194304 55.71 ms 47.50 ms -14.7%
🟢 gather/small_f32_threads=1-internal/4096 912.2 ns 765.8 ns -16.0%
🟢 reduce_mean/large_f32_threads=1-internal/262144 1.37 ms 1.10 ms -19.4%
🟢 gather/small_bf16_threads=1-internal/4096 707.9 ns 566.9 ns -19.9%
🟢 gather/large_f32_threads=1-internal/131072 59.83 µs 46.92 µs -21.6%
🟢 gather/medium_bf16_threads=1-internal/32768 3.78 µs 2.91 µs -23.1%
🟢 gather/medium_f16_threads=1-internal/32768 3.95 µs 2.92 µs -26.0%
🟢 gather/medium_f32_threads=1-internal/32768 6.49 µs 4.73 µs -27.1%
🟢 gather/small_f16_threads=1-internal/4096 735.5 ns 533.3 ns -27.5%
🟢 reduce_mean/medium_f32_threads=1-internal/65536 393.01 µs 272.63 µs -30.6%
🟢 add/small_f32_threads=1-internal/1024 391.6 ns 223.2 ns -43.0%
🟢 matmul/small_generic_f32_threads=8/1x256x256 103.44 µs 54.60 µs -47.2%

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: { 4.32 3.90 7.24 }
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)

justinchuby added a commit that referenced this pull request Aug 3, 2026
…+ GAP-3 decomposition (increment 1/N) (#613)

## What & why

Scope pass on **GAP-3 (native pipeline decode)**. The finding is that
GAP-3's core is
**already implemented and conformance-locked on `origin/main`** — the
task was authored
against a local checkout (`1ba215ee`) that is ~35+ merged PRs behind
`origin/main`
(`be6d4e34`, #612). So this PR does **not** add new decode
functionality. It is a
bounded, **zero-behavior-change** truth-up + a design/decomposition
drop.

This is **GAP-3 housekeeping increment 1 of N** (see the design note in
this diff:

`.squad/decisions/inbox/cohaagen-gap3-native-pipeline-decode-design.md`).
The
substantive increments already landed:

- Backend-neutral component ownership seam — #546
- `NativePipelineDecoder` via `PipelineDecoderComponent` (Inc2b) — #479
- Pure-native multi-component decode wiring (Inc-A) — #565
- Native present-KV mirroring, paged (Inc-C) — #566
- Device-resident present-KV read-out (Inc-D / D.1) — #567/#568
- rank-3 mrope native positions — #543 · text-only decode pipeline —
#535 · fp16 TopK MoE router — #612

The pipeline decode loop (`PipelineDecodeLoopBackend`) already owns
`Box<dyn PipelineDecoderComponent>` + `Box<dyn ComponentSession>`, not
an ORT `Session`;
`DecodeState`/ORT `Value` are confined to `OrtPipelineDecoder`. The
Qwen3.5-0.8B hybrid
(same class as Qwen3.6-35B-A3B) decodes natively with **token-for-token
parity vs ORT**
under `tests/qwen35_0_8b_hybrid_native_cuda_e2e.rs`.

## The change

`native_component.rs`'s module doc still claimed wiring native sessions
into *"the
ORT-owned pipeline decode loop is the remaining GAP 3 work"* — false
since #546/#565.
Corrected to describe the now-backend-neutral loop and the merged
Inc-A/C/D, and to name
the genuinely-remaining feature-sized gaps (non-flat plans, native
cross-attn/vision KV).

**Behavior-preserving:** comment-only in `native_component.rs` + a
tracked decision drop.
No code path, signature, or data change.

## Remaining decomposition (in the design note)

Each is feature-sized (needs op/attention support and/or fixtures —
**not** zero-behavior),
none blocks the text-only 35B-A3B native number: R2 native
sliding-window paged mirror ·
R3 Inc-D.2 discontinuous prefix reuse · R4 native cross-attn/vision KV
(Inc3) · R5 non-flat
plans native · R6 (optional) neutral host tensor in the shared pool.

## Verification for the first native pipeline model

Byte/token-exact differential vs an ORT-backend decode of the **same
artifact** (ORT
front-end for both arms, decoder EP isolated), greedy, token-for-token —
already in place
for the 0.8B hybrid; same harness pattern applied to the real 35B-A3B is
the next step.

## Checks

- `cargo fmt --all` clean
- `cargo clippy -p onnx-genai-engine --features "native-backend cuda" --
-D warnings` clean

Left **open for Harry review**; do not merge without it.

Co-authored-by: Copilot <223556219+Copilot@users.noreply.github.com>
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