Repository navigation
perf(cuda/moe): backend-aware default for the fused decode-MoE Dff threshold - #643
Merged
Merged
Conversation
…reshold The fused single-token decode-MoE kernel declines to gather_qmm when the expert intermediate (Dff) exceeds MLXCEL_FUSED_MOE_MAX_DFF, whose default was a flat 4096 tuned for Metal. On CUDA the crossover is higher, so mid-size-expert models such as phi-3.5-moe (Dff 6400) silently lost the fused kernel by default. Make the default backend-aware: 4096 when Metal is available, 8192 otherwise, with the env var still overriding both backends. The Dff-bound resolution moves into a pure, unit-tested fused_moe_max_dff_from(env, metal_available); the live backend is read through a new metal_is_available() FFI that mirrors the mlx::core::metal::is_available() gate the kernel dispatch already uses, so the default follows the actual device instead of a compile-time cfg. Re-measured the CUDA crossover on GB10 (DGX Spark, sm_121) under MLX pin e9463bb. The #625 pin bump moved gather_gemm to a JIT path that sped the gather_qmm fallback dramatically for the pathological 128-expert config (qwen3-30b-a3b fallback 58.2 -> 91.0 tok/s, erasing its old +55% fused win) while leaving low-expert configs unchanged (lfm2 still +12%). Fused now clearly wins through Dff 6400 (phi-3.5-moe +5%), is break-even at 8192 (llama-4-scout), and loses at 14336 (mixtral -1.2%), so the crossover collapsed from the old ~13-14k to ~8000. The CUDA default is set to the 8192 break-even boundary, which captures phi-3.5-moe and llama-4-scout while leaving mixtral on gather_qmm. Docs and the committed per-run sweep CSV under docs/benchmark_results/ record the new backend-specific defaults and the 2026-07-03 GB10 sweep.
4 tasks done
This was referenced Jul 9, 2026
inureyes
added a commit
that referenced
this pull request
Jul 9, 2026
…711) ## Summary Complete the backend-aware fused-MoE `Dff` cap for #330 by recording the honest provenance of the CUDA default and pinning the decline boundary with a test. The cap-resolution logic itself (Metal 4096 / CUDA 8192, the `MLXCEL_FUSED_MOE_MAX_DFF` override, the pure `fused_moe_max_dff_from`, and the 2026-07-03 GB10 sweep) already landed in #643; that work was measured under MLX pin `e9463bb` (#626) and was never linked to #330. Since #643 the MLX pin advanced to 0.32.1 (#703/#704, commit `57c66cac`), so the CUDA default is now a carried-forward documented-data value, not one confirmed on current binaries. This PR closes that honesty gap, documents the runtime-vs-`cfg!` backend choice, adds a decline-boundary test, and formally closes #330. The cap value 8192 is unchanged. ## What changed - `src/models/switch_layers.rs`: expanded the `FUSED_MOE_MAX_DFF_CUDA` doc comment with a provenance caveat (the 8192 sweep ran on MLX pin `e9463bb`/#626; the pin has since moved to 0.32.1/#703/#704, commit `57c66cac`; 8192 is the conservative crossover floor kept so pin drift cannot flip it into a regression; a GB10 re-validation on 0.32.1 is pending). - `src/models/switch_layers.rs`: documented on `fused_moe_max_dff_from` that the cap is family-agnostic (governs every SwitchGLU MoE model per the module `Used by:` list) and that the backend is resolved at runtime via `metal_is_available()` rather than a `cfg!(feature = "cuda")` switch, matching the fused kernel's own runtime dispatch (`run_fused_moe_two_kernel` picks the `cuda_kernel` port when `metal::is_available()` is false). - `src/models/switch_layers.rs`: added `fused_moe_dff_above_cap_declines_and_at_cap_dispatches`, a pure test pinning that `forward_fused_kernel` declines exactly when `dff > max_dff` at the cap boundary on both backends and under an explicit override. - `docs/benchmark_results/fused-moe-decode-kernel-design.md`: added the 2026-07-09 addendum (CUDA default provenance, why keeping 8192 is safe under pin drift, and the pending GB10 re-validation on MLX 0.32.1 with the exact `mlxcel-bench-decode` harness to run it), and cross-referenced it from the `MLXCEL_FUSED_MOE_MAX_DFF` table row. ## Design notes - The value 8192 is unchanged. Under pin drift the most conservative floor of the measured crossover (break-even at 8192, regression only at 14336) is the correct choice, so no re-tuning is done without hardware re-validation. - Metal behavior is byte-identical by construction (still 4096). The diff is runtime-inert (doc comments, a doc file, and a test only); no compiled non-test code path changed. ## Test plan Ran here (Apple M1 Ultra, Metal): - [x] `cargo fmt --all -- --check` - [x] `cargo test --release --features metal,accelerate -p mlxcel --lib switch_layers` (8/8 pass, including the new `fused_moe_dff_above_cap_declines_and_at_cap_dispatches`; the successful release build subsumes `cargo check --lib --tests`) - [x] `cargo clippy --features metal,accelerate -p mlxcel --lib --tests -- -D warnings` (clean) - [x] Metal MoE decode smoke, dispatch unchanged: `mlxcel generate -m models/qwen1.5-moe-a2.7b-4bit -p "Hello" -n 32 --temp 0` produced coherent output at 96.58 tok/s on Apple GPU (Metal), exercising the fused MoE decode path (Dff below the 4096 Metal cap, so it dispatches). Documented-data (not re-run in this PR): the CUDA default 8192 rests on the GB10 sweep in the 2026-07-03 addendum under MLX pin `e9463bb` (#626), landed in #643. Pending (hardware-blocked): GB10 re-validation of the crossover on the current MLX 0.32.1 pin. No CUDA/GB10 hardware was reachable for this change; the how-to-run harness is recorded in the 2026-07-09 addendum. Closes #330
This was referenced Sep 14, 2026
inureyes
added a commit
that referenced
this pull request
Sep 14, 2026
## Summary `qwen3_moe` kept its own `SwitchLinear`, `SwitchGLU` and `forward_fused_kernel`, and `qwen3_vl_moe` (which also builds the Qwen3-Omni-MoE thinker) built that copy by hand. The copy never got the fused-kernel `dff` bound the shared type gained in #311 and #643, so `MLXCEL_FUSED_MOE_MAX_DFF` was silently ignored. Both families now hold `switch_layers::SwitchGLU` (`SwitchGLU::from_weights`, affine) and the copy is deleted rather than patched. - Loader changes: bits are inferred from tensor shapes (the copy stored the configured value), a missing `.biases` plane fails through `validate_quantization_biases` instead of `Weight not found`, and unstacked `experts.{idx}` checkpoints now reach the shared stacking fallback. Every expert plane of the four local Qwen3-MoE checkpoints (safetensors headers) solves to 4-bit with biases shaped like the scales. - Tests: the #958 guard tests drive each family's `SparseMoeBlock::from_weights`; new tests pin the Dff 8256 decline through both loaders (positive control at Dff 64 only when `custom_kernels_available()`, #1803) and the mxfp4 decline. - Docs: `switch_layers.rs`, `docs/environment-variables.md` and the fused-MoE design note now say which families read the bound, fix the family counts, and qualify the qwen3-30b-a3b greedy parity claim. ## Verification (GB10 sm_121, release `--features cuda`, base `b8d10fb1`) - Greedy ids (64 tokens, `-t 0`, as ids): identical base vs branch on `qwen3-30b-a3b-4bit`, `qwen3-vl-30b-a3b-4bit` and the `qwen3-omni-30b-a3b-instruct-4bit` thinker, default fused and `MLXCEL_FUSED_MOE=0`. Fused and gather diverge from each other at token 39 (qwen3-30b) and 26 (omni) on both binaries, so the fused match also rules out a silent fallback. - Trace (`MLXCEL_PROFILE_BLOCKS=1 MLXCEL_PROFILE_QWEN3_MOE_DETAIL=1`, `-n 4`): 144 `path=fused tokens=1`, 0 `path=gather_qmm tokens=1` on both binaries. - The bound now applies: with `MLXCEL_FUSED_MOE_MAX_DFF=512` the base trace stays 144 fused and the branch is 144 `path=gather_qmm tokens=1`; omni base ids equal its fused ids and branch ids equal its `MLXCEL_FUSED_MOE=0` ids; VL 64-token runs stay at 2.31 to 2.35 s on base (kernel launched) and drop to 1.28 to 1.33 s on the branch. - `mlxcel-server`, temperature 0: identical content and reasoning base vs branch on qwen3-30b and VL; the VL completion is coherent and opens with the CLI's tokens. - Decode tok/s, qwen3-30b (`mlxcel-bench-decode -n 128 --ignore-eos --warmup-tokens 20`), 15 interleaved runs per arm: base 66.21 to 87.76 (median 78.00), branch 61.27 to 86.28 (median 73.63). Decode here is bimodal; high-mode means 81.06 vs 79.65, Mann-Whitney p = 0.21. - 38 of 38 tests under `models::qwen3_moe`, `models::qwen3_vl_moe`, `models::switch_layers` pass, one process each; both new decline tests fail against a mutant with the bound disabled. Scoped clippy (`-p mlxcel --lib --tests --features cuda -D warnings`) and fmt are clean. Ids, trace, cap and server checks were repeated with binaries rebuilt at the final commit. - Not run: `metal,accelerate` gate (Linux host), workspace-wide `verify-test-cuda`. The `OpenXLA feature compile` CI failure also fails on `main` at `b8d10fb1` (unused imports in `models/mod.rs` and server modules) and is unrelated. ## Follow-up candidates `qwen3_next.rs` (same kernel copy and missing guards), `gemma4.rs` (GeGLU, no bound), NemotronH's opt-in `MLXCEL_FUSED_MOE_RELU2` path, ten families with a local `SwitchGLU` and no fused path, a cached bound instead of a per-call env read, and pre-existing shared-loader gaps found in review (no cross-plane or router-vs-expert-count check; the unstacked fallback does not compare shapes before `stack`). Closes #1884
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
Summary
The fused single-token decode-MoE kernel declines to
gather_qmmwhen the expert intermediate (Dff) exceedsMLXCEL_FUSED_MOE_MAX_DFF, whose default was a flat 4096 tuned for Metal. On CUDA the crossover is higher, so mid-size-expert models silently lost the kernel. This makes the default backend-aware: 4096 on Metal, 8192 on CUDA (re-measured on GB10), with the env var still overriding both backends.What changed
src/models/switch_layers.rs: split the Dff-bound resolution into a pure, unit-testedfused_moe_max_dff_from(env, metal_available); the default is nowFUSED_MOE_MAX_DFF_METAL(4096) when Metal is available andFUSED_MOE_MAX_DFF_CUDA(8192) otherwise. An explicitMLXCEL_FUSED_MOE_MAX_DFFstill wins on both backends. Added a unit test for the backend-aware default and env override.src/lib/mlxcel-core: newmetal_is_available()FFI (cxx bridge decl insrc/lib.rs, header decl incpp/mlx_cxx_bridge.h, impl incpp/mlx_cxx_kernels.cpp) mirroring themlx::core::metal::is_available()gate the kernel dispatch already uses, so the Rust default follows the live backend without a compile-time cfg.docs/environment-variables.mdanddocs/benchmark_results/fused-moe-decode-kernel-design.mdupdated with the backend-specific defaults and a dated (2026-07-03) sweep addendum; raw per-run data committed atdocs/benchmark_results/fused-moe-decode-dff-sweep-gb10-2026-07-03.csv.Re-measured CUDA crossover (GB10, DGX Spark, sm_121, MLX pin e9463bb)
mlxcel-bench-decode, 100 decode tokens after 20-token warmup, median of 3; fallback forced withMLXCEL_FUSED_MOE_MAX_DFF=1, fused with=20000.The crossover collapsed from the old ~13-14k to ~8000: the #625 pin bump moved
gather_gemmto JIT, which sped the fallback dramatically for the pathological 128-expert config (qwen3-30b-a3b fallback 58.2 -> 91.0, erasing its old +55% fused win) while leaving low-expert configs unchanged (lfm2 still +12%). Fused clearly wins through Dff 6400, is break-even at 8192, and loses at 14336. CUDA default set to the 8192 break-even boundary (rounded down from the 14336 regression); Metal stays 4096.Before / after by default (CUDA)
Models with Dff in (4096, 8192] that the old 4096 default declined and the new 8192 default now fuses by default:
mixtral-8x7b (Dff 14336) stays on
gather_qmm(past the crossover, -1.2% if forced).Parity
#319-style fused-vs-fallback greedy (temp 0) parity re-run on qwen3-30b-a3b. This change touches only the threshold default, not the kernel math: models that now fuse (phi-3.5-moe, llama-4-scout) use the identical pre-existing fused kernel. On GB10/CUDA greedy temp-0 output is non-deterministic run-to-run for both paths (whole-model GPU FP-reduction order): two identical gather_qmm runs already diverge, as do two identical fused runs. All four runs share the same ~30-token prefix and then diverge at the same near-tie cascade point, so the fused kernel stays inside the same envelope gather_qmm has with itself, i.e. within the documented f16 jitter class. Byte-identical greedy parity is not achievable on CUDA for either path and is unaffected by this change.
Test plan
cargo build --release --features cudafused_moe_max_dff_default_is_backend_aware_and_env_overridespasses (cargo test --release --features cuda -p mlxcel fused_moe_max_dff -- --test-threads=1: 1 passed)cargo test --release --features cuda -p mlxcel-core -- --test-threads=1: three failures pre-date this PR and are unrelated to the additivemetal_is_available()FFI (which no failing test calls):steel_gemm_edge_tile_safe_load_matches_reference(numerical drift, max_abs 0.0198),test_from_bytes_fp16_native_dtype(dtype enum code 9 vs 10) andtest_fused_paged_decode_gqa_and_batched(a Metal-only kernel that throws[metal_kernel] No Metal back-endand aborts on the CUDA build). All three are MLX pin e9463bb / CUDA-environment issues introduced by the chore(mlx): upgrade pinned MLX commit, rebase CUDA overlays, GB10 re-baseline and outlier triage #625 pin bump;ffi_tests.rsis untouched here.Closes #626