[CUDA] QMoE: fused int2 mixed-input GEMM + GeGLU activation + fractional zero-point - #32198
Thiago Pereira Rocha (thpereir) wants to merge 4 commits into
Conversation
|
Azure Pipelines: There may be pipelines that require an authorized user to comment /azp run to run. |
There was a problem hiding this comment.
Pull request overview
Adds fused int2 QMoE CUDA execution, GeGLU activation, and fractional zero-point support.
Changes:
- Adds int2 preprocessing, GEMM/GEMV dispatch, conversion, and pipeline handling.
- Adds tanh-approximated GeGLU activation.
- Adds fractional zero-point handling, quantization support, and CUDA tests.
Reviewed changes
Copilot reviewed 32 out of 32 changed files in this pull request and generated 6 comments.
Show a summary per file
| File | Description |
|---|---|
onnxruntime/test/python/transformers/test_qmoe_cuda.py |
Adds int2, GeGLU, and fractional-ZP tests. |
onnxruntime/python/tools/quantization/cuda_quantizer.py |
Adds 2-bit packing and quantization. |
onnxruntime/core/graph/contrib_ops/contrib_defs.cc |
Extends QMoE/MoE schemas. |
onnxruntime/contrib_ops/cuda/moe/qmoe_kernels.h |
Declares constant-bias launchers. |
onnxruntime/contrib_ops/cuda/moe/qmoe_kernels.cu |
Implements fractional-ZP bias kernels. |
onnxruntime/contrib_ops/cuda/moe/moe_quantization.h |
Adds fractional-ZP state and helpers. |
onnxruntime/contrib_ops/cuda/moe/moe_quantization.cc |
Wires int2 runners, preprocessing, and bias handling. |
onnxruntime/contrib_ops/cuda/moe/moe_base.h |
Parses GeGLU activation. |
onnxruntime/contrib_ops/cuda/llm/moe_gemm/moe_kernels.cu |
Extends MoE dispatch to int2. |
onnxruntime/contrib_ops/cuda/llm/moe_gemm/moe_gemv.h |
Documents int2 GEMV support. |
onnxruntime/contrib_ops/cuda/llm/moe_gemm/moe_gemv.cu |
Implements int2 GEMV dispatch. |
onnxruntime/contrib_ops/cuda/llm/moe_gemm/moe_gemm_template_dispatch.h |
Enables int2 GEMM templates. |
onnxruntime/contrib_ops/cuda/llm/moe_gemm/moe_gemm_kernels_fp16_uint2.cu |
Instantiates FP16/int2 GEMM. |
onnxruntime/contrib_ops/cuda/llm/moe_gemm/moe_gemm_kernels_bf16_uint2.cu |
Instantiates BF16/int2 GEMM. |
onnxruntime/contrib_ops/cuda/llm/moe_gemm/moe_gemm_activation_kernels.cuh |
Uses tanh GELU for GeGLU. |
onnxruntime/contrib_ops/cuda/llm/fpA_intB_gemv/details.h |
Adds int2 GEMV conversion traits. |
onnxruntime/contrib_ops/cuda/llm/fpA_intB_gemm_preprocessors.h |
Defines W2 quantization type. |
onnxruntime/contrib_ops/cuda/llm/fpA_intB_gemm_preprocessors_impl.h |
Adds W2 layout and permutation metadata. |
onnxruntime/contrib_ops/cuda/llm/fpA_intB_gemm_preprocessors_impl.cu |
Implements int2 transpose/interleave preprocessing. |
onnxruntime/contrib_ops/cuda/llm/fpA_intB_gemm_adaptor.h |
Declares int2 transpose adaptor. |
onnxruntime/contrib_ops/cuda/llm/fpA_intB_gemm_adaptor.cu |
Implements packed int2 transposition. |
onnxruntime/contrib_ops/cuda/llm/cutlass_extensions/interleaved_numeric_conversion.h |
Adds int2-to-FP16/BF16 converters. |
onnxruntime/contrib_ops/cuda/llm/cutlass_extensions/gemm/threadblock/dq_mma_pipelined_percol.h |
Fixes B-fragment parity. |
onnxruntime/contrib_ops/cuda/llm/cutlass_extensions/gemm/threadblock/dq_mma_pipelined_finegrained.h |
Fixes fine-grained pipeline parity. |
onnxruntime/contrib_ops/cuda/llm/cutlass_extensions/gemm/threadblock/dq_mma_multistage_percol.h |
Fixes multistage B-fragment parity. |
onnxruntime/contrib_ops/cuda/llm/cutlass_extensions/gemm/threadblock/dq_mma_multistage_finegrained.h |
Fixes fine-grained multistage parity. |
onnxruntime/contrib_ops/cuda/llm/cutlass_extensions/gemm/threadblock/default_mma.h |
Adds FP16/int2 MMA specializations. |
onnxruntime/contrib_ops/cuda/llm/cutlass_extensions/gemm/threadblock/default_mma_bf16.h |
Adds BF16/int2 MMA specializations. |
onnxruntime/contrib_ops/cuda/llm/cutlass_extensions/gemm/threadblock/default_dq_mma_pipelined.h |
Allows int2 pipelined dequant MMA. |
onnxruntime/contrib_ops/cuda/llm/cutlass_extensions/gemm/threadblock/default_dq_mma_multistage.h |
Allows int2 multistage dequant MMA. |
onnxruntime/contrib_ops/cuda/llm/cutlass_extensions/gemm/kernel/moe_cutlass_kernel.h |
Requires scales for int2 kernels. |
onnxruntime/contrib_ops/cuda/llm/cutlass_extensions/gemm/kernel/mixed_gemm_B_layout.h |
Defines the int2 mixed-GEMM layout. |
Suppressed comments (1)
onnxruntime/python/tools/quantization/cuda_quantizer.py:573
- This comment incorrectly says int2 has no CUTLASS mixed-input GEMM and uses a dequant fallback. The raw tensor is needed because W2 is transformed by QMoE's internal PrePack rather than this offline helper; update the explanation and error text accordingly.
💡 Configure MCP servers for context-aware, tailored reviews. Learn more in the docs.
|
Azure Pipelines: There may be pipelines that require an authorized user to comment /azp run to run. |
Add 2-bit (uint2b_t) weight support to the QMoE CUDA EP as a fused mixed-input GEMM: weights stay packed in HBM and are dequantized in-register inside the SM80 DqMma pipeline, mirroring the int4 path (no full-precision weight materialization). Includes the in-register interleaved converter, W2_A16 weight preprocessor (bias+interleave, LDSM permutation), LayoutDetailsB<uint2b_t>, DqMma/DefaultMma wiring, runner instantiations, decode GEMV, and the QMoE op wiring that replaces the phase-1 dequant fallback. The key correctness fix: the warp_frag_B double-buffer index was derived from the per-CTA-K-stage load offset, which stays coherent only when kWarpGemmIterationsForB is even (int8=4, int4=2). For int2 it is odd (=1), so every stage re-read stage-0's B fragment (tile0 counted twice, tile1 dropped), producing an ORT[m]=ref[m]+ref[m+64] K-fold for K>64. Fixed with a persistent cross-stage parity in all four DqMma variants (multistage/pipelined x finegrained/percol); the change is a no-op for even kWarpGemmIterationsForB (int4/int8). Also add the GeGLU gated activation (gelu-tanh gate), for Gemma4 MoE. Validation (H100, fp16, hidden=inter=2048, E=8, top2, block=128): - 2-bit parity max_diff 0.000488/0.003906 (tolerance 0.35). - Full test_qmoe_cuda.py: 122 passed, 0 failed (no int4/int8/fp4/fp8 regression). - int2 decode 1.019 -> 0.037 ms (tracks int4 0.036); prefill 1.047 -> 0.101 ms (beats int4 0.161).
Add an optional ``zero_point_offset`` float attribute to the QMoE op for
integer block-wise quantization when no uint8 fc*_zero_points tensor is
provided. It selects an asymmetric dequant (code - zero_point_offset) *
scale using a single fractional center, enabling balanced schemes whose
zero-point is not integer-representable by a uint8 tensor -- e.g. a 2-bit
checkpoint quantized around the midpoint 1.5, giving codes {0,1,2,3} ->
{-1.5,-0.5,0.5,1.5}*scale.
Implementation: the fused int2 weight converter already subtracts the
symmetric center 2^(bits-1); a constant per-element bias =
(2^(bits-1) - zero_point_offset) * scale corrects it to the requested
fractional center. The bias is built at PrePack time from the packed
scales (PrePackConstantBiasFromScales) and consumed via the same slot
int4-asymmetric uses; ComputeInternal has a fallback that builds it from
the live scales when prepacking is disabled.
- contrib_defs.cc: new OPTIONAL_VALUE ``zero_point_offset`` FLOAT attr.
- qmoe_kernels.cu/.h: LaunchQMoEConstantBias (float/half/bf16) computing
bias = delta * scale.
- moe_quantization.cc/.h: attr parse + ORT_ENFORCE(int && block_size>0),
PrePack wiring for fc1/fc2, ComputeInternal fallback branch.
- test_qmoe_cuda.py: TestQMoEFractionalZeroPoint 2-bit case (offset 1.5).
…fset - Reject packed int2 fc*_zero_points (unsupported 4-per-byte layout) in both the PrePack and Compute paths instead of misreading an undersized bias buffer. - Reject the CUDA-only zero_point_offset attribute on the CPU and WebGPU QMoE kernels so those EPs cannot silently ignore the documented dequant bias. - Document geglu in the QMoE activation_type schema and add the zero_point_offset attribute to docs/ContribOperators.md. - Add a live-scale (disable_prepacking) fractional-zp test covering the on-the-fly bias branch, and correct the stale int2 dequant rationale in cuda_quantizer.py.
…nputs The prior version set session.disable_prepacking to reach the on-the-fly bias branch, but int2 mandates weight PrePack (the CUTLASS layout transform), so the kernel correctly rejects weights_prepacked=0. Instead feed fc*_scales as runtime graph inputs (weights stay initializers -> int2 PrePack still runs, live scales are not consumed -> the constant bias is built on the fly from live scales). Validated on H100: full test_qmoe_cuda.py 125 passed / 15 skipped / 0 failed.
eaeb8af to
233e7ad
Compare
|
Heads-up on the latest force-push: I rebased this PR onto current
This PR originally introduced its own copies of those int2 primitives from scratch. Since they now exist upstream, I dropped the 15 duplicated primitive files in favor of What remains here is the piece Validation on SM90: clean build (0 errors), full QMoE test suite 125 passed / 12 skipped / 0 failed. |
Description
Adds three related capabilities to the QMoE CUDA execution provider:
cutlass::uint2b_t) mixed-input GEMM. 2-bit weights stay packed in HBM and are dequantized in-register inside the SM80 DqMma pipeline, mirroring the existing int4 (uint4b_t) path — no full-precision weight materialization. This replaces the phase-1 dequant-to-fp16 fallback that materialized[E, N, K]fp16 weights every call.zero_point_offsetattribute for integer block-wise quant, enabling balanced asymmetric schemes whose zero-point is not integer-representable by a uint8 tensor (e.g. a 2-bit checkpoint quantized around the 1.5 midpoint: codes{0,1,2,3}→{-1.5,-0.5,0.5,1.5}*scale).Motivation
The prior 2-bit path fell back to dequantizing the full weight tensor to fp16 in HBM, forfeiting the entire point of 2-bit (memory bandwidth). It ran ~24x slower than int4 at decode. The fused kernel keeps weights packed and dequantizes in-register.
Changes
In-register converter / layout / DqMma wiring (patterned on the
uint4b_tsites):cutlass_extensions/interleaved_numeric_conversion.h: from-scratchFastInterleavedAndBiasedNumericArrayConverter<half_t/bfloat16_t, uint2b_t, N>(16 codes/word, symmetric +2 bias, fp16/bf16 magic).cutlass_extensions/gemm/kernel/mixed_gemm_B_layout.h:LayoutDetailsB<uint2b_t>(ThreadblockK=64,ColumnsInterleaved=8).gemm/threadblock/default_mma*.h,default_dq_mma_*.h, the fourdq_mma_*variants:uint2b_tspecializations + static_assert extensions.Correctness fix (the crux): the
warp_frag_Bdouble-buffer index was derived from the per-CTA-K-stage load offset, which stays coherent only whenkWarpGemmIterationsForBis even (int8=4, int4=2). For int2 it is odd (=1), so every stage re-read stage-0's B fragment (tile0 counted twice, tile1 dropped) — anORT[m]=ref[m]+ref[m+64]K-fold for K>64. Fixed with a persistent cross-stage parity carried in all four DqMma variants; the change is a no-op for evenkWarpGemmIterationsForB(int4/int8 unaffected).Weight preprocessor / runner / decode GEMV:
fpA_intB_gemm_preprocessors*:W2_A16quant type, 2-bit subbyte transpose,add_bias_and_interleave_int2s, LDSM permutation map.moe_gemm/moe_gemm_kernels_{fp16,bf16}_uint2.cu,moe_kernels.cu,moe_gemm_template_dispatch.h:uint2b_trunner instantiations.fpA_intB_gemv/details.h,moe_gemv.cu: int2 decode GEMV path.QMoE op wiring:
moe/moe_quantization.cc/.h: construct the fuseduint2b_trunner for 2-bit; re-enable int2 PrePack (W2_A16); addzero_point_offsetparse +PrePackConstantBiasFromScales(bias =(2^(bits-1) - zero_point_offset) * scale) and a ComputeInternal fallback when prepacking is disabled.moe/qmoe_kernels.cu/.h:LaunchQMoEConstantBias(float/half/bf16).core/graph/contrib_ops/contrib_defs.cc:zero_point_offsetOPTIONAL_VALUE FLOAT attr on the QMoE schema.Quantizer / tests:
python/tools/quantization/cuda_quantizer.py: 2-bit packing support.test/python/transformers/test_qmoe_cuda.py: fused int2/int4/int8 GeGLU coverage +TestQMoEFractionalZeroPoint(2-bit, offset 1.5).Testing
Validated on H100 (SM90), fp16, hidden=inter=2048, E=8, top_k=2, block_size=128:
test_qmoe_cuda.py: 122 passed, 15 skipped, 0 failed (no int4/int8/fp4/fp8 regression).Opening as draft for review.