Skip to content

perf(moe): wire mixtral to the fused decode-MoE kernel - #311

Merged
inureyes merged 2 commits into
mainfrom
perf/issue-301-mixtral-fused-kernel
Jun 16, 2026
Merged

inureyes merged 2 commits into
mainfrom
perf/issue-301-mixtral-fused-kernel

Conversation

@inureyes

Copy link
Copy Markdown
Member

Summary

Group B migration for the fused decode-MoE kernel (#268 series, default-on since #282), chain 4/6 under epic #307 and sibling to #286 (qwen2_moe). mixtral's MoE experts move from its local SwitchGLU/SwitchLinear to the shared crate::models::switch_layers::SwitchGLU, and SparseMoeBlock::forward gains the single-token decode fused-kernel dispatch.

What changed

  • src/models/switch_layers.rs: add SwitchGLU::from_weights_with_proj_names(weights, prefix, group_size, bits, [gate, up, down]). The default from_weights now delegates to it with the standard ["gate_proj", "up_proj", "down_proj"] names, so every existing shared caller (qwen2_moe, mistral4, dots1, kimi_linear, minimax, longcat) is byte-for-byte unchanged. The per-expert stacker already derives the expert leaf name from the projection segment of the prefix, so overriding the leaf names is all that is needed; no checkpoint-specific strings live in the shared loader. Doc comment on stack_individual_experts updated to note the override path; new unit test switch_glu_loads_overridden_proj_leaf_names exercises the w1/w2/w3 path.
  • src/models/mixtral.rs: remove the local SwitchGLU/SwitchLinear (nothing else referenced them), point SparseMoeBlock.experts at the shared SwitchGLU, and load via from_weights_with_proj_names(.., "{prefix}.switch_mlp", .., ["w1", "w3", "w2"]) (gate=w1, up=w3, down=w2). SparseMoeBlock::forward now tries forward_fused_kernel for single-token decode when fused_moe_enabled(), falling back to experts.forward + moe_weighted_sum otherwise. mixtral uses softmax over the top-k logits with no shared expert and no renorm, so scores already carries the full combine weights; everything else in the forward path is unchanged.
  • docs/benchmark_results/fused-moe-decode-kernel-design.md: move mixtral into the wired "Models covered" list and document the w1/w2/w3 mapping and the overridable-leaf-name loader.

Weight-loading parity

The old local loader read block_sparse_moe.experts.{idx}.{w1,w2,w3}.{weight,scales,biases} for idx ascending and stacked them along axis 0 with utils::stack_arrays, with gate_proj=w1, up_proj=w3, down_proj=w2. The shared per-expert stacker reads the same keys ({root}.experts.{idx}.{leaf} where the .switch_mlp virtual prefix resolves root to block_sparse_moe, identical to the qwen2_moe path), in the same ascending index order, and stacks along axis 0 with stack_owned. Both stackers reduce to the same ffi::stack(&ptrs, 0) call, so the resulting [num_experts, out, in_packed] tensors are identical. Bits: the local loader hard-set args.bits(); the shared loader infers per-tensor bits from (packed_in * 32) / (num_groups * group_size), which for Mixtral-8x7B-4bit (group_size 64) is (512*32)/(64*64) = 4 for w1/w3 and (1792*32)/(224*64) = 4 for w2, matching args.bits(). mode is affine (default), biases are present, so quantized_parts() is Some and the fused kernel qualifies.

Test plan

  • cargo check --lib --tests --features metal,accelerate
  • cargo clippy --lib --tests --features metal,accelerate -- -D warnings (clean; dead-code from the removed local structs would have failed -D warnings)
  • cargo fmt -- --check
  • Real-checkpoint validation pending the orchestrator on models/mixtral-8x7b-4bit: greedy temp-0 output coherent and byte-identical or within the documented f16 jitter class vs MLXCEL_FUSED_MOE=0, no NaN, no garbage; decode tok/s fused-on vs off recorded with no regression. A wrong w1/w2/w3 mapping or a bits/shape mismatch would surface as garbage decode here.

Closes #301

Migrate mixtral's MoE experts from its local SwitchGLU/SwitchLinear to the shared crate::models::switch_layers::SwitchGLU, then add the single-token decode fused-kernel dispatch (#268) mirroring qwen2_moe/dots1/qwen3_moe. With the kernel on by default (#282), mixtral single-token decode now runs forward_fused_kernel and falls back to gather_qmm + moe_weighted_sum when MLXCEL_FUSED_MOE=0 or when the config is unsupported. Mixtral routes softmax over the top-k logits with no shared expert and no renormalization, so the scores already carry the full combine weights.

Mixtral stores experts under the w1/w2/w3 convention at block_sparse_moe.experts.{idx}.{w1,w2,w3}, mapping gate=w1, up=w3, down=w2. The shared per-expert stacker derives the expert leaf name from the projection segment of the prefix, so a new SwitchGLU::from_weights_with_proj_names(weights, prefix, group_size, bits, [gate, up, down]) lets mixtral pass ["w1","w3","w2"] without baking those names into the shared loader. The default from_weights delegates with the standard gate_proj/up_proj/down_proj names, so existing shared callers are unchanged.

The shared loader stacks per-expert tensors along axis 0 (ffi::stack, identical to the removed local loader's stack_arrays) and infers per-tensor bits from the packed-weight/scales shapes, which yields 4 for Mixtral-8x7B-4bit's uniform-precision experts. The local SwitchGLU/SwitchLinear are removed; no other module referenced them.

Docs: move mixtral to the wired "Models covered" list in the fused-moe decode-kernel design doc.
@inureyes inureyes added type:refactor Code restructuring without changing functionality type:performance Performance improvements priority:medium Medium priority area:models Model architectures, weights, loading, metadata area:inference Generation, sampling, decoding (incl. speculative, DRY) platform:macos macOS (Apple Silicon) specific status:review Under review labels Jun 16, 2026
The two-kernel fused decode-MoE path wins only while gather_qmm underutilizes the GPU (small experts). For large experts gather_qmm already saturates the GPU, so the fused path's extra dispatch plus global-memory activation staging becomes a net loss. Measured on M1 Ultra: Dff 704..2560 gain +3.5% to +15.4%, Dff 6400 (phi-3.5-moe) loses 5.9%, Dff 14336 (mixtral) loses 21%.

forward_fused_kernel now declines (the caller falls back to gather_qmm) when the expert intermediate exceeds MLXCEL_FUSED_MOE_MAX_DFF (default 4096). This keeps every measured winner, including qwen3_next at Dff 2560, and excludes the large-expert regressions. The bound is env-tunable because the break-even is hardware-dependent. With the default, mixtral (Dff 14336) stays on gather_qmm: the migration removes duplication while decode is unchanged from the gather_qmm baseline.
@inureyes

Copy link
Copy Markdown
Member Author

Orchestrator validation + Dff guard (M1 Ultra)

Real-checkpoint validation on mlx-community/Mixtral-8x7B-Instruct-v0.1-4bit surfaced a regression: the fused decode kernel is ~21% slower than gather_qmm for mixtral (expert intermediate Dff=14336). The output is correct (coherent, w1/w2/w3 mapping verified), but the two-kernel fused path only wins 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.

To resolve this without regressing the small-expert winners, this PR adds a Dff upper-bound guard to forward_fused_kernel (MLXCEL_FUSED_MOE_MAX_DFF, default 4096; env-tunable since the break-even is hardware-dependent). The threshold was calibrated against measured break-even data:

Dff model fused vs gather_qmm
704 gemma4-26b-a4b +13%
768 qwen3-30b-a3b / qwen3-vl-30b-a3b +3.5% / +15.4%
1408 qwen2_moe / dots.llm1 +5.8%
1792 lfm2 +6.6%
2560 qwen3_next +8.7%
6400 phi-3.5-moe -5.9%
14336 mixtral -21%

4096 keeps every measured winner (including qwen3_next at 2560) and excludes both losers (6400, 14336).

Validation with the guard:

  • mixtral default (guard on): 51.8-52.5 tok/s == MLXCEL_FUSED_MOE=0 baseline 52.1 tok/s (falls back, no regression).
  • mixtral MLXCEL_FUSED_MOE_MAX_DFF=99999: 41.2 tok/s (fused path re-arms, the -21% path returns, proving the guard is the gate).
  • qwen2_moe (Dff 1408, below the bound): fused on 146.6 vs off 134.3 tok/s (still gains; guard does not disable small winners).

Net for mixtral: the local SwitchGLU/SwitchLinear migration to the shared struct removes duplication, and decode stays on the proven gather_qmm path.

@inureyes inureyes added status:done Completed and removed status:review Under review labels Jun 16, 2026
@inureyes
inureyes merged commit d846fd4 into main Jun 16, 2026
4 checks passed
@inureyes
inureyes deleted the perf/issue-301-mixtral-fused-kernel branch June 16, 2026 13:43
@inureyes inureyes added this to the 0.3 milestone Jun 21, 2026
@inureyes inureyes self-assigned this Aug 31, 2026
@inureyes inureyes removed the type:refactor Code restructuring without changing functionality label Sep 7, 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:inference Generation, sampling, decoding (incl. speculative, DRY) area:models Model architectures, weights, loading, metadata platform:macos macOS (Apple Silicon) specific 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(moe): wire mixtral to the fused decode-MoE kernel

1 participant