Skip to content

fix(metal): recompute M after broadcast in sorted-rhs quantized gather - #12

Merged
davidtai merged 1 commit into
qwen3.8-125b-a6b-mlx-v1from
fix/gather-qmm-rhs-stale-M
Aug 29, 2026
Merged

davidtai merged 1 commit into
qwen3.8-125b-a6b-mlx-v1from
fix/gather-qmm-rhs-stale-M

Conversation

@davidtai

Copy link
Copy Markdown

Summary

GatherQMM::eval_gpu dispatches the sorted-RHS quantized route
(takes_sorted_rhs_route: M==1 && B>=16 && right_sorted && B/E>=4) by
calling gather_qmm_rhs / gather_qmm_rhs_nax with M = x.size() / K
computed on the pre-broadcast x. Inside those functions
broadcast_with_indices grows an index-unaligned x to indices.size() rows,
but M was never recomputed. The dispatch grid ((M+bm-1)/bm) and the kernel
row bound (set_bytes(M)) therefore used the stale, smaller M, so only
x.size()/K of the indices.size() output rows were written. The tail was
left as uninitialized pool memory — deterministic garbage on reuse, wrong
from the first call, no NaN. The stale M also undercounts the Gemma4 expert
route's descriptor tile count (which reads M as the assignment count via
set_bytes(M, 3)).

Reference for the correct pattern: gather_mm_rhs (matmul.cpp) broadcasts
first, then derives int M = a.size() / K from the broadcast array. MLX's own
vjp already uses the aligned predicate rhs_indices.size()*M*K == x.size().

Fix

Recompute M = x.size() / K from the post-broadcast x in both sorted-RHS
quantized paths (gather_qmm_rhs and gather_qmm_rhs_nax), placed before every
downstream use of M (Gemma4 classifier + tile build, grid sizing, kernel M
argument). K == x.shape(-1) is unchanged by the broadcast, so this mirrors
gather_mm_rhs exactly. Minimal and local to the dispatch; no .metal source
change.

RED / GREEN receipts

New test test_gather_qmm_sorted_nested_broadcast pins the 125B-A6B batched-
prefill firing geometry (128 experts, weight [704, 2816], 8192 assignments,
affine/4-bit/group-64) and asserts the sorted-route result matches an
independent dense-dequant reference. It drives:

  • (a) index-aligned x — correct today; used to measure the tolerance.
  • (b) nested / non-aligned x — the fault; output pool poisoned with a
    sentinel so the uninitialized tail is unambiguous.

Built against this repo's own Metal test harness (python/tests/test_quantized.py)
on Apple M5 / Metal, Xcode 26.6.

RED (fix reverted, both recompute lines commented):

[aligned] max|y-ref| = 4.882812e-04
[nested ] max|y-ref| = 9.871361e+05
[nested ] rows with error > 8*tol : 8191 / 8192
[nested ] row_err[0:4]   = [0.00024, 987136.06, 987136.06, 987136.06]
[nested ] row_err[A-4:A] = [987136.06, 987136.06, 987136.06, 987136.06]
-> RED (tail rows corrupt)
AssertionError: 987136.0625 not less than 0.00390625

Row 0 (the single row the stale M==1 grid writes) is correct; rows 1..8191
hold the poison sentinel — exactly the unwritten tail.

GREEN (fix applied):

[aligned] max|y-ref| = 4.882812e-04
[nested ] max|y-ref| = 4.882812e-04
[nested ] rows with error > 8*tol : 0 / 8192
[nested ] row_err[A-4:A] = [0.00024, 0.00024, 0.00024, 0.00024]
-> GREEN (all rows within tol)

Existing test_gather_qmm, test_gather_qmm_sorted, test_gather_qmm_matrix_path,
test_gather_qmm_grad, test_gather_matmul_grad all still pass.

Tolerance justification (measured, not guessed)

The dense-dequant reference differs from the sorted-quantized kernel only by the
affine/4-bit/group-64 quantization-error class plus fp accumulation order. That
error is measured on the index-aligned case (correct today):
max|y - y_ref| = 4.88e-4. The test threshold is 8 x that = 3.9e-3. The
faulted tail sits at ~9.9e5 (sentinel), so the RED/GREEN separation is ~9
orders of magnitude — the multiplier is not load-bearing.

Measured prefill delta

Microbenchmark of the quantized gather at the pinned firing shape, sorted hint
forwarded (sorted_indices=True, now correct on non-aligned x) vs hint-off
(sorted_indices=False, the current workaround), 100 iters after 20 warmup,
Apple M5 GPU:

path us/call
hint_off (workaround) ~42,500
hint_on (fixed sorted route) ~3,800
speedup ~11x (10.99x / 11.37x across two runs)

Both paths now produce correct output (max|y-ref| 4.9e-4 / 7.3e-4). The 125B
track scores prefill^0.25, so forwarding the hint on the batched-prefill path is
the point of the fix.

Lane B (separate, not in this PR)

This PR does not bump any submodule pointer. The follow-up lane pins
mlx-swift's Source/Cmlx/mlx submodule + the engine Package.resolved to this
fix and restores the sorted hint on the engine's batched-prefill path.
Fixed commit for Lane B to pin: 07896d9bf65f0ae57ca3f5aebe3e6c1dcbdb1eda
(base 9b0d1b4cbb9924b5075098d2fa71a25891c89e8f).

gather_qmm_rhs and gather_qmm_rhs_nax received M = x.size()/K computed by
the caller (GatherQMM::eval_gpu) on the PRE-broadcast x. When the sorted-rhs
route fires (takes_sorted_rhs_route: M==1, B>=16, right_sorted, B/E>=4) on an
index-unaligned x, broadcast_with_indices grows x to indices.size() rows but
M was never recomputed, so the dispatch grid ((M+bm-1)/bm) and the kernel row
bound (set_bytes(M)) still used the stale, smaller M. Only x.size()/K of the
indices.size() output rows were written; the tail was left as uninitialized
pool memory (deterministic garbage on reuse, wrong from the first call, no
NaN). The stale M also undercounts the Gemma4 expert route's descriptor tile
count, which consumes M as the assignment count.

Fix: recompute M = x.size() / K from the POST-broadcast x in both sorted-rhs
quantized paths, mirroring gather_mm_rhs (matmul.cpp), which broadcasts first
then derives M from the broadcast array. K == x.shape(-1) is unchanged by the
broadcast. This precedes every downstream use of M (Gemma4 classifier + tile
build, grid sizing, kernel M argument).

Adds test_gather_qmm_sorted_nested_broadcast pinning the 125B-A6B batched-
prefill firing geometry (128 experts, [704,2816], 8192 assignments). It drives
the sorted route on an index-aligned x (to measure the dense-dequant
quantization-error tolerance ~4.9e-4) and on a nested/non-aligned x with the
output pool poisoned; before the fix the tail rows read the poison sentinel,
after the fix every row matches the dense reference.

Co-Authored-By: Claude Fable 5 <noreply@anthropic.com>
@davidtai
davidtai changed the base branch from main to codex/gemma4-autoresearch-v0.8.2 August 29, 2026 14:21
@davidtai

Copy link
Copy Markdown
Author

Base retargeted from main to codex/gemma4-autoresearch-v0.8.2. The base commit 9b0d1b4c (what the mlx-swift Source/Cmlx/mlx submodule pins, and what Lane B pins to) is the tip of codex/gemma4-autoresearch-v0.8.2, not an ancestor of main — main (734241b, MLX 0.32.2 upgrade #10) is the upstream line and has rewritten quantized.cpp, so a PR against main conflicts. Against the codex branch this is a clean single-commit fast-forward (parent == 9b0d1b4).

@davidtai
davidtai changed the base branch from codex/gemma4-autoresearch-v0.8.2 to qwen3.8-125b-a6b-mlx-v1 August 29, 2026 14:28
@davidtai

Copy link
Copy Markdown
Author

Per lane ruling (David, 2026-08-29): merge target is the dedicated branch qwen3.8-125b-a6b-mlx-v1, cut from 9b0d1b4 (the exact base the 125B engine's mlx-swift pins). This PR is fix-only — a single commit (07896d9, 'fix(metal): recompute M after broadcast in sorted-rhs quantized gather') directly on 9b0d1b4, no gemma4-autoresearch / qwen36 feature commits riding along. gemma4 cuts its own branch for the same fix rather than sharing this one. Landing held on both two-verdict reviews clearing.

@davidtai
davidtai merged commit d1080f9 into qwen3.8-125b-a6b-mlx-v1 Aug 29, 2026
4 checks passed
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