Repository navigation
feat(ep-cuda): fp16 TopK kernel (unblocks dense_fallback MoE routers on CUDA) - #612
Conversation
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 Report✅ All modified and coverable lines are covered by tests. Additional details and impacted files@@ 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
Flags with carried forward coverage won't be shown. Click here to find out more. 🚀 New features to boost your workflow:
|
|
VERDICT: APPROVE Independent review by Harry (reviewer, opus-4.8) — reviewed a detached worktree off 2 — fp16 write-back & compare (byte-exactness): CONFIRMED
4 — Non-final-axis oracle honesty: author is CORRECT (ruling against CPU)ONNX
3 — Tests non-vacuous / tight: CONFIRMED via mutation
5 — Scope + hygiene
CPU non-final-axis bug — follow-up?Yes, worth a separate, low-priority follow-up PR against Not merging (review-only). |
|
| 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:
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)
…+ 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>
Summary
The Qwen3.6-35B-A3B (and any
dense_fallbackMoE decoder) 256-expert / top-8router runs its gate
TopKin fp16. The CUDA EP used to reject fp16 TopK forall 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 remainingconformance-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 theexisting total-order
before()compare (equivalent tof32::total_cmp, matching theCPU 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
%37tie-heavy pattern, asserting GPU == CPUbyte-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 nowCLAIMS 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 input0elem_type=10 (Float16) — confirms the exact op gap.
require_dtype(.., Float32, "X")): 40×dtype Float16 unsupported; expected Float32→ whole-session CPU fallback.require_one_of(.., CUDA_FLOAT_DTYPES, "X")): fp16 TopK is claimed;the new claim-proof test exercises this through the public
supports_oppath.Validation
cargo test -p onnx-runtime-ep-cuda --features cuda --test indexing_gpu— 15/15pass 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 bytopk_non_final_axes_*). They disagree fora non-final axis with
inner>1andk>1. This is independent of dtype (affectsf32 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.