Repository navigation
perf(cuda/moe): fix MoE prefill collapse on CUDA (gather_qmm / grouped GEMM path) #629
Copy link
Copy link
Closed
Labels
area:coremlxcel-core: MLX FFI, primitives, KV cache, layersmlxcel-core: MLX FFI, primitives, KV cache, layersarea:modelsModel architectures, weights, loading, metadataModel architectures, weights, loading, metadatabugpriority:mediumMedium priorityMedium prioritystatus:doneCompletedCompletedtype:performancePerformance improvementsPerformance improvements
Description
Activity
- addedtype:performancePerformance improvementsPerformance improvementsstatus:readyReady to be worked onReady to be worked onpriority:mediumMedium priorityMedium priorityarea:coremlxcel-core: MLX FFI, primitives, KV cache, layersmlxcel-core: MLX FFI, primitives, KV cache, layersarea:modelsModel architectures, weights, loading, metadataModel architectures, weights, loading, metadata
on Jul 2, 2026 - added a commit that references this issue
on Jul 2, 2026 Phase 1 (#625) triage confirms this is still required (docs/benchmark_results/mlx-pin-upgrade-2026-07-03.md, PR #642). MoE prefill is unchanged by the pin bump: mixtral-8x7b 12.6 tok/s (baseline 12.5, Metal 81), llama-4-scout 27.5 (baseline 28.2, Metal 119), phi-3.5-moe 29.1, gpt-oss-20b 127. Root cause pointer: our
patches/mlx/backend/cuda/matmul.cppoverlay routes GatherMM through CUTLASS and bypasses upstream MLX #3706 which JIT-compiledgather_gemm, so the upstream rework never reaches this path. This issue should revisit the CUTLASS-GatherMM vs upstream-JIT-gather_gemm decision on MoE prefill.- added a commit that references this issue
on Jul 2, 2026 - addedstatus:in-progressCurrently being worked onCurrently being worked onstatus:reviewUnder reviewUnder reviewand removedstatus:readyReady to be worked onReady to be worked onstatus:in-progressCurrently being worked onCurrently being worked on
on Jul 10, 2026 - added a commit that references this issue
on Jul 10, 2026 - addedstatus:doneCompletedCompletedand removedstatus:reviewUnder reviewUnder review
on Jul 10, 2026 - added a commit that references this issue
on Jul 30, 2026 - added a commit that references this issue
on Aug 20, 2026 - added a commit that references this issue
on Sep 1, 2026
Metadata
Metadata
Assignees
Labels
area:coremlxcel-core: MLX FFI, primitives, KV cache, layersmlxcel-core: MLX FFI, primitives, KV cache, layersarea:modelsModel architectures, weights, loading, metadataModel architectures, weights, loading, metadatabugpriority:mediumMedium priorityMedium prioritystatus:doneCompletedCompletedtype:performancePerformance improvementsPerformance improvements
Part of #623. Depends on #625: re-triage first. Upstream #3706 (JIT gather_gemm) and #3632 (gather_qmm NAX name fix) may partially or fully resolve this; the pin-upgrade triage doc states what remains.
Context
MoE prefill on CUDA is 5-10x slower than on Metal M1 Ultra even though GB10 has more compute. From
benchmarks/cuda_gb10_2026-06-17.csvvs the 06-15 Metal sweep (prefill tok/s):Dense prefill on the same hardware is fine (CUDA beats Metal by ~1.6x median across 133 common models), so this is specific to the expert-routing matmul path:
gather_qmm(src/lib/mlxcel-core/cpp/mlx_cxx_bridge.cpp:2496-2522, model call sites e.g.src/models/deepseek_v3.rs:663-702,qwen3_moe.rs:134) and the CUDA grouped-GEMM overlay (src/lib/mlx-cpp/patches/mlx/backend/cuda/gemms/grouped_gemm_unaligned.cu,matmul.cpp). Note prefill hits the matrix-matrix shape (many tokens per expert), unlike decode's GEMV shape, and the quantized dispatch on CUDA falls through toqmm_sm80/qmm_naiveon GB10 becauseqmm_sm90is Hopper-gated (patches/mlx/backend/cuda/quantized/quantized.cpp:53-113).Scope
Bring CUDA MoE prefill to at least Metal-M1-Ultra absolute throughput on the table above (GB10 should beat M1 Ultra on any compute-bound path). Decode-side MoE is covered by #626 and #636, not here.
Implementation plan
nsys profile ./target/release/mlxcel-bench-decode ...) and inspect: which kernels dominate (qmm_naive? dequant + plain gemm? per-expert loop of small qmm launches?), whether experts are processed as one grouped GEMM or as a Python-style per-expert loop, and whether host-side sorting/segmentation is serializing.scripts/bench_mlxlm.py) to separate "MLX kernel problem" from "mlxcel bridge problem":gather_qmminvocation (argument layout,ensure_row_contiguousforcing copies of expert weight tensors per call, sorted-indices contract) is the suspect.a. Adopt/extend upstream JIT
gather_gemm(post-#3706) so prefill uses a real grouped GEMM per expert group.b. Token-sorting by expert on device before the grouped call so each expert processes a contiguous token block (standard vLLM-style segmentation), if the current path degenerates to per-token gathers.
c. Repair our
grouped_gemm_unaligned.cuoverlay for sm_121 tile configurations if it is being skipped and silently falling back toqmm_naive.Acceptance criteria
docs/benchmark_results/note with nsys evidence of the removed bottleneck.cargo test --features cudapasses.Validation
References
src/lib/mlxcel-core/cpp/mlx_cxx_bridge.cpp:2496-2522(gather_qmm), fused switch-GeGLU1805-1891.src/lib/mlx-cpp/patches/mlx/backend/cuda/gemms/,patches/mlx/backend/cuda/quantized/quantized.cpp:53-113.scripts/gather_qmm_shape_microbench.py.docs/benchmark_results/moe-decode-gap-investigation.md.