Skip to content

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

Merged
inureyes merged 1 commit into
mainfrom
feature/issue-626-cuda-fused-moe-dff
Jul 2, 2026
Merged

inureyes merged 1 commit into
mainfrom
feature/issue-626-cuda-fused-moe-dff

Conversation

@inureyes

@inureyes inureyes commented Jul 2, 2026

Copy link
Copy Markdown
Member

Summary

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 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-tested fused_moe_max_dff_from(env, metal_available); the default is now FUSED_MOE_MAX_DFF_METAL (4096) when Metal is available and FUSED_MOE_MAX_DFF_CUDA (8192) otherwise. An explicit MLXCEL_FUSED_MOE_MAX_DFF still wins on both backends. Added a unit test for the backend-aware default and env override.
  • src/lib/mlxcel-core: new metal_is_available() FFI (cxx bridge decl in src/lib.rs, header decl in cpp/mlx_cxx_bridge.h, impl in cpp/mlx_cxx_kernels.cpp) mirroring the mlx::core::metal::is_available() gate the kernel dispatch already uses, so the Rust default follows the live backend without a compile-time cfg.
  • Docs: docs/environment-variables.md and docs/benchmark_results/fused-moe-decode-kernel-design.md updated with the backend-specific defaults and a dated (2026-07-03) sweep addendum; raw per-run data committed at docs/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 with MLXCEL_FUSED_MOE_MAX_DFF=1, fused with =20000.

Dff model gather_qmm fused fused/fallback
768 qwen3-30b-a3b 91.00 91.34 1.00
1792 lfm2-8b-a1b 140.35 157.76 1.12
2880 gpt-oss-20b (mxfp4 control) 78.59 78.00 0.99
6400 phi-3.5-moe 51.09 53.62 1.05
8192 llama-4-scout-17b 21.23 21.08 0.99
14336 mixtral-8x7b 28.23 27.89 0.99

The crossover collapsed from the old ~13-14k to ~8000: the #625 pin bump moved gather_gemm to 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:

model Dff before (4096 default, gather_qmm) after (8192 default, fused) delta
phi-3.5-moe-4bit 6400 51.09 53.62 +5.0%
llama-4-scout-17b-4bit 8192 21.23 21.08 break-even

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 cuda
  • New unit test fused_moe_max_dff_default_is_backend_aware_and_env_overrides passes (cargo test --release --features cuda -p mlxcel fused_moe_max_dff -- --test-threads=1: 1 passed)
  • Fused-vs-fallback greedy parity characterized on qwen3-30b-a3b (see Parity)
  • cargo test --release --features cuda -p mlxcel-core -- --test-threads=1: three failures pre-date this PR and are unrelated to the additive metal_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) and test_fused_paged_decode_gqa_and_batched (a Metal-only kernel that throws [metal_kernel] No Metal back-end and 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.rs is untouched here.

Closes #626

…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.
@inureyes inureyes added type:performance Performance improvements priority:medium Medium priority area:models Model architectures, weights, loading, metadata area:core mlxcel-core: MLX FFI, primitives, KV cache, layers status:done Completed labels Jul 2, 2026
@inureyes
inureyes merged commit 37931d7 into main Jul 2, 2026
5 checks passed
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
@inureyes
inureyes deleted the feature/issue-626-cuda-fused-moe-dff branch August 4, 2026 12:07
@inureyes inureyes self-assigned this Aug 31, 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
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

area:core mlxcel-core: MLX FFI, primitives, KV cache, layers area:models Model architectures, weights, loading, metadata priority:medium Medium priority status:done Completed type:performance Performance improvements

Projects

None yet

Development

Successfully merging this pull request may close these issues.

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

1 participant