Skip to content

perf(cuda/moe): backend-specific default for the fused decode-MoE Dff threshold #626

Description

@inureyes

Part of #623. Depends on #625 (the crossover point must be re-measured after the MLX pin upgrade, because upstream gather_gemm JIT work moves the fallback's cost).

Context

The fused single-token decode-MoE kernel (issue #319 line of work) runs on CUDA via mlx::core::fast::cuda_kernel and is selected at runtime by !metal::is_available() (src/lib/mlxcel-core/cpp/mlx_cxx_kernels.cpp, use_cuda selection around lines 158 and 1643). It declines and falls back to gather_qmm when the expert FFN dimension exceeds MLXCEL_FUSED_MOE_MAX_DFF, whose default is 4096.

That default was tuned for Metal. The CUDA crossover was measured at Dff ~13-14k (docs/benchmark_results/fused-moe-decode-kernel-design.md, around lines 222-257), and the measured wins on GB10 where the kernel does engage are large: qwen3-30b-a3b 58.2 -> 90.3 tok/s (+55%), qwen3.5-35b +41%, lfm2-8b +13%. With the default at 4096, CUDA users silently lose the kernel on models whose Dff falls between 4096 and the true crossover.

Scope

Make the threshold default backend-aware: keep 4096 on Metal, raise the CUDA default to the measured crossover. Keep MLXCEL_FUSED_MOE_MAX_DFF as an explicit override for both backends.

Implementation plan

  1. Locate the threshold read (env MLXCEL_FUSED_MOE_MAX_DFF, consumed in src/lib/mlxcel-core/cpp/mlx_cxx_kernels.cpp near the fused-MoE gate; also documented in docs/benchmark_results/fused-moe-decode-kernel-design.md line ~171 and docs/environment-variables.md).
  2. Change the default resolution to: env var if set, else 4096 when Metal is available, else the CUDA crossover value determined in step 3.
  3. Re-measure the crossover on GB10 post pin-upgrade: sweep Dff by benchmarking the fused kernel vs the gather_qmm fallback on the MoE set (qwen3-30b-a3b-4bit, qwen3.5-35b-a3b-4bit, gpt-oss-20b-mxfp4, lfm2-8b-a1b-4bit, phi-3.5-moe-4bit, mixtral-8x7b-4bit). Use scripts/bench_decode.sh and, where useful, scripts/capture_moe_decode_trace.sh. Force each side with the env var to get both curves.
  4. Set the CUDA default from the measurement (round down conservatively). If the crossover moved because of upstream changes, document the new number.
  5. Update docs/environment-variables.md and the design doc with the backend-specific defaults.

Acceptance criteria

  • On CUDA the new 8192 default fuses two MoE models in (4096, 8192] by default: phi-3.5-moe (Dff 6400, 51.09 -> 53.62 tok/s, +5%) and llama-4-scout-17b (Dff 8192, 21.23 -> 21.08, break-even). The crossover re-measured to ~8000 (collapsed from ~13-14k after the chore(mlx): upgrade pinned MLX commit, rebase CUDA overlays, GB10 re-baseline and outlier triage #625 JIT gather_gemm), so mixtral (14336) now regresses if forced and is left on gather_qmm; before/after recorded in PR perf(cuda/moe): backend-aware default for the fused decode-MoE Dff threshold #643.
  • Metal default unchanged (4096); the env var overrides both backends (unit test fused_moe_max_dff_default_is_backend_aware_and_env_overrides).
  • Sweep data committed: docs/benchmark_results/fused-moe-decode-dff-sweep-gb10-2026-07-03.csv plus a dated addendum in the design doc.
  • Parity re-run on qwen3-30b-a3b: fused stays within gather_qmm's own run-to-run envelope (both non-deterministic on GB10/CUDA; identical ~30-token prefix). New unit test passes. The full -p mlxcel-core suite has three failures that pre-date this PR and are unrelated to the additive metal_is_available() FFI (pin e9463bb steel_gemm drift, fp16 dtype-code 9 vs 10, and a Metal-only paged-decode test that aborts on CUDA); see PR perf(cuda/moe): backend-aware default for the fused decode-MoE Dff threshold #643.

Validation

cargo build --release --features cuda
MLXCEL_FUSED_MOE_MAX_DFF=4096  ./target/release/mlxcel-bench-decode --model ./models/qwen3.5-35b-a3b-4bit   # fallback side
MLXCEL_FUSED_MOE_MAX_DFF=16384 ./target/release/mlxcel-bench-decode --model ./models/qwen3.5-35b-a3b-4bit   # fused side
./target/release/mlxcel-bench-decode --model ./models/qwen3.5-35b-a3b-4bit                                   # new default = fused

References

  • Kernel and gate: src/lib/mlxcel-core/cpp/mlx_cxx_kernels.cpp (cuda_kernel port, use_cuda gate, Dff decline).
  • Measurements: docs/benchmark_results/fused-moe-decode-kernel-design.md (CUDA gains and crossover), docs/benchmark_results/moe-decode-gap-investigation.md (why MoE decode is GPU-bound).

Activity

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, metadatapriority: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