Skip to content

One-kernel MoE combine for Qwen3.5/3.6 MoE prefill on the invariant lane (opt-in, depends on #549) - #558

Open
jvmenen wants to merge 4 commits into
youssofal:mainfrom
jvmenen:pr/a3b-moe-prefill-combine
Open

jvmenen wants to merge 4 commits into
youssofal:mainfrom
jvmenen:pr/a3b-moe-prefill-combine

Conversation

@jvmenen

@jvmenen jvmenen commented Sep 28, 2026

Copy link
Copy Markdown

Depends on #549. This branch is #549 plus one commit; only the last commit (perf(prefill): one-kernel MoE combine for Qwen3.5/3.6 MoE on the invariant lane (opt-in)) is new here. Once #549 is merged, the diff reduces to that commit.

Summary

Opt-in MTPLX_A3B_MOE_PREFILL_COMBINE=1: on the batch-invariant prefill lane from #549, the tail of mlx-lm's Qwen3NextSparseMoeBlock (Qwen3.5/3.6 MoE) runs as one kernel. The routed experts' sorted output, its inverse permutation, the routing scores and the gated shared expert go straight into the existing Flash-Next combine kernel (mtplx/kernels/qwen4_moe_prefill_combine.py) instead of materializing three [rows, top_k, hidden] tensors. Bit-identical to the stock block. Default off, and only installed together with the lane.

Motivation

After the routed experts, the stock block unsorts the expert outputs to token order, multiplies by the scores and sums over top_k, each step writing a [rows, top_k, hidden] BF16 tensor, and then adds the shared expert. On the invariant lane every prefill forward from the lane's sorted minimum (128 tokens at 256 experts, top-8) already runs the experts expert-sorted, so the sorted output and its inverse permutation are available anyway. Flash-Next's prefill combine kernel does unsort, weighting, sum and shared add in one pass with the same rounding as the stock tail, and is generic in top_k and hidden.

Change

  • mtplx/batch_invariant_prefill.py: BatchInvariantSwitchGLU gains sorted_route_applies(x, top_k) (prefill phase and at least the sorted minimum, i.e. no padding) and sorted_experts(x, indices): the stock sorted chain (_gather_sort, up/gate/down with sorted_indices=True) without the final unsort, returning the expert-sorted outputs and the inverse permutation. __call__ is unchanged apart from sharing the minimum computation.
  • New mtplx/a3b_moe_prefill_combine.py: a subclass of the stock Qwen3NextSparseMoeBlock. When the lane's SwitchGLU reports sorted_route_applies, it computes the routing exactly as the stock block does, runs the stock shared expert and gate, sorted_experts, and the combine kernel; otherwise it calls the stock block. Installed by class swap; the parameter tree is unchanged.
  • mtplx/runtime.py _install_batch_invariant_lane: installs the combine only when the switch is on and the lane was admitted; the lane's install report gains moe_prefill_combine_blocks.
  • mtplx/kernel_selfcheck.py: new bitwise self-check lane a3b_moe_prefill_combine (128 rows, top-8 of hidden 2048, the lane's narrowest sorted forward at that geometry) against the stock unsort-weight-sum-add tail. Any difference or failure turns the route off for the process.
  • Narrower forwards, decode, verify and the lone final prompt token keep the stock block.

How to enable: MTPLX_BATCH_INVARIANT_PREFILL=1 MTPLX_A3B_MOE_PREFILL_COMBINE=1. Without the lane the switch does nothing.

Tests

  • New tests/test_a3b_moe_prefill_combine.py (Metal, 15 tests): the block is bit-identical to the stock block at 256 experts, top-8, at 128, 331, 2,048 and 4,096 rows; whole-prefill logits, hidden and cache bit-identical with and without the combine on a tiny Qwen3.5-MoE (tests/a3b_tiny_synth.py, the same helper as in the draft-head history PR); forwards below the sorted minimum, decode and the final token take the stock block; install only on the lane and only with the switch; off by default; the self-check validates the kernel at this geometry, and a failed self-check turns the route off.
  • tests/test_batch_invariant_prefill.py passes.
  • Full suite with MTPLX_CONFIG=/nonexistent: 9,494 passed, 67 skipped, 1 failed. The failure is test_laguna_model.py::test_laguna_s_2_1_ar_route_skips_qwen_performance_hooks, which fails the same way on main on a Mac with less than 85.3 GiB (addressed in Tests no longer read the user's ~/.mtplx/config.toml or depend on the Mac's memory size #535). Ruff: new files clean, no new findings in the changed files.

🤖 Generated with Claude Code

Jeroen van Menen and others added 4 commits September 27, 2026 17:32
MTPLX_BATCH_INVARIANT_PREFILL=1 (default off) makes every prefill row
independent of how many rows share its forward. On MLX 0.32 three kernel
choices follow the row count and round differently: split-K quantized
matmul for narrow outputs (router N=256, shared expert, GDN a/b, k/v),
gather_qmv against gather_qmm_rhs for the routed experts, and the unfused
against the fused SDPA below 1024 query rows. On Qwen3.6-35B-A3B the 8th
and 9th expert often tie within a bf16 step, so those differences change
the routing and move scores between block layouts.

In the prefill phase only: QuantizedLinear runs as a two-batch matmul of
at least 33 rows per batch (no split-K, one NAX tile shape), SwitchGLU
pads to 4 token rows per expert (always the sorted tiled kernel), and
mlx-lm's scaled_dot_product_attention (hooked in qwen3_next and in base,
where attention_split imports it per call) always runs the fused causal
kernel, with zero query rows in front below 9 rows. The lone final prompt
token runs on the stock kernels (stock_prefill_kernels): padded to the
lane's minimums it cost TTFT, and no caller compares that row with a wider
forward. Class swaps and one function hook; decode and verify keep the
stock kernels.

The lane is installed only where it covers every quantized projection
(affine QuantizedLinear, stock SwitchGLU with at least 16 experts) and
only on MoE models; a dense model is refused with "dense_model" unless
MTPLX_BATCH_INVARIANT_PREFILL_DENSE=1. On a refusal nothing is swapped or
hooked. /health reports the install report or the refusal reason under
degradation.batch_invariant_prefill.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
Prompt scoring (/v1/completions with echo and max_tokens 0) ran its trunk
in forwards of 256 rows, the width of its logits slices. With the
batch-invariant prefill lane installed the result no longer depends on
the width, so the trunk now runs at the prefill chunk a generation would
use (the server's live setting, else the profile) and the target lm_head
is applied per 256-row slice of the post-norm rows, inside the prefill
phase: logits residency stays at 256 rows. mtp_patch exposes the head as
logits_from_post_norm, the same call __call__ uses.

Without the lane the default stays 256 rows, because on MoE models a
wider trunk changes the routing and the scores.
MTPLX_PROMPT_SCORE_TRUNK_CHUNK names a width explicitly.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
…the invariant lane

With the batch-invariant prefill lane installed, a cold prompt shorter
than the block-restore floor (MTPLX_SESSION_BLOCK_PREFIX_MIN_MATCH_TOKENS,
512) is prefilled without the tail boundary grid. On Qwen3.6-35B-A3B the
grid cuts one extra forward off the last chunk (~0.1 s, all experts read
again) whenever a session bank is attached, e.g. [0,331],[331,395] for a
396-token prompt. With the lane both layouts give bitwise the same states,
so only the cost changes; without it the cut stays because removing it
would change the routing and the scores.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
…riant lane (opt-in)

After the routed experts, mlx-lm's Qwen3NextSparseMoeBlock (Qwen3.5/3.6
MoE) materializes three [rows, top_k, hidden] BF16 tensors (unsort, weight,
column sum) before adding the shared expert. On the batch-invariant prefill
lane every forward from 128 tokens already runs the experts expert-sorted, so
the lane's SwitchGLU now hands out that sorted output and its inverse
permutation (sorted_experts: the stock chain without the unsort), and a
subclass of the stock block feeds both, with the scores and the gated shared
expert, to Flash-Next's combine kernel (kernels/qwen4_moe_prefill_combine.py,
generic in top_k and hidden, same rounding as the stock tail). Routing,
experts and shared expert are the stock calls.

MTPLX_A3B_MOE_PREFILL_COMBINE=1 (default off), installed only together with
the invariant lane. Narrower forwards (below the lane's sorted minimum),
decode, verify and the lone final token keep the stock block. New self-check
lane a3b_moe_prefill_combine (bitwise, 128 rows, top-8 of 2048) turns the
route off for the process on any difference or failure.

Tests (Metal): tests/test_a3b_moe_prefill_combine.py, 15 tests: bit-identical
to the stock block at 256 experts, top-8, from 128 to 4,096 rows; whole-prefill
logits, hidden and cache bit-identical on a tiny Qwen3.5-MoE
(tests/a3b_tiny_synth.py); install and self-check behaviour.

Co-Authored-By: Claude Opus 5.5 (1M context) <noreply@anthropic.com>
@jvmenen
jvmenen requested a review from youssofal as a code owner September 28, 2026 09:51

This branch has not been deployed

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

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

1 participant