SIGN IN SIGN UP

[CUDA] Fix QMoE MXFP4 W4A16 on SM120/SM121 and speed up FP4 decode (#33067)

### Description

QMoE with MXFP4 weights and FP16/BF16 activations (`quant_type="fp4"`,
W4A16, e.g. GPT-OSS) fails at session creation on Blackwell consumer
GPUs (SM120/SM121), on every platform, and also on SM80-89 (and SM90 on
Windows):

```
TMA WS grouped MoE GEMM for SM%d is not compiled, and this QMoE configuration has no SM80 fallback90
```

#### Root causes

1. **Invalid SM120 -> SM90 remap.** `getTmaWarpSpecializedConfigs`
routed WFP4A16 on `sm >= 120` to the SM90 mixed-input TMA
warp-specialized kernel family. Those kernels use `sm_90a` WGMMA and are
compiled only for `90a-real`, so they can never run on SM100/SM120 (no
WGMMA, `sm_90a` SASS is not loadable). Windows/MSVC just failed earlier
because `COMPILE_HOPPER_TMA_GROUPED_GEMMS` is omitted there.
2. **Runner constructor throws before the SM80 decision is applied.**
`CutlassMoeFCRunner`'s constructor queries `getTactics(sm)` with
`use_sm80_fp4 = false` before `setUseSm80Fp4()` runs. For WFP4A16 on any
SM without a compiled SM90 TMA kernel (SM80-89 everywhere, SM90 on MSVC)
this threw, so the default `ORT_FP4_SM80_GEMM=1` regime could never be
constructed there.
3. **CUDA plugin EP read FLOAT8E8M0 tensors as INT2.**
`ep::adapter::CreateTensorFromApiValue` passed the C API
`ONNXTensorElementDataType` straight to
`DataTypeImpl::TensorTypeFromONNXEnum`, which expects a `TensorProto`
enum. The two enums diverge after `FLOAT4E2M1` (C API `FLOAT8E8M0 = 26`
vs proto `FLOAT8E8M0 = 24`, proto `26 = INT2`), so MXFP4 block scales
failed validation with `fc1_scales must be a float8e8m0 MXFP block-scale
tensor` in plugin builds.
4. The FP4 Python tests converted these failures into skips (any message
containing "FP4" — which includes the node name `QMoE_FP4` — or "SM"),
and the SM80-regime test class skipped SM120 entirely.

#### Fix

- MXFP4 W4A16 now uses the SM80 fused-dequant grouped GEMM (prefill) +
fused MXFP4 GEMV (decode) on **every SM >= 80**, including SM120/SM121.
This is the existing default regime on H200; the Ampere `mma.sync`
kernel runs natively on SM120 (the kernel's arch guard already maps
`__CUDA_ARCH__ >= 900` to the SM80 path), consumes FP4 weights directly
with a single pre-packed weight copy, and needs no TMA, so it also works
with MSVC. The SM90-only opt-in native path
(`ORT_ENABLE_FP4_CUTLASS_GEMM=1`) is unchanged.
- Removed the `sm >= 120 -> 90` remap. `getTmaWarpSpecializedConfigs`
returns no WFP4A16 TMA configs unless `sm == 90` and the kernel is
compiled (no throw, since the constructor queries before
`setUseSm80Fp4`). The QMoE constructor now enforces that the selected
runner has at least one tactic, with a clear error.
- Plugin EP adapter converts the C API element type to the `TensorProto`
type (`OnnxTensorProtoTypeFromApiElementType`). The QMoE block-scale
type error now reports the received element type.
- Tests: `TestQMoEFP4Sm80SingleWeightCopy` runs on SM120; the FP4 tests
only skip on genuine "not built" messages.

SM120 block-scaled FP4 tensor cores require both operands to be
block-scaled (FP4xFP4 / FP8xFP4), so a W4A16 kernel on them would need
activation quantization and different numerics; that is out of scope
here.

#### Kernel selection for SM < 100

| SM / platform | Before | After |
|---|---|---|
| SM < 80 | dequant fallback | unchanged |
| SM80-89 (all platforms) | session creation throws | SM80 grouped GEMM
+ fused GEMV (intended `ORT_FP4_SM80_GEMM=1` default) |
| SM90 Linux | SM80 grouped GEMM + GEMV (default) / SM90 TMA (opt-in) |
unchanged |
| SM90 Windows | session creation throws | SM80 grouped GEMM + GEMV;
opt-in native fails with a clear error |
| any SM with `ORT_FP4_SM80_GEMM=0` | dequant fallback + GEMV |
unchanged |

#### Performance: faster FP4 (e2m1) decode

Profiling the SM80 grouped GEMM and the pair-interleaved decode GEMV on
RTX 5060 Ti showed the FP4 -> FP16/BF16 conversion dominating: the
grouped GEMM converter decoded each nibble with scalar bit-field logic
(~90 instructions per 8 values), and the GEMV's pair-interleaved path
(the default decode layout in the SM80 regime) fell back to per-element
decode. Both now restore linear nibble order per 32-bit word with one
`prmt` + a nibble swap (`cutlass::detail::fp4_e2m1x8_uninterleave`) and
decode with packed `prmt` table lookups (~18 instructions per 8 values).
The output is bit-identical (exhaustively verified over all 2^32 packed
words for half and bf16 on SM86 and SM120).

GPT-OSS-20B-sized layer (32 experts, hidden = inter = 2880, top-4,
FP16), RTX 5060 Ti, CUDA 13.0, Nsight Systems GPU kernel time per call:

| tokens | kernel | before | after | change |
|---:|---|---:|---:|---:|
| 1 | GEMV fc1 / fc2 | 93.9 / 48.6 us | 92.1 / 47.4 us | -2%
(DRAM-bound) |
| 4 | GEMV fc1 / fc2 | 301.7 / 151.1 us | 295.2 / 147.8 us | -2% |
| 16 | GEMV fc1 / fc2 | 949.3 / 460.1 us | 642.3 / 318.9 us | -32% /
-31% |
| 128 | grouped GEMM | 951.3 us | 667.8 us | -30% |
| 512 | grouped GEMM | 2227.7 us | 1611.5 us | -28% |
| 2048 | grouped GEMM | 6159.2 us | 5273.4 us | -14% |

### Testing

- Windows, CUDA 13.0, CUDA plugin EP build
(`CMAKE_CUDA_ARCHITECTURES=120-real`), RTX 5060 Ti (SM120):
`test_qmoe_fp4_cuda.py` 31 passed, 1 skipped (SM90-only native test).
Before the fix every inference test failed at session creation.
- Linux, CUDA 13.0, in-process CUDA EP build
(`CMAKE_CUDA_ARCHITECTURES=80`), A100 (SM80): `test_qmoe_fp4_cuda.py` 31
passed, 1 skipped (SM90-only native test); the FP4 test skip guards now
cover SM80+.
- Reproduced the pre-fix failure with an in-process CUDA EP build on RTX
3060 (SM86) and RTX 5060 Ti (SM120).
- Exhaustive bit-exactness check of the new FP4 converters (all 2^32
inputs, half and bf16) on SM86 and SM120.

### Motivation and Context

GPT-OSS MXFP4 models could not run on RTX 50-series / DGX Spark (SM121)
GPUs.
T
Tianlei Wu committed
3886abd4824fb04fdf9219cfb4bf2b861d833288
Parent: c971ca2
Committed by GitHub <noreply@github.com> on 10/2/2026, 6:20:29 AM