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
Open
One-kernel MoE combine for Qwen3.5/3.6 MoE prefill on the invariant lane (opt-in, depends on #549)#558jvmenen wants to merge 4 commits into
jvmenen wants to merge 4 commits into
Conversation
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>
This branch has not been deployed
This file contains hidden or bidirectional Unicode text that may be interpreted or compiled differently than what appears below. To review, open the file in an editor that reveals hidden Unicode characters.
Learn more about bidirectional Unicode characters
Sign up for free
to join this conversation on GitHub.
Already have an account?
Sign in to comment
Add this suggestion to a batch that can be applied as a single commit.This suggestion is invalid because no changes were made to the code.Suggestions cannot be applied while the pull request is closed.Suggestions cannot be applied while viewing a subset of changes.Only one suggestion per line can be applied in a batch.Add this suggestion to a batch that can be applied as a single commit.Applying suggestions on deleted lines is not supported.You must change the existing code in this line in order to create a valid suggestion.Outdated suggestions cannot be applied.This suggestion has been applied or marked resolved.Suggestions cannot be applied from pending reviews.Suggestions cannot be applied on multi-line comments.Suggestions cannot be applied while the pull request is queued to merge.Suggestion cannot be applied right now. Please check back later.
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'sQwen3NextSparseMoeBlock(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 intop_kandhidden.Change
mtplx/batch_invariant_prefill.py:BatchInvariantSwitchGLUgainssorted_route_applies(x, top_k)(prefill phase and at least the sorted minimum, i.e. no padding) andsorted_experts(x, indices): the stock sorted chain (_gather_sort, up/gate/down withsorted_indices=True) without the final unsort, returning the expert-sorted outputs and the inverse permutation.__call__is unchanged apart from sharing the minimum computation.mtplx/a3b_moe_prefill_combine.py: a subclass of the stockQwen3NextSparseMoeBlock. When the lane'sSwitchGLUreportssorted_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 gainsmoe_prefill_combine_blocks.mtplx/kernel_selfcheck.py: new bitwise self-check lanea3b_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.How to enable:
MTPLX_BATCH_INVARIANT_PREFILL=1 MTPLX_A3B_MOE_PREFILL_COMBINE=1. Without the lane the switch does nothing.Tests
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.pypasses.MTPLX_CONFIG=/nonexistent: 9,494 passed, 67 skipped, 1 failed. The failure istest_laguna_model.py::test_laguna_s_2_1_ar_route_skips_qwen_performance_hooks, which fails the same way onmainon 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