Repository navigation
perf(moe): wire mixtral to the fused decode-MoE kernel - #311
Conversation
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.
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.
Orchestrator validation + Dff guard (M1 Ultra)Real-checkpoint validation on To resolve this without regressing the small-expert winners, this PR adds a Dff upper-bound guard to
4096 keeps every measured winner (including qwen3_next at 2560) and excludes both losers (6400, 14336). Validation with the guard:
Net for mixtral: the local |
## 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
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/SwitchLinearto the sharedcrate::models::switch_layers::SwitchGLU, andSparseMoeBlock::forwardgains the single-token decode fused-kernel dispatch.What changed
src/models/switch_layers.rs: addSwitchGLU::from_weights_with_proj_names(weights, prefix, group_size, bits, [gate, up, down]). The defaultfrom_weightsnow 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 onstack_individual_expertsupdated to note the override path; new unit testswitch_glu_loads_overridden_proj_leaf_namesexercises thew1/w2/w3path.src/models/mixtral.rs: remove the localSwitchGLU/SwitchLinear(nothing else referenced them), pointSparseMoeBlock.expertsat the sharedSwitchGLU, and load viafrom_weights_with_proj_names(.., "{prefix}.switch_mlp", .., ["w1", "w3", "w2"])(gate=w1, up=w3, down=w2).SparseMoeBlock::forwardnow triesforward_fused_kernelfor single-token decode whenfused_moe_enabled(), falling back toexperts.forward+moe_weighted_sumotherwise. mixtral uses softmax over the top-k logits with no shared expert and no renorm, soscoresalready 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 thew1/w2/w3mapping 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 withutils::stack_arrays, withgate_proj=w1,up_proj=w3,down_proj=w2. The shared per-expert stacker reads the same keys ({root}.experts.{idx}.{leaf}where the.switch_mlpvirtual prefix resolvesroottoblock_sparse_moe, identical to the qwen2_moe path), in the same ascending index order, and stacks along axis 0 withstack_owned. Both stackers reduce to the sameffi::stack(&ptrs, 0)call, so the resulting[num_experts, out, in_packed]tensors are identical. Bits: the local loader hard-setargs.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) = 4for w1/w3 and(1792*32)/(224*64) = 4for w2, matchingargs.bits(). mode isaffine(default), biases are present, soquantized_parts()isSomeand the fused kernel qualifies.Test plan
cargo check --lib --tests --features metal,acceleratecargo clippy --lib --tests --features metal,accelerate -- -D warnings(clean; dead-code from the removed local structs would have failed-D warnings)cargo fmt -- --checkmodels/mixtral-8x7b-4bit: greedy temp-0 output coherent and byte-identical or within the documented f16 jitter class vsMLXCEL_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