You signed in with another tab or window. Reload to refresh your session.You signed out in another tab or window. Reload to refresh your session.You switched accounts on another tab or window. Reload to refresh your session.Dismiss alert
{{ message }}
Repository navigation
perf(cuda/moe): backend-specific default for the fused decode-MoE Dff threshold #626
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
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).
Change the default resolution to: env var if set, else 4096 when Metal is available, else the CUDA crossover value determined in step 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.
Set the CUDA default from the measurement (round down conservatively). If the crossover moved because of upstream changes, document the new number.
Update docs/environment-variables.md and the design doc with the backend-specific defaults.
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).
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_kerneland is selected at runtime by!metal::is_available()(src/lib/mlxcel-core/cpp/mlx_cxx_kernels.cpp,use_cudaselection around lines 158 and 1643). It declines and falls back togather_qmmwhen the expert FFN dimension exceedsMLXCEL_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_DFFas an explicit override for both backends.Implementation plan
MLXCEL_FUSED_MOE_MAX_DFF, consumed insrc/lib/mlxcel-core/cpp/mlx_cxx_kernels.cppnear the fused-MoE gate; also documented indocs/benchmark_results/fused-moe-decode-kernel-design.mdline ~171 anddocs/environment-variables.md).gather_qmmfallback 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). Usescripts/bench_decode.shand, where useful,scripts/capture_moe_decode_trace.sh. Force each side with the env var to get both curves.docs/environment-variables.mdand the design doc with the backend-specific defaults.Acceptance criteria
(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.fused_moe_max_dff_default_is_backend_aware_and_env_overrides).docs/benchmark_results/fused-moe-decode-dff-sweep-gb10-2026-07-03.csvplus a dated addendum in the design doc.-p mlxcel-coresuite has three failures that pre-date this PR and are unrelated to the additivemetal_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
References
src/lib/mlxcel-core/cpp/mlx_cxx_kernels.cpp(cuda_kernel port,use_cudagate, Dff decline).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).