Skip to content

perf(cuda/moe): fix MoE prefill collapse on CUDA (gather_qmm / grouped GEMM path) #629

Description

@inureyes

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.csv vs the 06-15 Metal sweep (prefill tok/s):

Model GB10 CUDA M1 Ultra Metal
mixtral-8x7b-4bit 12.5 81.0
llama-4-scout-17b-4bit 28.2 119.5
phi-3.5-moe-4bit 28.9 114.0
minimax-m2-3bit 26.9 68.0
gpt-oss-20b-mxfp4 126.2 282.4
solar-open-100b-4bit 46.8 71.0

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 to qmm_sm80/qmm_naive on GB10 because qmm_sm90 is 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

  1. Reproduce with the long-prompt harness (perf(bench): long-prompt prefill benchmark coverage and serving TTFT/decode-rate telemetry #624): mixtral-8x7b-4bit at 2048-token prompt is the primary repro; smaller repro phi-3.5-moe-4bit.
  2. Profile with Nsight Systems (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.
  3. Cross-check against Python mlx-lm on the same machine and pin (scripts/bench_mlxlm.py) to separate "MLX kernel problem" from "mlxcel bridge problem":
    • if mlx-lm is equally slow, the fix belongs in the CUDA kernel layer (upstream or our overlay);
    • if mlx-lm is fast, the bridge's gather_qmm invocation (argument layout, ensure_row_contiguous forcing copies of expert weight tensors per call, sorted-indices contract) is the suspect.
  4. Fix accordingly. Candidate directions, in preference order:
    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.cu overlay for sm_121 tile configurations if it is being skipped and silently falling back to qmm_naive.
  5. Parity: greedy 40-token continuation on mixtral and phi-3.5-moe must match the pre-change output (or HF reference where an oracle already exists in the test suite).

Acceptance criteria

  • mixtral-8x7b-4bit prefill >= 81 tok/s on GB10 at a 2048-token prompt (parity with M1 Ultra; stretch goal 2x that, consistent with the dense CUDA/Metal prefill ratio of ~1.6x).
  • All six models in the table improve; deltas recorded in a docs/benchmark_results/ note with nsys evidence of the removed bottleneck.
  • No decode regression on the MoE set (fused decode-MoE path untouched or re-verified).
  • Parity checks pass; cargo test --features cuda passes.

Validation

cargo build --release --features cuda
nsys profile -o moe_prefill ./target/release/mlxcel-bench-decode --model ./models/mixtral-8x7b-4bit --prompt-tokens 2048 --max-tokens 8
python3 scripts/bench_mlxlm.py --model ./models/mixtral-8x7b-4bit   # cross-check

References

  • Bridge: src/lib/mlxcel-core/cpp/mlx_cxx_bridge.cpp:2496-2522 (gather_qmm), fused switch-GeGLU 1805-1891.
  • CUDA overlays: src/lib/mlx-cpp/patches/mlx/backend/cuda/gemms/, patches/mlx/backend/cuda/quantized/quantized.cpp:53-113.
  • Shape microbench: scripts/gather_qmm_shape_microbench.py.
  • Prior investigation (decode side, methodology reusable): docs/benchmark_results/moe-decode-gap-investigation.md.

Activity

  1. added
    area:coremlxcel-core: MLX FFI, primitives, KV cache, layers
    area:modelsModel architectures, weights, loading, metadata
    on Jul 2, 2026
  2. inureyes commented on Jul 2, 2026

    @inureyes
    MemberAuthor

    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.cpp overlay routes GatherMM through CUTLASS and bypasses upstream MLX #3706 which JIT-compiled gather_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.

  3. added a commit that references this issue on Jul 30, 2026
    3b770b5
  4. self-assigned this
    on Aug 31, 2026
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

Labels

area:coremlxcel-core: MLX FFI, primitives, KV cache, layersarea:modelsModel architectures, weights, loading, metadatabugpriority:mediumMedium prioritystatus:doneCompletedtype:performancePerformance improvements

Type

No type

Projects

No projects

    Milestone

    No milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions