Skip to content

perf(moe): implement the fused decode-MoE kernel on CUDA - #319

Merged
inureyes merged 2 commits into
mainfrom
fix/313-cuda-fused-moe-default-off
Jun 17, 2026
Merged

inureyes merged 2 commits into
mainfrom
fix/313-cuda-fused-moe-default-off

Conversation

@inureyes

@inureyes inureyes commented Jun 17, 2026 •

Copy link
Copy Markdown
Member

Summary

The fused single-token decode-MoE kernel (#268) was Metal-only (mx.fast.metal_kernel), so on the CUDA backend it threw [metal_kernel] No Metal back-end and aborted for every small-expert MoE the dispatch selected (qwen2_moe, qwen3_moe, lfm2, qwen3_vl_moe). Larger experts fell through the dff > MLXCEL_FUSED_MOE_MAX_DFF break-even to gather_qmm and survived, so the failure tracked expert size, not architecture.

This PR ports the kernel to CUDA instead of disabling it, so the affected models run faster on CUDA rather than merely working.

What changed

Port both halves to mx.fast.cuda_kernel (the CUDA analogue of metal_kernel):

  • Kernel A: gate/up GEMV + SwiGLU/GeGLU → act_g[K, Dff].
  • Kernel B: down GEMV × score → summed over K.
  • One warp owns each output row; simd_sum becomes a __shfl_down_sync warp reduction; the quantized-affine dequant (4/8-bit gate/up, 4/6/8-bit down) is unchanged.
  • run_fused_moe_two_kernel selects the cuda_kernel port when metal::is_available() is false. no_cuda.cpp / no_metal.cpp stub the unused side, so both link on either backend.
  • Precise expf in the SwiGLU holds greedy parity.

Only src/lib/mlxcel-core/cpp/mlx_cxx_kernels.cpp changes; fused_moe_enabled stays default-on for both backends.

Validation (GB10 / DGX Spark, CUDA 13.0)

Greedy byte-identical to the gather_qmm path on qwen3-30b-a3b.

Model gather_qmm fused (CUDA) speedup
qwen3-moe-4bit 58.15 90.34 1.55x
qwen3.5-35b-a3b-4bit 45.84 64.69 1.41x
lfm2-8b-a1b-4bit 140.91 158.22 1.13x
qwen1.5-moe-a2.7b-4bit 113.03 124.73 1.10x

The previously-crashing MoE models now run at these fused rates by default (no env var).

Supersedes the earlier CUDA fused-off workaround. Relevant to the #307 fused-MoE wiring epic: the wired families now work on CUDA.

Closes #313

@inureyes inureyes added area:inference Generation, sampling, decoding (incl. speculative, DRY) priority:high High priority type:bug Bug fixes, error corrections, or issue resolutions labels Jun 17, 2026
The fused single-token decode-MoE kernel (#268) was Metal-only
(mx.fast.metal_kernel), so on the CUDA backend it threw
"[metal_kernel] No Metal back-end" and aborted for every small-expert
MoE that the dispatch selected (qwen2_moe, qwen3_moe, lfm2,
qwen3_vl_moe); larger experts fell through the dff break-even to
gather_qmm and survived.

Port both halves of the kernel (gate/up + SwiGLU/GeGLU, then down +
score) to mx.fast.cuda_kernel: one warp owns each output row, the
simd_sum reduction becomes a __shfl_down_sync warp reduction, and the
quantized-affine dequant is unchanged. run_fused_moe_two_kernel selects
the cuda_kernel port when metal::is_available() is false (no_cuda.cpp /
no_metal.cpp stub the unused side, so both link on either backend).
Precise expf in the SwiGLU holds greedy parity.

Greedy byte-identical to the gather_qmm path on qwen3-30b-a3b. Measured
on GB10 (DGX Spark, CUDA 13.0), decode tok/s vs gather_qmm:
  qwen3-moe-4bit        58.15 -> 90.34  (1.55x)
  qwen3.5-35b-a3b-4bit  45.84 -> 64.69  (1.41x)
  lfm2-8b-a1b-4bit     140.91 -> 158.22 (1.13x)
  qwen1.5-moe-a2.7b-4bit 113.03 -> 124.73 (1.10x)

This supersedes the earlier CUDA fused-off workaround: fused stays
default-on for both backends, and the previously-crashing MoE models now
run faster instead of falling back.

Closes #313
@inureyes
inureyes force-pushed the fix/313-cuda-fused-moe-default-off branch from e857236 to 7880e4d Compare June 17, 2026 00:56
@inureyes inureyes changed the title fix(moe): default the fused decode-MoE kernel off on CUDA perf(moe): implement the fused decode-MoE kernel on CUDA Jun 17, 2026
@inureyes

Copy link
Copy Markdown
Member Author

Force-pushed: this PR now implements the CUDA fused decode-MoE kernel instead of the earlier default-off workaround. Porting the Metal kernel to mx.fast.cuda_kernel is greedy byte-identical and faster on GB10 (+10% to +55%), so fused stays default-on for both backends and the previously-crashing MoE models run faster rather than falling back to gather_qmm. See the updated description.

Record the CUDA port in the kernel design doc: a GB10 measured-gains
section (qwen3-30b-a3b +55%, qwen3.5-35b-a3b +41%, lfm2 +13%,
qwen1.5-moe +10%), the break-even sweep (CUDA crossover ~13-14k vs
Metal's ~4096, so the 4096 cap is conservative on CUDA but kept as the
shared default), a roadmap entry, and the default-on-both-backends note.

Update environment-variables.md: MLXCEL_FUSED_MOE now spans both
backends, and document MLXCEL_FUSED_MOE_MAX_DFF with the per-backend
break-even.
@inureyes

Copy link
Copy Markdown
Member Author

Kept MLXCEL_FUSED_MOE_MAX_DFF at 4096 (shared default). The break-even probe shows the CUDA crossover is ~13-14k (Dff 768 +60%, 1792 +13%, 6400 +2%, 14336 -2%), so 4096 is conservative on CUDA, but the meaningful wins are all below it and the 4096-14k range gains at most ~2% (phi-3.5-moe) while mixtral slightly prefers gather_qmm. Documented the CUDA backend + per-backend break-even in the kernel design doc and environment-variables.md (cbcce6c).

@inureyes
inureyes merged commit 06a8d09 into main Jun 17, 2026
5 checks passed
@inureyes
inureyes deleted the fix/313-cuda-fused-moe-default-off branch June 17, 2026 01:22
inureyes added a commit that referenced this pull request Jun 17, 2026
First GB10 benchmark on the merged CUDA fused decode-MoE kernel (#319): 148 text models (133 pass / 14 fail / 1 OOM-skip), 53 VLM rows. Nine fused-MoE models flipped FAIL→pass; qwen3-moe-30b (89.84) edges past M1 Ultra (83.75). Regenerates model_tests_gb10.md + index for 0.3.1.
@inureyes inureyes added this to the 0.3 milestone Jun 21, 2026
@inureyes inureyes self-assigned this Aug 31, 2026
@inureyes inureyes added the status:done Completed label Sep 6, 2026
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) priority:high High priority status:done Completed type:bug Bug fixes, error corrections, or issue resolutions

Projects

None yet

Development

Successfully merging this pull request may close these issues.

GB10 CUDA: Qwen fused-MoE models fail warmup at 0.3.0 (regression vs 0.1.0)

1 participant