Repository navigation
perf(moe): implement the fused decode-MoE kernel on CUDA - #319
Conversation
The fused single-token decode-MoE kernel (#268) was Metal-only (mx.fast.metal_kernel), so on the CUDA backend it threw "[metal_kernel] No Metal back-end" and aborted for every small-expert MoE that the dispatch selected (qwen2_moe, qwen3_moe, lfm2, qwen3_vl_moe); larger experts fell through the dff break-even to gather_qmm and survived. Port both halves of the kernel (gate/up + SwiGLU/GeGLU, then down + score) to mx.fast.cuda_kernel: one warp owns each output row, the simd_sum reduction becomes a __shfl_down_sync warp reduction, and the quantized-affine dequant is unchanged. run_fused_moe_two_kernel selects the cuda_kernel port when metal::is_available() is false (no_cuda.cpp / no_metal.cpp stub the unused side, so both link on either backend). Precise expf in the SwiGLU holds greedy parity. Greedy byte-identical to the gather_qmm path on qwen3-30b-a3b. Measured on GB10 (DGX Spark, CUDA 13.0), decode tok/s vs gather_qmm: qwen3-moe-4bit 58.15 -> 90.34 (1.55x) qwen3.5-35b-a3b-4bit 45.84 -> 64.69 (1.41x) lfm2-8b-a1b-4bit 140.91 -> 158.22 (1.13x) qwen1.5-moe-a2.7b-4bit 113.03 -> 124.73 (1.10x) This supersedes the earlier CUDA fused-off workaround: fused stays default-on for both backends, and the previously-crashing MoE models now run faster instead of falling back. Closes #313
e857236 to
7880e4d
Compare
|
Force-pushed: this PR now implements the CUDA fused decode-MoE kernel instead of the earlier default-off workaround. Porting the Metal kernel to |
Record the CUDA port in the kernel design doc: a GB10 measured-gains section (qwen3-30b-a3b +55%, qwen3.5-35b-a3b +41%, lfm2 +13%, qwen1.5-moe +10%), the break-even sweep (CUDA crossover ~13-14k vs Metal's ~4096, so the 4096 cap is conservative on CUDA but kept as the shared default), a roadmap entry, and the default-on-both-backends note. Update environment-variables.md: MLXCEL_FUSED_MOE now spans both backends, and document MLXCEL_FUSED_MOE_MAX_DFF with the per-backend break-even.
|
Kept |
First GB10 benchmark on the merged CUDA fused decode-MoE kernel (#319): 148 text models (133 pass / 14 fail / 1 OOM-skip), 53 VLM rows. Nine fused-MoE models flipped FAIL→pass; qwen3-moe-30b (89.84) edges past M1 Ultra (83.75). Regenerates model_tests_gb10.md + index for 0.3.1.
Summary
The fused single-token decode-MoE kernel (#268) was Metal-only (
mx.fast.metal_kernel), so on the CUDA backend it threw[metal_kernel] No Metal back-endand aborted for every small-expert MoE the dispatch selected (qwen2_moe, qwen3_moe, lfm2, qwen3_vl_moe). Larger experts fell through thedff > MLXCEL_FUSED_MOE_MAX_DFFbreak-even togather_qmmand survived, so the failure tracked expert size, not architecture.This PR ports the kernel to CUDA instead of disabling it, so the affected models run faster on CUDA rather than merely working.
What changed
Port both halves to
mx.fast.cuda_kernel(the CUDA analogue ofmetal_kernel):act_g[K, Dff].simd_sumbecomes a__shfl_down_syncwarp reduction; the quantized-affine dequant (4/8-bit gate/up, 4/6/8-bit down) is unchanged.run_fused_moe_two_kernelselects the cuda_kernel port whenmetal::is_available()is false.no_cuda.cpp/no_metal.cppstub the unused side, so both link on either backend.expfin the SwiGLU holds greedy parity.Only
src/lib/mlxcel-core/cpp/mlx_cxx_kernels.cppchanges;fused_moe_enabledstays default-on for both backends.Validation (GB10 / DGX Spark, CUDA 13.0)
Greedy byte-identical to the
gather_qmmpath on qwen3-30b-a3b.The previously-crashing MoE models now run at these fused rates by default (no env var).
Supersedes the earlier CUDA fused-off workaround. Relevant to the #307 fused-MoE wiring epic: the wired families now work on CUDA.
Closes #313