Skip to content

Rebase onto upstream main: take upstream 1-bit (#3161), re-apply the fork as 10 patches (gemma4-qat#282) - #10

Open
jeethu wants to merge 18 commits into
prism-main-06eb7483-fixesfrom
fork/282-sync-upstream-main
Open

jeethu wants to merge 18 commits into
prism-main-06eb7483-fixesfrom
fork/282-sync-upstream-main

Conversation

@jeethu

@jeethu jeethu commented Oct 11, 2026 •

Copy link
Copy Markdown
Member

Rebases the fork onto upstream main and drops the February PrismML snapshot of 1-bit affine quantization in favour of upstream's merged version (ml-explore#3161). The fork's own work is re-applied on top as a curated series of 10 patches (18 commits: patch 1 is two, plus one comment-only test commit, four Codex-round fixes and two PR-review fixes), so the fork becomes "upstream main + 10 patches" and later syncs stay small. Tracking: numen-tech/gemma4-qat#282.

Upstream pin: ml-explore/mlx main 06eb7483f03bca49a80e596f4089d02a507a2cf9 ("Fix lost rank output in the distributed launcher (ml-explore#4615)", 2026-10-10). It includes ml-explore#3161, ml-explore#4516, ml-explore#4584, ml-explore#4629 and ml-explore#4641.

  • Base branch: prism-main-06eb7483-fixes, at the pin.
  • Pre-sync tip: prism-0.32.2-fixes at 43e6613b, tagged pre-sync-2026-10-11 (v0.32.2 1f8e74e3 + 59 commits).

Patches

Each line gives the patch's commit(s) on this branch, then the fork commits it was built from.

  1. 1-bit deltas over Add 1-bit affine quantization support (Metal) ml-explore/mlx#3161.
    • Commits: 7c1362e04, 11fde19d0.
    • Sources: b7e6d4d2 (uint32 wide-load qdot<bits=1>), 20bf2eb7 (the 1-pack-per-thread comment; upstream already has the expression), ab914a5b (tail), e11c4bbf (1-bit load comment).
    • Adds the non-aligned-K qmv_fast tail test.
  2. affine_sym (bias-free symmetric 1/2-bit affine), rebuilt on Fix non transposed affine qmm dispatch logic ml-explore/mlx#4392/Fix fp quantized matmul corruption when the quantized dim is not a multiple of 32 ml-explore/mlx#3912.
  3. Implied-bias affine quantized_matmul / dequantize.
  4. qmv_fast routing for affine 4-bit K % 256 and 1-bit K % 512, reconciled with Use the fast qmv kernel for outputs not divisible by 8 ml-explore/mlx#4516/Tune M1 quantized matmul dispatch for six-row large shapes ml-explore/mlx#4629.
  5. qmv_wide routing by bit width and batch, with the single-tile verify for 6–7 rows.
  6. Dispatch/commit/sync/wait counters.
  7. Per-variant kernel build/hash memo, bounded, merged with Sanitize '-' in make_template_hash for metal kernel ml-explore/mlx#4584 (both behaviors kept).
  8. NAX gated off for gen 17 (M5-class) at every NAX dispatch site.
    • Commit: 0bf8503e0.
    • Source: 73db231d.
  9. Precise transcendentals unless MathMode::Fast.
    • Commit: faa5eccf2.
    • Source: c7316268.
  10. spec_decode_verify fast op (CPU composition + fused Metal), carried unchanged.
    • Commit: 385e7031b.
    • Sources: 3e1d206b, 28efbe63, 36030f6f, 49840c9e, e1971918.

Plus bcc992261, a comment-only commit on the qmv tail and qmv_wide test comments.

Codex round 1 fixes. Four new commits on top of bcc992261. All four fix defects in spec_decode_verify that predate this series: they were in the old fork, and patch 10 carried them over unchanged.

Dropped commits

All 59 fork commits (1f8e74e3..43e6613b) are accounted for below. git range-diff pairs only 5 commits once patches are squashed, so the mapping was done by hand and checked by content (see Content check):

Category Count Commits
Mapped to a patch 44 listed per patch above
Dropped, snapshot of ml-explore#3161 (superseded by upstream's merged ml-explore#3161) 5 2fb3d1ad, 9a28ad1a, 73fb1928, 0b631f52, 20bf2eb7 (expression already in ml-explore#3161; its comment is carried in patch 1)
Dropped, merge commits (each has an empty --cc diff: no resolution content) 10 08121fd1, 0f4257d0, 2f0e44f9, 15b0677f, 76e3003b, 4c629475, a7ed9286, 071a92a8, d23d9b96, 43e6613b
Dropped, superseded by upstream beyond the snapshot 0 none

Review-churn commits (Codex rounds, comment condensations) are folded into their patches; none is dropped.

20bf2eb7's expression (packs_per_thread = bits <= 2 ? 1 : 2) is already in upstream ml-explore#3161, and its comment is carried in patch 1.

Content check. Of the 3372 non-trivial lines the fork added (git diff 1f8e74e3 43e6613b), 62 are not verbatim in this branch. Every one is accounted for:

The only fork-touched file this branch leaves alone is benchmarks/python/comparative/bench_mlx.py, which is snapshot-only; ml-explore#3161 has the same change.

The brief range 20bf2eb7..08121fd1 (21 commits) is fully carried: 13 in patch 2, cddb7bed in 5, 73db231d in 8, c7316268 in 9, and five in 10. Its 36 lines that are not verbatim here are:

  • 21 that later fork commits had already replaced by 43e6613b;
  • 15 restructured or reworded as above.

No semantic content is dropped.

qmv_fast tail predicate

Decision: keep upstream's lane predicate aligned_end + simd_lid * values_per_thread < K and ml-explore#4516's partial_rows row clamp. The fork's aligned_end + (simd_lid + 1) * values_per_thread <= K is not used.

The two predicates are equivalent whenever K % values_per_thread == 0, which the host gate guarantees (the half block is 16 × vpt, and K % group_size == 0). The fork's comment explaining reads past K stays.

  • GREEN. The tail shapes route to affine_qmv_fast (pr_1 at N = 67), and both tests pass:

    • test_qmv_fast_tail_non_aligned_k;
    • test_qmv_fast_half_block, whose fast-call x has non-zero memory past K.
  • RED. With a one-extra-lane predicate (< K + vpt), the same tests fail:

    • test_qmv_fast_tail_non_aligned_k: 20 of 20 tail subtests;
    • test_qmv_fast_half_block: 162 subtests, plus 3 gather subtests.

    The patch-4 commit body says "165 + 3"; the correct count is 162 + 3.

NAX gate (patch 8): audited sites

is_nax_available() requires gen >= 18 on every device class. Upstream's rule was gen >= (arch == 'p' ? 18 : 17).

Audit method: git grep -n -iE nax mlx/backend/metal, covering kernels/*.h, jit/ and the JIT and no-JIT kernel getters. Every NAX dispatch decision asks that single predicate:

Site Path Predicate
matmul.cpp:964 steel_matmul_axpby use_nax: thin ml-explore#4654 :995, NAX split-K :1050, fused :1087 is_nax_available()
matmul.cpp:2964 gather_mm_rhs_nax is_nax_available()
matmul.cpp:3043 segmented mm use_nax (:3060, :3103) is_nax_available()
quantized.cpp:1243 qmm → qmm_nax (qmm_t_nax / qmm_n_nax, including the _ib and affine_sym variants) is_nax_available()
quantized.cpp:1522 gather_qmm → gather_qmm_nax is_nax_available()
quantized.cpp:1917 gather_qmm_rhs → gather_qmm_rhs_nax is_nax_available()
quantized.cpp:2468 global-scale gather_qqmm K alignment (ml-explore#4481), which then dispatches via :1917 is_nax_available()
scaled_dot_product_attention.cpp:29 pre-M5 D512 non-NAX path (ml-explore#4518) !is_nax_available()
scaled_dot_product_attention.cpp:412 sdpa_full_self_attention_nax is_nax_available()
scaled_dot_product_attention.cpp:436 D72/D80 padded to the NAX kernel (ml-explore#4455) is_nax_available()
scaled_dot_product_attention.cpp:964 D512 support check (ml-explore#4487) is_nax_available()
scaled_dot_product_attention.cpp:1436 D512 fused eligibility (ml-explore#4487) is_nax_available()
scaled_dot_product_attention.cpp:1445 D256 masked/causal fused NAX (ml-explore#3842; array mask ml-explore#4416) is_nax_available()
scaled_dot_product_attention.cpp:1452 D256 non-NAX fused (ml-explore#4505) !is_nax_available()
gated_delta_update.cpp:18, :24 chunk size C = 16 only with NAX. C = 16 is the only value that selects the gated_delta_*_nax kernels at :133, :335, :369, :416 (ml-explore#4020, ml-explore#4565) is_nax_available()

Nothing else can select a NAX kernel:

  • The get_*_nax_kernel getters are called only from the sites above.
  • get_sdpa_vjp_nax_kernel has no caller.
  • MPP tensor_ops code exists only in *_nax.h / *_nax.metal.
  • Nothing outside mlx/backend/metal mentions NAX.
  • The remaining raw get_architecture_gen() / get_architecture() uses only choose tiles or non-NAX kernels:

This is a static audit on an M3 Max (gen 15, no NAX). The gen-17 runtime check is the iPad M5 leg of gemma4-qat#282.

Precise transcendentals (patch 9) vs upstream ml-explore#4461

ml-explore#4461 changes one op: it respells Sigmoid's exp as metal::precise::exp. It does not touch the compile options.

The fork's rule is wider. On macOS 15 / iOS 18+, mathFloatingPointFunctions is set to Precise for every Safe or Relaxed runtime compile. Other runtime-compiled ops still use unqualified metal:: transcendentals, for example:

  • LogAddExp's exp / log1p;
  • Power's pow;
  • complex exp / log / sin / cos / atan2.

So the fork's rule is kept. The two changes compose: an explicit precise:: call stays precise under every mode.

Tests (Step 6, M3 Max applegpu_g15s)

Build Command Result
metallib (Python, JIT OFF) pytest python/tests/test_quantized.py python/tests/test_fast.py -q 108 passed, 2 skipped, 10659 subtests passed
metallib (Python, JIT OFF) pytest python/tests/test_ops.py -k "exp or log or sin or cos or tanh or erf or softmax or sigmoid or power or sqrt" 29 passed
JIT (Python, -DMLX_METAL_JIT=ON) the same two commands 108 passed, 2 skipped, 10659 subtests passed; 29 passed
JIT doctest (build/) ./build/tests/tests 311 / 311 test cases, 6534 / 6534 assertions
metallib doctest (build-lib/) cmake -B build-lib -DMLX_BUILD_TESTS=ON -DMLX_METAL_JIT=OFF -DCMAKE_BUILD_TYPE=Release, then ./build-lib/tests/tests 311 / 311 test cases, 6534 / 6534 assertions

The gate was rerun on HEAD 2fc35a344, after the Codex round 1 and PR-review fixes. The +4 tests and +31 subtests over the first run are the four new spec_decode tests.

CUDA is not built here. The CUDA edits in patches 3, 6, 7 and 10 come from the fork unchanged.

Codex round

  • P1, strided drafts. The fused Metal kernel indexed dense rows but was bound to strided drafts as given. Transposed, broadcast and sliced drafts gave wrong results: a transposed draft returned n_accepted [1,0] where CPU returns [2,2], and a broadcast draft read past the stored elements.
    • Fix: 8b46f3f58. A non-row-contiguous draft or target is copied to a dense temporary before dispatch. A row-contiguous input takes the same path as before, with no extra copy or dispatch.
    • Test: test_spec_decode_verify_strided_inputs. RED: 4 failed (every Metal draft-view subtest). GREEN: -k spec 2 passed, 86 subtests.
  • P2, K = 0. With K = 0, the CPU composition took min over an empty axis and threw ([min] Cannot min reduce over axis 1 with size 0), while Metal returns a result.
    • Fix: 52d3ef54b. The CPU composition now matches Metal's measured K = 0 result: n_accepted = 0 and committed [B, 1] = the argmax bonus token.
    • Test: test_spec_decode_verify_empty_draft. RED: 2 failed (both CPU subtests). GREEN: -k spec 3 passed, 90 subtests.
  • B = 0 (review follow-up). With B = 0, the op threw on every backend before reaching the kernel ([argmax] Cannot argmax reduce zero size array.). The Metal validation layer reported nothing else.
    • Fix: 2e4e81660. After the shape checks, the op returns int32 zeros of shape [0] and [0, K+1], with no argmax and no primitive, so nothing is dispatched.
    • Test: test_spec_decode_verify_empty_batch, for K = 3 and K = 0, checking shapes and dtypes on Metal and on CPU, and that the two backends match. RED: 6 failed. GREEN: -k spec 4 passed, 96 subtests.
  • K = 0 on Metal (found under the validation layer). With K = 0 and B > 0, the fused kernel was dispatched with an empty draft and therefore no buffer bound at index 0. The result was still correct, but MTL_DEBUG_LAYER=1 MTL_SHADER_VALIDATION=1 aborts with missing Buffer binding at index 0 for draft[0].
    • Fix: 78bdb20eb. K = 0 now takes the composition on every backend.
    • Test: test_spec_decode_verify_empty_draft, run under the validation layer. RED: abort (rc 134). GREEN: the whole -k spec suite passes under the validation layer, 4 passed, 96 subtests.

PR review round

  • P1, NAX gate comment too long (device.cpp). Fix: 87d00a604 condenses it to two lines. The measurements stay in 0bf8503e0 and 73db231d.
  • P2, non-integer drafts truncated. astype(int32) silently turned a floating draft 1.9 into token 1, and accepted bool. Fix: 2fc35a344 rejects non-integer draft_tokens with ValueError; integer widths are still cast to int32. Test: test_spec_decode_verify_rejects_non_integer_drafts. RED: 8 failed (float32 / float16 / bfloat16 / bool × default / CPU). GREEN: -k spec 5 passed, 107 subtests.

jeethu added 16 commits October 11, 2026 11:49
…-aligned-K qmv_fast tail test (gemma4-qat#282)

Patch 1 of the numen-tech/mlx rebase onto upstream main (06eb748), which
carries upstream's final 1-bit affine quantization (ml-explore#3161,
e0408d4). It replaces the fork's five pre-review ml-explore#3161 snapshot commits
(2fb3d1a..20bf2eb) with only the deltas we still carry over upstream.

- qdot<bits=1>: read the packed codes as uint32 words, 32 values per load,
  instead of upstream's byte-at-a-time loop. Same little-endian bit order
  and accumulation order, so bit-exact with the per-byte path.
  Provenance: mlx-swift bc96a51 "metal: uint32 wide load in 1-bit qmv qdot
  (large mobile decode win)", sourced in the fork by b7e6d4d
  (numen-tech/gemma4-qat#179). Every 1-bit qdot caller has
  values_per_thread == 32 and 4-byte aligned weights (qmv_fast, qmv,
  qmv_quad at D == 128); qdot_safe, qouter and dequantize keep upstream's
  byte path.
- qmv_fast_impl packs_per_thread: upstream already took the fork's
  `bits <= 2 ? 1 : 2`; restore the fork's comment (#3).
- qmv_fast_impl partial tail: keep upstream's lane predicate
  `aligned_end + simd_lid * vpt < K` and the partial_rows clamp from
  ml-explore#4516. The fork's `aligned_end + (simd_lid + 1) * vpt <= K`
  is equivalent whenever K % vpt == 0, which the host K gate guarantees.
  Upstream's comment there was stale (it said the 1-bit block is 2048; it
  is 1024 with 1 pack/thread); replace it with the fork's OOB reasoning,
  reworded for this predicate.
- dispatch_qmv: document why the qmv_quad K == 64 route excludes 1-bit
  (a quad lane would hold 16 values, less than one uint32 pack). The
  fork's affine_sym fall-through comment lands with affine_sym (patch 2).

Test: test_qmv_fast_tail_non_aligned_k (GPU vs CPU, 1/4-bit, K = 17/33
groups, N = 67/256). It passes before and after the qdot port. Note that
on this base the qmv_fast tail is unreachable from the host (the 1-bit K
gate is 2048 = 2 * block; other widths gate at block), so these shapes
route to qmv. With an experimental (uncommitted) 1-bit K gate of 512 the
tail runs, the ported qdot plus upstream's predicate matched the CPU on
1-bit K = 1536/2560 tail shapes, and a predicate letting one lane read
past K failed them.
…-proof (gemma4-qat#282)

The comment claimed the shapes run qmv_fast's partial tail and the
partial_rows clamp. They do not on this base: the host K gate is 2048 for
1-bit and 512 for 4-bit, so every shape here routes to plain qmv (or
qmv_wide for M = 3 on gen 15+). Say so, and say when it becomes the tail
guard: once the routing patch widens the gate to 1-bit K % 512 and affine
4-bit K % 256 (gemma4-qat#192).

Add the shapes that the widened gate makes tail-reachable, keeping the
existing ones: 1-bit K in {1536, 2560} x gs {32, 64, 128} and 4-bit
K in {768, 1280} x gs {32, 64}, with N in {67, 256} and M in {1, 3}.
Same GPU-vs-CPU comparison and tolerances (atol = rtol = 1e-2). Skip
without Metal like the neighbouring tests; ruff-format.
…tream main (gemma4-qat#192, ml-explore#282)

Port of numen-tech/mlx PR #8 (923aef6, 7d565ec) onto upstream main.
qmv_fast_k_alignment takes the mode and returns half a kernel block for
affine / affine_sym 1-bit and 4-bit (1-bit K % 512, 4-bit K % 256), the
whole block otherwise. nvfp4 / mxfp4 share bits == 4 but fp_qmv_fast_impl
has no tail, so they keep the whole block; every other width keeps
upstream's routing. qmv keeps upstream's gate shape from ml-explore#4516
(fast = K-aligned, partial_rows = fast && N % 8 != 0); gather_qmv keeps
N % 8 == 0.

Tail predicate: kept upstream's lane predicate
  aligned_end + simd_lid * values_per_thread < K
and the ml-explore#4516 partial_rows row clamp, not the fork's
  aligned_end + (simd_lid + 1) * values_per_thread <= K.
The two agree whenever K % values_per_thread == 0, which the host gate
guarantees (block / 2 = 16 * values_per_thread, and K % group_size
already implies it). Evidence on M3 Max (applegpu_g15s):
- GREEN: test_qmv_fast_tail_non_aligned_k and test_qmv_fast_half_block
  pass; a kernel-name probe (removed) showed the 1-bit K 1536/2560 and
  4-bit K 768/1280 M = 1 subtests on affine_qmv_fast (pr_1 at N = 67).
- RED: with the predicate loosened to `< K + values_per_thread` (one extra
  lane reads past K) both tests fail: 20/20 tail-reachable subtests of
  test_qmv_fast_tail_non_aligned_k and every tail subtest of
  test_qmv_fast_half_block (165 + 3 gather), errors up to ~78.
No read past K shows with upstream's predicate, so the fork's stricter
form is not needed; the OOB-reasoning comment stays.

Tests: test_qmv_fast_half_block (fork's, with the generic baseline moved
to K + gs zero-padded because N % 8 != 0 no longer leaves qmv_fast after
ml-explore#4516; the fast call's x is a view with non-zero memory past K so the test
is a tail guard; fp32 allows 4 ulps for accumulation order, measured
<= 1.1), K = 768 in test_implied_bias_matmul, the "affine quantized_matmul
qmv_fast tail block" doctest (renamed from the fork's "test affine
qmv_fast tail block" so the *quantiz* filter runs it), updated bias-free
doctest case comments. Also: the stale path="qmv_generic" subTest label
is now "n_not_mult_8", and affine_sym_qmv[_fast] note their unused
has_global_scale / results_per_simdgroup slots.
…w verify (gemma4-qat#191, ml-explore#282)

Port of cddb7be (route affine qmv_wide by bit width and batch) and
numen-tech/mlx PR #9 (071a92a..43e6613: run an affine qmv_wide verify of
6 or 7 rows as one tile) onto upstream main.

- use_qmv_wide takes bits and M: affine 1-bit, and affine 2-bit at M < 3,
  stay on qmv; other affine widths take qmv_wide on gen 15+ as before; fp
  modes unchanged; affine_sym still has no qmv_wide kernel.
- qmv_wide runs affine M <= 7 as one tile (vecs_per_tg = M); M >= 8 keeps
  upstream's ceil(M / 5) tiling. quantized.metal instantiates the affine
  and implied-bias qmv_wide kernels for vecs_per_tg 6 and 7.

No conflict with ml-explore#4516, ml-explore#4629 or ml-explore#4641: none touches use_qmv_wide or the
qmv_wide tiling, and ml-explore#4629's M1 batch limit only applies below gen 15,
where affine never takes qmv_wide.

Tests: test_qmv_wide gains 1-bit affine (full sweep and tiny shapes);
test_qmv_wide_tile_matches_split (M = 6/7/8/11/12 bit-identical to the
split calls, verbatim from the fork); the "affine quantized_matmul qmv_wide
single tile of 6 and 7 vectors" doctest (renamed from the fork's "test
affine qmv_wide single tile of 6 and 7 vectors" so the *quantiz* filter
runs it). Routing comments in the tests updated for the 1/2-bit gate.
…#282)

Process-wide Metal work counters (metal::counters() / metal::reset()):
dispatches, commits, explicit stream synchronizations, and host waits on
GPU work. Source: numen-tech/mlx PR #3, `git diff 08121fd 0f4257d`
(5447675, aaefa50, 2e5d6cc, 1fee32b, 9cb5069).

Conflicts with upstream main since v0.32.2:
- fence.cpp: upstream dropped ~FenceImpl (fence is now an
  std::optional<allocator::Data>) and Fence::wait takes the target value
  as a parameter. The fork's `Device device` member and the host-wait
  count in the CPU spin-wait are kept, spinning on upstream's `value`.
- tests/CMakeLists.txt: upstream renamed METAL_TEST_SOURCES to
  GPU_TEST_SOURCES; metal_counters_tests.cpp is appended to the new list.

No new host-wait or dispatch site upstream: CommandEncoder::
dispatch_threadgroups/dispatch_threads and CommandEncoder::synchronize
are still the only dispatch and waitUntilCompleted sites in
mlx/backend/metal, and eval error recovery still flushes via
gpu::synchronize(s, false).

Tests: ./build/tests/tests -tc='*counters*' -> 4 cases, 26 assertions
pass. test_quantized.py + test_fast.py: 102 passed, 2 skipped.
…gemma4-qat#189, ml-explore#282)

fast::metal_kernel's closure keeps a per-kernel memo (shared by all copies
of the closure, mutex-guarded) from each call variant (input dtypes and
pass mode, output dtypes, template argument kind + name + value) to its
kernel name, source and source hash. CustomKernel holds the source in a
shared_ptr<const std::string> with a precomputed hash. The memo holds at
most metal_kernel_max_cached_variants (64) entries and is cleared when
full; entries are shared_ptrs, so a pending array whose entry was evicted
still evaluates. Source: numen-tech/mlx PR #5, f5428f8 + b2c93bd.

Merged with upstream ml-explore#4584 (908148e, `-` sanitized in
make_template_hash): the memo calls make_template_hash unchanged when it
builds a variant, so negative template values still give a valid kernel
name, and the memo key encodes the value's bytes, so -1 and 1 stay
separate variants. Both behaviors are kept; the only textual conflict
was tests/CMakeLists.txt (upstream's METAL_TEST_SOURCES ->
GPU_TEST_SOURCES rename).

Tests: test_fast.py -k "custom_kernel or metal_kernel" -> 11 passed
(incl. ml-explore#4584's test_custom_metal_kernel_negative_template_int and the
fork's test_custom_kernel_many_template_values). ./build/tests/tests
-tc='*metal kernel*,*counters*' -> 8 cases, 778 assertions pass.
test_quantized.py + test_fast.py: 103 passed, 2 skipped.
…l-explore#282)

Source: numen-tech/mlx 73db231 ("gate nax off for gen-17 devices",
fork PR #4). M5-class gen-17 parts (applegpu_g17*) compute wrong
results in the NAX steel-gemm and qmm_t kernels; the fallback bit-matches
stock mlx. is_nax_available() now requires gen >= 18 on every device
class (upstream: gen >= (arch == 'p' ? 18 : 17)). The unused `arch`
local is dropped instead of the fork's `(void)arch`.

The gate is one predicate. Audit of every NAX dispatch decision on
upstream main 06eb748 + patches 1-7 (`git grep -n -iE nax
mlx/backend/metal`, including kernels/*.h, jit/ and the JIT/no-JIT
kernel getters), each site -> the predicate it asks:

  matmul.cpp:964    steel_matmul_axpby use_nax (thin ml-explore#4654 :995, NAX
                    split-K :1050, regular fused :1087) -> is_nax_available()
  matmul.cpp:2964   gather_mm_rhs_nax                    -> is_nax_available()
  matmul.cpp:3043   segmented_mm use_nax (:3060, :3103)  -> is_nax_available()
  quantized.cpp:1243 qmm -> qmm_nax (qmm_t_nax/qmm_n_nax, incl. the _ib
                    and affine_sym variants from patches 2-3)
                                                        -> is_nax_available()
  quantized.cpp:1522 gather_qmm -> gather_qmm_nax      -> is_nax_available()
  quantized.cpp:1917 gather_qmm_rhs -> gather_qmm_rhs_nax
                                                        -> is_nax_available()
  quantized.cpp:2468 global-scale gather_qqmm K alignment (ml-explore#4481; the
                    matrix kernels then dispatch via :1917)
                                                        -> is_nax_available()
  scaled_dot_product_attention.cpp:29   pre-M5 D512 non-NAX path (ml-explore#4518)
                                                        -> !is_nax_available()
  scaled_dot_product_attention.cpp:412  sdpa_full_self_attention_nax
                                                        -> is_nax_available()
  scaled_dot_product_attention.cpp:436  D72/D80 pad to NAX (ml-explore#4455)
                                                        -> is_nax_available()
  scaled_dot_product_attention.cpp:964  D512 support check (ml-explore#4487)
                                                        -> is_nax_available()
  scaled_dot_product_attention.cpp:1436 D512 fused eligibility (ml-explore#4487)
                                                        -> is_nax_available()
  scaled_dot_product_attention.cpp:1445 D256 masked/causal fused NAX
                    (ml-explore#3842, array mask ml-explore#4416)           -> is_nax_available()
  scaled_dot_product_attention.cpp:1452 D256 non-NAX fused (ml-explore#4505)
                                                        -> !is_nax_available()
  gated_delta_update.cpp:18, :24  chunk C = 16 (the only value that
                    selects gated_delta_*_nax kernels, :133/:335/:369/:416)
                    only when NAX (ml-explore#4020, ml-explore#4565)        -> is_nax_available()

No other path selects a NAX kernel: the get_*_nax_kernel getters
(jit_kernels.cpp, nojit_kernels.cpp) are only called from the functions
above, get_sdpa_vjp_nax_kernel has no caller, the MPP tensor_ops code in
kernels/ lives only in *_nax.h/.metal, and nothing outside
mlx/backend/metal mentions NAX. The remaining raw
get_architecture_gen()/get_architecture() uses pick tiles or non-NAX
kernels (qmv batch limits quantized.cpp:99-134, nvfp4 narrow qmv
quantized.cpp:550, affine qmv_wide gen 15+ quantized.cpp:633, gemv
stream matmul.cpp:1416, GEMM_TPARAM/devc tile tunings incl. ml-explore#4447).

This host is an M3 Max (gen 15, no NAX), so the change is a no-op here;
the gen-17 runtime check is the iPad M5 leg (gemma4-qat ml-explore#282 T9/T11).
test_quantized/test_fast/test_blas/test_fast_sdpa: 162 passed, 5 skipped.
…#282)

Source: numen-tech/mlx c731626 (applied unchanged). On macOS 15 /
iOS 18+, set_compile_options set only mathMode, and
MTLCompileOptions.mathFloatingPointFunctions (independent, default Fast)
stayed Fast, so a Safe/Relaxed runtime compile (JIT kernels, compiled
fusions, custom kernels) still emitted the approximate exp/log/sin/cos/pow.
It is now Fast only for MathMode::Fast and Precise otherwise; the
pre-macOS-15 branch keeps setFastMathEnabled.

Checked against upstream ml-explore#4461 (365bd0f, "Use precise::exp in Sigmoid
so compiled and eager sigmoid agree"): that fix only respells Sigmoid's
exp as metal::precise::exp. It does not change compile options, and
other runtime-compiled ops still spell unqualified metal:: transcendentals
(binary_ops.h LogAddExp exp/log1p, Power pow, complex log/exp/sin/cos/
atan2, ...), so upstream covers one op of what this rule covers. Nothing
is dropped. The two compose: an explicit metal::precise:: call stays
precise under every mode, including Fast.

Tests: test_ops.py -k "exp or log or sin or cos or tanh or erf or softmax
or sigmoid or power or sqrt" -> 29 passed. test_compile.py (incl. ml-explore#4461's
test_compile_sigmoid_matches_eager) + test_quantized.py + test_fast.py
(incl. test_custom_metal_kernel_math_mode) -> 174 passed, 2 skipped.
…re#282)

mx.fast.spec_decode_verify(draft_tokens [B, K], target_logits
[B, K+1, V]) -> (n_accepted [B], committed [B, K+1]): greedy speculative
verify. The target argmax runs as a regular op; the SpecDecodeVerify
primitive does the per-row prefix match in a one-thread-per-row Metal
kernel, with an op-composition fallback on CPU / no-GPU / CUDA. Source:
3e1d206 (v1 composition), 28efbe6 (fused Metal v2), 36030f6
(docs), 49840c9 (fuzz), e197191 (test moved into test_fast.py),
taken as `git diff 3e1d206^ e197191` minus quantized.cpp /
test_quantized.py (cddb7be in that range is patch 5).

Applied unchanged, except:
- add/add conflicts with upstream: kernels/CMakeLists.txt (after the
  sdpa_vjp / gated_delta build_kernel lines), no_gpu/primitives.cpp
  (after GatedDeltaUpdate[VJP]::use_fallback), test_fast.py (after
  ml-explore#4584's test_custom_metal_kernel_negative_template_int); both sides kept;
- clang-format (fast.cpp, fast.h) and ruff 0.16.10 format (test_fast.py)
  of the new code, whitespace only.

No caller in NumenKit or Python; carried as is (spec §2.10).
Tests: test_fast.py -k spec -> 1 passed, 76 subtests. test_quantized.py
+ test_fast.py -> 104 passed, 2 skipped.
Comment-only (gemma4-qat ml-explore#282 ruling R12, Task 3 review minors):
- test_qmv_fast_half_block: the "x is a view whose memory past K is
  non-zero" claim holds for B = 0 only. At B = 2
  ensure_row_contiguous_matrix copies the (2, 1, K) slice, so past K
  batch 0 reads batch 1's x (still non-zero) and batch 1 reads past the
  buffer.
- test_implied_bias_matmul: affine qmv_wide routing needs gen 15+.
…xplore#282)

Codex round 1, P1 (pre-existing in the fork's 28efbe6). The fused
kernel indexes dense rows (draft + row * K, target + row * (K + 1)), but
eval_gpu bound its inputs as given, and the wrapper's astype(int32) keeps
an int32 view a view. A transposed draft (logically [[1,2],[3,4]]) gave
n_accepted [1, 0] on Metal vs [2, 2] on CPU; a broadcast [1,2] -> [2,2]
draft made row 1 read past the stored elements; a column slice read the
wrong elements.

eval_gpu now copies a non-row-contiguous input (draft or target) into a
dense temporary (contiguous_copy_gpu + add_temporary, as in
gated_delta_update.cpp) before setting the kernel's pipeline. A
row-contiguous input is returned untouched before the encoder is used:
same path as before, no extra copy or dispatch. Targets come from argmax,
which already returns dense rows from strided logits; the check covers
them anyway.

Test: test_spec_decode_verify_strided_inputs (transposed, broadcast,
sliced and mismatching transposed drafts, plus a transposed logits view;
Metal and CPU vs explicit expected results).
RED (before): 4 failed -- every Metal draft-view subtest, e.g.
  case='transposed' path='default': Lists differ: [1, 0] != [2, 2]
  case='broadcast'  path='default': Lists differ: [2, 0] != [2, 2]
  case='sliced'     path='default': Lists differ: [2, 0] != [2, 2]
  case='mismatch-transposed' path='default': [1, 0] != [1, 2]
(CPU and the transposed-logits subtests passed.)
GREEN: test_fast.py -k spec -> 2 passed, 86 subtests passed.
…e#282)

Codex round 1, P2 (pre-existing in the fork's 3e1d206). With K = 0
the CPU / no-GPU composition took min over an empty axis, which
ops.cpp rejects, so draft [B, 0] with logits [B, 1, V] threw on
stream=mx.cpu while the fused kernel returns a result. Per ruling R16,
K = 0 is not rejected: the composition now returns what Metal returns,
measured first: n_accepted = 0 for every row and committed [B, 1] = the
target argmax token (Metal: B=1 -> [0], [[5]]; B=2 -> [0, 0], [[5], [3]]).

Test: test_spec_decode_verify_empty_draft (B = 1 and 2; Metal and CPU
vs the explicit result, shapes included).
RED (before): 2 failed -- both CPU subtests:
  ValueError: [min] Cannot min reduce over axis 1 with size 0.
(The Metal subtests passed.)
GREEN: test_fast.py -k spec -> 3 passed, 90 subtests passed; CPU and
Metal both give int32 [0] / [[5]].
…atching (ml-explore#282)

Ruling R17 (fix round 1). With B = 0 (drafts [0, K], logits
[0, K+1, V], K = 3 and K = 0) the op threw on every backend before
reaching the fused kernel: `[argmax] Cannot argmax reduce zero size
array.` (also under MTL_DEBUG_LAYER=1 MTL_SHADER_VALIDATION=1, where
nothing else is reported: the throw precedes any encode). The op now
returns int32 zeros [0] and [0, K+1] right after the shape checks, with
no argmax and no SpecDecodeVerify primitive, so nothing is dispatched.
B > 0 inputs take exactly the previous path.

Test: test_spec_decode_verify_empty_batch (K = 3 and 0; default device
and CPU; shape and dtype of both outputs; default vs CPU equal).
RED (before): 6 failed -- the 4 path subtests with the argmax ValueError,
the 2 compare subtests with KeyError (no result to compare).
GREEN: test_fast.py -k spec -> 4 passed, 96 subtests passed; both
backends give n (0,) int32 and committed (0, K+1) int32, also clean
under the Metal debug/shader-validation layers.
…kend (ml-explore#282)

Found while running the R17 B = 0 check under the Metal validation
layer (fix round 1; pre-existing in the fork's 28efbe6). With K = 0
and B > 0 the fused kernel was dispatched with an empty draft, which has
no buffer, so index 0 was left unbound. Without validation the kernel
never reads it (K = 0) and the result was right, but
MTL_DEBUG_LAYER=1 MTL_SHADER_VALIDATION=1 aborts:
  validateComputeFunctionArguments:1043: failed assertion `Compute
  Function(spec_decode_verify): missing Buffer binding at index 0 for
  draft[0].'
K = 0 now takes the op composition on every backend (n_accepted = 0,
committed = the argmax bonus token, the K = 0 result from 52d3ef5),
so the kernel is never bound to an empty draft. K > 0 is unchanged.

Test: test_spec_decode_verify_empty_draft (from 52d3ef5), run under
the validation layer.
RED (before): pytest -k empty_draft under MTL_DEBUG_LAYER=1
MTL_SHADER_VALIDATION=1 -> rc 134, the assertion above.
GREEN: the whole -k spec suite under the validation layer -> 4 passed,
96 subtests (rc 0); without it -> 4 passed, 96 subtests.
@chatgpt-codex-connector

chatgpt-codex-connector Bot commented Oct 11, 2026 •

Copy link
Copy Markdown

Codex Review Summary

This comment shows the latest Codex review activity on this pull request.

Review Status Commit Review trigger
📝 Code Review ✅ Completed 2026-10-11T14:27:41.349894Z 2fc35a3 New commits
ℹ️ About Codex in GitHub

Your team has set up Codex to review pull requests in this repo. Reviews are triggered when you

  • Open a pull request for review
  • Mark a draft as ready
  • Comment "@codex review" or "@codex security review".

Codex reacts with 👀 while any review is running, comments if it has suggestions, and reacts with 👍 once all reviews finish with no findings.

@chatgpt-codex-connector chatgpt-codex-connector Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

💡 Codex Review

Here are some automated review suggestions for this pull request.

Reviewed commit: 78bdb20eb1

ℹ️ About Codex in GitHub

Your team has set up Codex to review pull requests in this repo. Reviews are triggered when you

  • Open a pull request for review
  • Mark a draft as ready
  • Comment "@codex review".

If Codex has suggestions, it will comment; otherwise it will react with 👍.

Codex can also answer questions or update the PR. Try commenting "@codex address that feedback".

Comment thread mlx/backend/metal/device.cpp Outdated
Comment on lines +1006 to +1010
// gen-17 desktop parts (M5-class, applegpu_g17*) take the nax path under
// the old `gen >= 17` gate but compute WRONG results in the nax steel-gemm
// and qmm_t kernels (measured 2026-07-05: fp16 GEMM max|err| ~4 vs fp32 for
// every shape routed to steel_gemm_fused_nax; quantized qmm_t ~400 abs err;
// non-nax fallback bit-matches stock mlx). Require gen >= 18 on every

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Condense the NAX gate comment

Replace this eight-line historical and benchmark narrative with a concise one- or two-line explanation of the gen >= 18 invariant, and keep detailed measurements in issue or commit history. The current block is the excessive inline commentary that this repository explicitly prohibits.

AGENTS.md reference: AGENTS.md:L29-L33

Useful? React with 👍 / 👎.

Comment thread mlx/fast.cpp
// Reuse mlx's fast argmax on-device (result stays on the GPU); the fused
// primitive does only the tiny per-row prefix-match. Primitive inputs are
// int32 token arrays, so the kernel needs no large-V reduction.
auto dft = astype(draft_tokens, int32, s);

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Badge Reject non-integer draft tokens before casting

When draft_tokens has a floating dtype, this unconditional cast silently truncates values before verification; for example, a draft token 1.9 is treated as token 1 and may be accepted instead of producing an input error. Since the API documents token IDs as integer data, validate that the input has an integer dtype before converting integer widths to int32, as the other fast token/index operations do.

Useful? React with 👍 / 👎.

Review feedback on #10 (AGENTS.md code standards: 1-2 line comments). The measurements behind the gate are in 0bf8503 and the source commit 73db231.
…#282)

Review feedback on #10: astype(int32) silently truncated floating drafts (1.9 -> 1) and accepted bool. Integer widths are still cast to int32. test_spec_decode_verify_rejects_non_integer_drafts: RED 8 failed (float32/float16/bfloat16/bool x default/cpu), GREEN after the fix.

@chatgpt-codex-connector chatgpt-codex-connector Bot left a comment

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

💡 Codex Review

Here are some automated review suggestions for this pull request.

Reviewed commit: 2fc35a3447

ℹ️ About Codex in GitHub

Your team has set up Codex to review pull requests in this repo. Reviews are triggered when you

  • Open a pull request for review
  • Mark a draft as ready
  • Comment "@codex review".

If Codex has suggestions, it will comment; otherwise it will react with 👍.

Codex can also answer questions or update the PR. Try commenting "@codex address that feedback".

Comment on lines +983 to +985
// Partial last block (in_vec_size % block_size != 0), reached only when the
// host's K gate (qmv_fast_k_alignment) is finer than block_size: affine
// 1/4-bit take K % (block_size / 2) == 0. A lane whose first value lies

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Condense the qmv tail comment

Replace this 13-line narrative with a one- or two-line statement of the lane-bound invariant, and move the dispatch history, rejected alternative, and issue references outside the source. This block is the excessive inline commentary that the repository's code standard explicitly prohibits.

AGENTS.md reference: AGENTS.md:L29-L33

Useful? React with 👍 / 👎.

Comment thread mlx/fast.cpp
return array::make_arrays(
{Shape{B}, Shape{B, K + 1}},
{int32, int32},
std::make_shared<SpecDecodeVerify>(s, fallback),

Copy link
Copy Markdown

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P1 Badge Register the fused verify primitive for export

When a function containing spec_decode_verify is exported on Metal with B > 0 and K > 0, this path leaves a SpecDecodeVerify primitive in the traced graph, but PrimitiveFactory::factory in mlx/export.cpp has no serializer entry for that primitive. Consequently, mx.export_function throws Unable to serialize primitive SpecDecodeVerify; register the primitive for serialization or use the composition fallback while export tracing.

Useful? React with 👍 / 👎.

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