Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
@@ -0,0 +1,15 @@
# Fused decode-MoE vs gather_qmm Dff crossover sweep (#626).
# Host: GB10 (DGX Spark), sm_121, CUDA 13.0. MLX pin e9463bb.
# Harness: mlxcel-bench-decode, prompt "Hello, how are you today?", 100 decode
# tokens after a 20-token in-process warmup, median of 3 invocations per side.
# Fallback side forced with MLXCEL_FUSED_MOE_MAX_DFF=1 (gather_qmm);
# fused side forced with MLXCEL_FUSED_MOE_MAX_DFF=20000 (fused two-kernel path).
# decode tok/s. NA = model does not run in the text decode bench harness.
model,dff,quant_mode,fallback_run1,fallback_run2,fallback_run3,fallback_median,fused_run1,fused_run2,fused_run3,fused_median,fused_over_fallback
qwen3.5-35b-a3b-4bit,512,affine,NA,NA,NA,NA,NA,NA,NA,NA,NA
qwen3-30b-a3b-4bit,768,affine,91.00,91.66,90.44,91.00,88.43,91.91,91.34,91.34,1.0037
lfm2-8b-a1b-4bit,1792,affine,140.65,139.82,140.35,140.35,157.36,157.80,157.76,157.76,1.1240
gpt-oss-20b-mxfp4,2880,mxfp4,78.93,78.59,77.10,78.59,78.00,78.90,76.57,78.00,0.9925
phi-3.5-moe-4bit,6400,affine,49.44,51.96,51.09,51.09,53.62,53.41,53.78,53.62,1.0495
llama-4-scout-17b-4bit,8192,affine,21.23,21.23,21.28,21.23,21.00,21.27,21.08,21.08,0.9929
mixtral-8x7b-4bit,14336,affine,28.15,28.23,28.30,28.23,27.88,27.89,27.90,27.89,0.9880
63 changes: 55 additions & 8 deletions docs/benchmark_results/fused-moe-decode-kernel-design.md
Original file line number Diff line number Diff line change
Expand Up @@ -168,7 +168,7 @@ falls back to the proven `gather_qmm` / `SwitchGLU` path automatically. Set
| `MLXCEL_FUSED_MOE` | unset (on) | On by default. Set to `0` (also `false`/`off`/`no`, case-insensitive) to force the proven `gather_qmm` / `SwitchGLU` path; any other value, or leaving it unset, keeps the kernel on. |
| `MLXCEL_FUSED_MOE_SGY` | 8 | Simdgroups per threadgroup (one output row each). Tune per hardware. |
| `MLXCEL_FUSED_MOE_RELU2` | unset (off) | Enable the squared-ReLU (fc1/relu²/fc2) fused path used by nemotron-class experts. Correct but measured performance-neutral on nemotron-h; kept for a future MoE-dominated squared-ReLU model. |
| `MLXCEL_FUSED_MOE_MAX_DFF` | 4096 | Expert-intermediate (Dff) upper bound. Above it, `forward_fused_kernel` declines and the caller falls back to `gather_qmm`. The fused path wins only while `gather_qmm` underutilizes the GPU (small experts); for large experts `gather_qmm` already saturates and the extra dispatch plus global-memory activation staging is a net loss. M1 Ultra measurements: Dff 704..2560 gain +3.5% to +15.4%, Dff 6400 (phi-3.5-moe) loses 5.9%, Dff 14336 (mixtral) loses 21%. On CUDA (GB10) the crossover is much higher, ~13-14k (Dff 768 +60%, 1792 +13%, 6400 +2%, 14336 -2%), so 4096 is conservative there but kept as the shared default. The break-even is hardware-dependent, so the bound is tunable. |
| `MLXCEL_FUSED_MOE_MAX_DFF` | 4096 (Metal) / 8192 (CUDA) | Expert-intermediate (Dff) upper bound. Above it, `forward_fused_kernel` declines and the caller falls back to `gather_qmm`. The fused path wins only while `gather_qmm` underutilizes the GPU (small experts); for large experts `gather_qmm` already saturates and the extra dispatch plus global-memory activation staging is a net loss. The break-even is backend-dependent, so the default is chosen from the live backend (`mlx::core::metal::is_available()`). M1 Ultra measurements: Dff 704..2560 gain +3.5% to +15.4%, Dff 6400 (phi-3.5-moe) loses 5.9%, Dff 14336 (mixtral) loses 21%, so Metal keeps 4096. On CUDA the crossover is higher; re-measured on GB10 under MLX pin e9463bb (#626) it collapsed from the old ~13-14k to ~8000 (Dff 6400 +5%, 8192 break-even, 14336 -1.2%; see the 2026-07-03 addendum), so the CUDA default is 8192. An explicit value overrides both backends. |

### Measured decode gains (M1 Ultra, `MLXCEL_FUSED_MOE=1`)

Expand Down Expand Up @@ -250,11 +250,13 @@ Forcing the fused path past the cap (`MLXCEL_FUSED_MOE_MAX_DFF=20000`):
| 6400 | phi-3.5-moe | +2% |
| 14336 | mixtral | −2% |

The CUDA crossover is ~13–14k (vs Metal's ~4096, where phi-3.5-moe already loses
5.9%). The 4096 cap is therefore conservative on CUDA but kept as the shared
default: the meaningful wins are all below it, and the 4096–14k range gains at
most ~2% (phi-3.5-moe) while mixtral (14336) slightly prefers `gather_qmm`. CUDA
users with mid-size experts can raise `MLXCEL_FUSED_MOE_MAX_DFF`.
The CUDA crossover measured here (#319) was ~13–14k (vs Metal's ~4096, where
phi-3.5-moe already loses 5.9%). At that time 4096 was kept as the shared default
even on CUDA. That was re-measured after the #625 MLX pin bump (which moved
`gather_gemm` to JIT and sped the fallback): the CUDA crossover collapsed to
~8000 and the default is now backend-aware, 4096 on Metal and 8192 on CUDA. See
the 2026-07-03 addendum below for the new sweep; CUDA users with larger experts
can still raise or lower `MLXCEL_FUSED_MOE_MAX_DFF`.

**Parity caveat.** The kernel accumulates the GEMV in f32 with a different
reduction order than `gather_qmm`'s tiling, so it is within the f16 jitter
Expand Down Expand Up @@ -282,8 +284,9 @@ SwiGLU activation, text-only decode path), phimoe (Phi-3.5-MoE; migrated from
its local `SwitchGLU`/`SwitchLinear` to the shared ones; checkpoints pre-stacked
under `block_sparse_moe.switch_mlp.{gate,up,down}_proj`; `sanitize_weights` still
handles the unstacked `experts.{i}.w1/w2/w3` layout for community checkpoints; the
expert intermediate is 6400, above `MLXCEL_FUSED_MOE_MAX_DFF`, so like mixtral the
kernel declines and decode stays on `gather_qmm`), olmoe (OLMoE; migrated from
expert intermediate is 6400: below the CUDA default (8192) so the kernel dispatches
on CUDA at a measured +5%, above the Metal default (4096) so it declines to
`gather_qmm` on Metal; see the 2026-07-03 addendum), olmoe (OLMoE; migrated from
its local `SwitchGLU`/`SwitchLinear` to the shared ones; SwiGLU, softmax-routed
with optional `norm_topk_prob`, weights stacked under `switch_mlp.{gate,up,down}_proj`
or joined from the per-expert `experts.{i}` layout by `sanitize_weights`; expert
Expand All @@ -301,3 +304,47 @@ to `gather_qmm`, generating coherent output with no crash or OOM; fused-path
throughput and output-parity validation remain blocked on a fitting 4-bit or 8-bit
minimax checkpoint). nemotron-h's MoE runs through the separate C++ `fused_moe_forward` and is wired
behind `MLXCEL_FUSED_MOE_RELU2` only.

### 2026-07-03 addendum: GB10 crossover re-measured under MLX pin e9463bb (#626)

The #625 MLX pin bump to `e9463bb` moved `gather_gemm` to a JIT path, changing
the `gather_qmm` fallback cost, so the CUDA crossover was re-measured on GB10 (DGX
Spark, sm_121, CUDA 13.0). Harness: `mlxcel-bench-decode`, prompt "Hello, how are
you today?", 100 decode tokens after a 20-token warmup, median of 3. Each side is
forced with the env var: `MLXCEL_FUSED_MOE_MAX_DFF=1` selects the `gather_qmm`
fallback, `=20000` selects the fused kernel. decode tok/s:

| 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) | 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 |

Findings:

- **The crossover moved down from ~13-14k to ~8000, but stays well above Metal's
4096.** The fused path clearly wins at Dff 6400 (phi-3.5-moe +5%) and below
(lfm2 +12%), is break-even at 8192 (llama-4-scout, within run-to-run noise: one
of three fused runs matched the fallback), and loses at 14336 (mixtral -1.2%).
Interpolating phi (+5%) and llama-4-scout (break-even) puts the ratio=1.0
crossover at ~8000.
- **The old +55% qwen3-30b-a3b win is gone**, not because the fused kernel
regressed but because the JIT `gather_gemm` made the fallback far faster for the
pathological 128-expert config (old fallback 58.2 -> now 91.0 tok/s). The fused
path itself is unchanged there (~91). Low-expert configs are unaffected: lfm2 is
still +12%, matching the #319 measurement.
- **gpt-oss-20b is a control**: it is `mxfp4`, which `forward_fused_kernel`
declines (non-affine), so the env var is a no-op and both sides match (0.99).
qwen3.5-35b-a3b (Dff 512, multimodal `Qwen3_5MoeForConditionalGeneration`) does
not run in the text decode bench harness and is below 4096 regardless.

**Default set to 8192 on CUDA** (the break-even boundary, rounded down from the
14336 regression), keeping Metal at 4096. Models with an expert intermediate in
`(4096, 8192]` now take the fused kernel by default on CUDA: phi-3.5-moe (Dff
6400, 51.09 -> 53.62 tok/s, +5%) and llama-4-scout-17b (Dff 8192, 21.23 -> 21.08
tok/s, break-even). Mixtral (Dff 14336) stays on `gather_qmm`. The env var
`MLXCEL_FUSED_MOE_MAX_DFF` still overrides both backends. Raw per-run data:
`fused-moe-decode-dff-sweep-gb10-2026-07-03.csv`.
2 changes: 1 addition & 1 deletion docs/environment-variables.md
Original file line number Diff line number Diff line change
Expand Up @@ -208,7 +208,7 @@ recommended as normal deployment settings.
| `MLXCEL_PIPELINE_GRANULARITY` | `off`, `layer`, `block:N` | `off` | Inserts layer-boundary async-eval hints for pipeline experiments. |
| `MLXCEL_FUSED_MOE` | `0`/`false`/`off`/`no` disable; any other value or unset enables | on | Fused single-token decode-MoE kernel (#268), on by default since #282 (Metal) and #319 (CUDA, via `mx.fast.cuda_kernel`); validated on M1 Ultra, M5, and GB10. Set to `0` to force the proven `gather_qmm`/`SwitchGLU` path. Active for qwen3_moe, qwen3_next, dots.llm1, gemma4, qwen2_moe, mixtral, phimoe, lfm2, qwen3_vl_moe, and olmoe decode. |
| `MLXCEL_FUSED_MOE_SGY` | `1`-`32` | `8` | Simdgroups (Metal) / warps-per-block (CUDA) per threadgroup for the fused decode-MoE kernel; tune per hardware. |
| `MLXCEL_FUSED_MOE_MAX_DFF` | positive int | `4096` | Expert-intermediate (Dff) upper bound for the fused path; above it the caller falls back to `gather_qmm`. The fused path wins only while `gather_qmm` underutilizes the GPU (small experts). The break-even is hardware-dependent: ~4096 on M1 Ultra, ~13-14k on GB10/CUDA. 4096 is the conservative shared default; raise it on CUDA for mid-size experts. |
| `MLXCEL_FUSED_MOE_MAX_DFF` | positive int | `4096` (Metal) / `8192` (CUDA) | Expert-intermediate (Dff) upper bound for the fused path; above it the caller falls back to `gather_qmm`. The fused path wins only while `gather_qmm` underutilizes the GPU (small experts), so the break-even is backend-dependent and the default is chosen from the live backend: `4096` on Metal (M1 Ultra tuning) and `8192` on CUDA (GB10 re-measured under MLX pin e9463bb, #626; fused wins through Dff 6400 and is break-even at 8192). An explicit value overrides the default on both backends: lower it to force `gather_qmm` sooner, raise it (e.g. `20000`) to force the fused kernel on larger experts such as mixtral (Dff 14336, where it is a slight net loss). |
| `MLXCEL_FUSED_MOE_RELU2` | presence enables | off | Enables the squared-ReLU fused MoE path for nemotron-class experts; performance-neutral on nemotron-h, kept for a future MoE-dominated squared-ReLU model. |
| `MLXCEL_FUSED_QK_NORM` | `1`/`true`/`on`/`yes` enable; any other value or unset disables | off (opt-in) | Fused single-token QKV projection + Q/K RMSNorm + RoPE kernel (#326) for Qwen3 and Qwen3-MoE decode. Opt-in: set to `1` to enable. Matches the graph path within RMS < 5e-3 (the reduction is over the transpose-invariant head_dim axis), but greedy temp-0 is not byte-identical over long generation; on CUDA the graph path is itself non-deterministic run-to-run from GPU FP-reduction order, while the fused path is deterministic, so its output stays inside the graph baseline's own envelope. The kernel cuts Rust/C++ FFI crossings rather than MLX op count, so it does not speed up the GPU/bandwidth-bound decode loop: on M1 Ultra it measured 1 to 3.4% slower (qwen3-0.6b 275 vs 284, qwen3-8b 82.3 vs 83.2 tok/s); on GB10/CUDA (SM 12.1) it is also slower (qwen3-0.6b 0.96x, qwen3-8b ~1.0x, qwen3-30b-a3b 0.92x fused/graph; see `docs/benchmark_results/fused-qk-norm-decode-gb10.md`), so there is no per-backend win. Ships as a reusable shared primitive for the deferred QK-norm families and stays opt-in (default off) on every measured backend (M1 Ultra, M5 Max, GB10/CUDA), mirroring the opt-in `MLXCEL_FUSED_MOE_RELU2`. Active only when `l == 1` (decode) and weights are quantized. |
| `MLXCEL_FUSED_XIELU` | `0`/`false`/`off`/`no` disable; any other value or unset enables | on | Fused single-launch Metal xIELU kernel for the Apertus MLP activation (#409), on by default since the M5 Max validation. `MLP::forward` routes through one Metal dispatch covering the ~11 elementwise ops in `apertus_xielu` (square, minimum, expm1, where, and neighbors) instead of the per-op graph. Greedy temp-0 decode is byte-identical to the elementwise path on Apple Silicon: every intermediate stays in the input dtype (bf16) and the kernel reproduces MLX's `expm1f` exactly. Measured decode speedup on M1 Ultra (+2.7%, Apertus-8B 83.4 to 85.7 tok/s) and M5 Max (+1.9%, 112.0 to 114.2 tok/s), with no regression. Set to `0` to force the elementwise path. On non-Metal back-ends the FFI falls back to an equivalent elementwise graph, so the flag is safe to set everywhere. Apertus only; no other model family is affected. |
Expand Down
5 changes: 5 additions & 0 deletions src/lib/mlxcel-core/cpp/mlx_cxx_bridge.h
Original file line number Diff line number Diff line change
Expand Up @@ -1558,6 +1558,11 @@ void fused_gated_delta_decode_step(
// Check if GatedDeltaNet Metal kernel is available (Metal GPU only)
bool gated_delta_kernel_available();

// True when the MLX Metal backend is available at runtime (macOS Apple
// Silicon); false on CUDA-only and CPU-only builds. Mirrors the
// metal::is_available() gate used to pick the Metal vs CUDA kernel port. #626.
bool metal_is_available();

// Start/stop Metal GPU capture. Requires the process to run under
// `MTL_CAPTURE_ENABLED=1`; otherwise Metal drops the capture silently.
void metal_start_capture(rust::Str path);
Expand Down
13 changes: 13 additions & 0 deletions src/lib/mlxcel-core/cpp/mlx_cxx_kernels.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -769,6 +769,19 @@ bool gated_delta_kernel_available() {
#endif
}

// True when the MLX Metal backend is available at runtime. Mirrors the
// metal::is_available() gate used to pick the CUDA vs Metal kernel port
// (run_fused_moe_two_kernel, bitlinear_matmul) so Rust callers can choose a
// backend-specific default without a compile-time cfg. False on CUDA-only and
// CPU-only builds. See issue #626.
bool metal_is_available() {
#ifdef __APPLE__
return mlx::core::metal::is_available();
#else
return false;
#endif
}

// Start a GPU trace capture and stop an active one. Mirrors the Python
// `mx.metal.start_capture` / `mx.metal.stop_capture` API so a mlxcel
// decode run can emit the same `.gputrace` files that `mlx-lm` produces
Expand Down
7 changes: 7 additions & 0 deletions src/lib/mlxcel-core/src/lib.rs
Original file line number Diff line number Diff line change
Expand Up @@ -1801,6 +1801,13 @@ mod ffi {
/// Check whether the current default device is GPU
fn is_gpu_available() -> bool;

/// True when the MLX Metal backend is available at runtime (macOS
/// Apple Silicon). False on CUDA-only and CPU-only builds. Mirrors the
/// `mlx::core::metal::is_available()` gate that picks the Metal vs CUDA
/// kernel port, so Rust callers can choose a backend-specific default
/// without a compile-time cfg. See issue #626.
fn metal_is_available() -> bool;

/// Fused sampling: temperature + top-k + top-p + min-p + categorical
/// in a single C++ call to minimize FFI round-trips.
/// Input: 2D logits [batch, vocab] (already sliced, penalties applied)
Expand Down
Loading