onnxruntime
97aba237 - [CUDA] Use FP8 e4m3 weights for the DSV4 DeepGEMM MoE path

Commit
61 days ago
[CUDA] Use FP8 e4m3 weights for the DSV4 DeepGEMM MoE path The DSV4 MoE decode GEMMs are DRAM bound, so the prepacked weight format sets their cost. TryBuildDsv4DeepGemmWeights dequantized the stored MXFP4 weights all the way to bf16, moving 2 bytes per weight. Convert to fp8 e4m3 with a per-[128 N, 128 K] block scale instead and run DeepGEMM's sm90 fp8 masked grouped kernel. The conversion is bit exact. An E2M1 code carries at most two significant bits and e4m3 carries four, so a *power-of-two* block scale only shifts exponents and never disturbs a mantissa. Checked over all 256 experts of all 46 layers: every weight reproduces the fp32 dequantization bitwise, with no underflow and no clipping. The headroom is the group exponent spread within a block, which may reach 14 binades; the measured maximum is 6. Note the conventional amax/448 block scale is *not* usable here: it is not a power of two, so every weight would acquire a full mantissa before being rounded back to three bits (4.8% max relative error). The sm90 fp8 kernel takes both operands in fp8, so activations are quantized too. That work folds into the existing pack and SwiGLU kernels, which already read and write exactly this data, so it is nearly free. Activations use an amax/448 scale per token per 128-channel chunk, where the full e4m3 range is worth more than an exact scale. Measured on 8xH200 at prompt 1024 / generation 512, per decode step: 27.59 ms -> 24.43 ms, a 11.5% reduction, and 32 GiB per rank of weights freed. Standalone the FC1 GEMM goes 70.9 -> 37.9 us and FC2 38.6 -> 22.3 us. Smaller tiles also soften the wave quantization cliff: going from 8 to 9 active experts costs bf16 35% but fp8 only 18%. MMLU-Pro (800 samples) is unchanged at 0.660 vs 0.6625, 2 disagreements, McNemar p=0.48. Rank argmax disagreements remain 0.
Author
Committer
Parents
Loading