Skip to content

perf(moe): validate fused decode-MoE on M5, then enable MLXCEL_FUSED_MOE by default #282

Description

@inureyes

Problem / Background

The fused single-token decode-MoE kernel effort (#268) is implemented and merged across PRs #276 (two-kernel SwiGLU), #278 (6-bit/mixed-bit, dots.llm1), #279 (qwen3_next / qwen3.5/3.6), #280 (squared-ReLU kernel preserved behind MLXCEL_FUSED_MOE_RELU2), and #281 (GeGLU, gemma4). The feature is currently off by default, opt-in via the MLXCEL_FUSED_MOE env flag, scoped to single-token decode only, with 4/8-bit gate-up + 4/6/8-bit down affine quantization.

Measured on M1 Ultra (the only hardware tested so far), decode throughput with the flag on, all with no regression:

Model Activation / quant Gain Numerical status
gemma-4-26b-a4b-it GeGLU 4-bit +13% (73.8 -> 83.2 tok/s) greedy byte-identical (with chat template)
qwen3.5/3.6-35b-a3b SwiGLU 4-bit +8.7% (68.7 -> 74.7 tok/s) within f16 jitter class (greedy may flip near-tie tokens; abs RMS ~1e-4 < 5e-3)
dots.llm1 SwiGLU mixed 4/6-bit +4.7% (13.1 -> 13.7 tok/s) byte-identical
qwen3-30b-a3b SwiGLU 4-bit +3.5% (47.3 -> 49.0 tok/s) byte-identical
qwen1.5-moe-a2.7b + other SwiGLU 4-bit SwiGLU 4-bit ~parity, no regression byte-identical
nemotron-h-30b squared-ReLU ~0% reverted from default path, kept behind MLXCEL_FUSED_MOE_RELU2 (MoE is a small slice of its Mamba-hybrid decode)

The f16 jitter on deep hybrids (qwen3.5/3.6) is the expected numerical consequence of kernel fusion, not a bug. The design harness accepts "byte-identical OR within f16 jitter class" (abs RMS ~1e-4 < 5e-3 threshold).

M5-specific risk (core reason for this issue)

The C++ fused_moe_forward (src/lib/mlxcel-core/cpp/mlx_cxx_kernels.cpp) carries a documented workaround: a mixed float32 x float16 multiply produces NaN on M5 Max (Metal GPU Family 4) via the NAx broadcast kernel. The new fused kernels (moe_gateup / moe_down / moe_fc1_relu2, and the GeGLU act path) cast scores to the activation dtype and accumulate in f32 / output f16, so they are expected to be safe, but they have never been run on M5. This must be validated before the flag can be flipped on by default.

Proposed Solution

Run the representative MoE set on M5 (Neural Accelerator) hardware with MLXCEL_FUSED_MOE=1, confirm correctness and gains, then flip the flag to default-on for supported configs (gating on hardware detection if M5 needs special handling).

Acceptance Criteria

  • Run the representative MoE set with MLXCEL_FUSED_MOE=1 on an M5 (Neural Accelerator) machine: qwen3-30b-a3b-4bit, dots.llm1 (mixed 4/6-bit), qwen3.5-35b-a3b-4bit, gemma-4-26b-a4b-it-4bit, qwen1.5-moe-a2.7b-4bit. (gemma4 must be checked with the chat template — raw prompts degenerate.)
  • Confirm no NaN / no garbage: greedy temp-0 output is coherent and within the f16 jitter class of the fused-off path (byte-identical where the M1 results were byte-identical).
  • Benchmark decode tok/s vs fused-off on M5; confirm no regression and record the M5 gains.
  • If M5 is clean: flip MLXCEL_FUSED_MOE to default-on for single-token decode + supported configs. Keep an env override to force it off (MLXCEL_FUSED_MOE=0). If M5 needs special handling, gate the default on hardware detection (hardware::is_m5_neural_accelerator()), mirroring the existing M5 gates in the codebase.
  • Update the feature docs (docs/benchmark_results/fused-moe-decode-kernel-design.md, the flags section) to reflect default-on.

Technical Considerations

Activity

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Metadata

Metadata

Assignees

Labels

area:inferenceGeneration, sampling, decoding (incl. speculative, DRY)area:modelsModel architectures, weights, loading, metadataplatform:macosmacOS (Apple Silicon) specificpriority:mediumMedium prioritystatus:doneCompletedtype:performancePerformance improvements

Type

No type

Projects

No projects

    Milestone

    Relationships

    None yet

    Development

    No branches or pull requests

    Issue actions