Repository navigation
fix(metal): recompute M after broadcast in sorted-rhs quantized gather - #12
Conversation
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>
|
Base retargeted from |
|
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. |
Summary
GatherQMM::eval_gpudispatches the sorted-RHS quantized route(
takes_sorted_rhs_route:M==1 && B>=16 && right_sorted && B/E>=4) bycalling
gather_qmm_rhs/gather_qmm_rhs_naxwithM = x.size() / Kcomputed on the pre-broadcast
x. Inside those functionsbroadcast_with_indicesgrows an index-unalignedxtoindices.size()rows,but
Mwas never recomputed. The dispatch grid ((M+bm-1)/bm) and the kernelrow bound (
set_bytes(M)) therefore used the stale, smallerM, so onlyx.size()/Kof theindices.size()output rows were written. The tail wasleft as uninitialized pool memory — deterministic garbage on reuse, wrong
from the first call, no NaN. The stale
Malso undercounts the Gemma4 expertroute's descriptor tile count (which reads
Mas the assignment count viaset_bytes(M, 3)).Reference for the correct pattern:
gather_mm_rhs(matmul.cpp) broadcastsfirst, then derives
int M = a.size() / Kfrom the broadcast array. MLX's ownvjp already uses the aligned predicate
rhs_indices.size()*M*K == x.size().Fix
Recompute
M = x.size() / Kfrom the post-broadcastxin both sorted-RHSquantized paths (
gather_qmm_rhsandgather_qmm_rhs_nax), placed before everydownstream use of
M(Gemma4 classifier + tile build, grid sizing, kernelMargument).
K == x.shape(-1)is unchanged by the broadcast, so this mirrorsgather_mm_rhsexactly. Minimal and local to the dispatch; no.metalsourcechange.
RED / GREEN receipts
New test
test_gather_qmm_sorted_nested_broadcastpins 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:
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):
Row 0 (the single row the stale
M==1grid writes) is correct; rows 1..8191hold the poison sentinel — exactly the unwritten tail.
GREEN (fix applied):
Existing
test_gather_qmm,test_gather_qmm_sorted,test_gather_qmm_matrix_path,test_gather_qmm_grad,test_gather_matmul_gradall 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 is8 xthat =3.9e-3. Thefaulted tail sits at
~9.9e5(sentinel), so the RED/GREEN separation is ~9orders 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:
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/mlxsubmodule + the enginePackage.resolvedto thisfix and restores the sorted hint on the engine's batched-prefill path.
Fixed commit for Lane B to pin:
07896d9bf65f0ae57ca3f5aebe3e6c1dcbdb1eda(base
9b0d1b4cbb9924b5075098d2fa71a25891c89e8f).