Repository navigation
Conversation
…-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.
…m main (gemma4-qat#179, ml-explore#282)
…m main (gemma4-qat#190, ml-explore#282)
…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.
Codex Review SummaryThis comment shows the latest Codex review activity on this pull request.
ℹ️ About Codex in GitHubYour team has set up Codex to review pull requests in this repo. Reviews are triggered when you
Codex reacts with 👀 while any review is running, comments if it has suggestions, and reacts with 👍 once all reviews finish with no findings. |
There was a problem hiding this comment.
💡 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".
| // 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 |
There was a problem hiding this 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 👍 / 👎.
| // 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); |
There was a problem hiding this comment.
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 👍 / 👎.
There was a problem hiding this comment.
💡 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".
| // 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 |
There was a problem hiding this 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 👍 / 👎.
| return array::make_arrays( | ||
| {Shape{B}, Shape{B, K + 1}}, | ||
| {int32, int32}, | ||
| std::make_shared<SpecDecodeVerify>(s, fallback), |
There was a problem hiding this comment.
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 👍 / 👎.
Rebases the fork onto upstream
mainand 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
main06eb7483f03bca49a80e596f4089d02a507a2cf9("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.prism-main-06eb7483-fixes, at the pin.prism-0.32.2-fixesat43e6613b, taggedpre-sync-2026-10-11(v0.32.21f8e74e3+ 59 commits).Patches
Each line gives the patch's commit(s) on this branch, then the fork commits it was built from.
7c1362e04,11fde19d0.b7e6d4d2(uint32 wide-loadqdot<bits=1>),20bf2eb7(the 1-pack-per-thread comment; upstream already has the expression),ab914a5b(tail),e11c4bbf(1-bit load comment).qmv_fasttail test.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.db2bcd84a.af01cc55,b7e6d4d2,2c2d8a4e,329b29e8,ab914a5b,ee91d6ac,7072c9a2,84b42176,ff1deee5,b4750de8,aa25bdde,e11c4bbf,62880057(e1971918..08121fd1minusdevice.cpp).quantized_matmul/dequantize.276d40f02.0e7a5b09,5a7aeace,6fa3f0cc,477f91ad,a84b8d07,499fc3e2; PR Fork/190 implied bias review fixes #7 =d0f0bb76,4e7ceeab,d3445fa8,555c386a.qmv_fastrouting 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.bc90a8d72.923aef64,7d565ec4.qmv_widerouting by bit width and batch, with the single-tile verify for 6–7 rows.b4f0dd842.cddb7bed; fork PR Run an affine qmv_wide verify of 6 or 7 rows as one tile (gemma4-qat #191) #9 =202b82fa,c74cbf19,13d62578,d5da0546.5f515eac4.54476750,aaefa500,2e5d6ccd,1fee32b6,9cb5069a.f0980d577.f5428f85,b2c93bd1.0bf8503e0.73db231d.MathMode::Fast.faa5eccf2.c7316268.spec_decode_verifyfast op (CPU composition + fused Metal), carried unchanged.385e7031b.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 inspec_decode_verifythat predate this series: they were in the old fork, and patch 10 carried them over unchanged.8b46f3f58fast: spec_decode_verify materializes non-row-contiguous drafts (Scatter gradient implementations ml-explore/mlx#282)52d3ef54bfast: spec_decode_verify CPU path handles K = 0 like Metal (Scatter gradient implementations ml-explore/mlx#282)2e4e81660fast: spec_decode_verify returns empty results for B = 0 without dispatching (Scatter gradient implementations ml-explore/mlx#282)78bdb20ebfast: spec_decode_verify takes the composition for K = 0 on every backend (Scatter gradient implementations ml-explore/mlx#282)Dropped commits
All 59 fork commits (
1f8e74e3..43e6613b) are accounted for below.git range-diffpairs only 5 commits once patches are squashed, so the mapping was done by hand and checked by content (see Content check):2fb3d1ad,9a28ad1a,73fb1928,0b631f52,20bf2eb7(expression already in ml-explore#3161; its comment is carried in patch 1)--ccdiff: no resolution content)08121fd1,0f4257d0,2f0e44f9,15b0677f,76e3003b,4c629475,a7ed9286,071a92a8,d23d9b96,43e6613bReview-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:partial_rows/has_global_scale(Use the fast qmv kernel for outputs not divisible by 8 ml-explore/mlx#4516, [Metal] global scale for qmm ml-explore/mlx#4458);test_qmv_fast_half_blockbaseline, reworked because Use the fast qmv kernel for outputs not divisible by 8 ml-explore/mlx#4516 made N % 8 != 0 takeqmv_fast;METAL_TEST_SOURCES→GPU_TEST_SOURCESrename;gather_qmmcall taking upstream's newglobal_scaleargument;spec_decode_verify;(void)arch;, dropped together with the unused local;e11c4bbf's two-line condensation of the patch-9 compile-option comment. Patch 9 keepsc7316268's original four-line comment instead. This is comment-only.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,cddb7bedin 5,73db231din 8,c7316268in 9, and five in 10. Its 36 lines that are not verbatim here are:43e6613b;No semantic content is dropped.
qmv_fast tail predicate
Decision: keep upstream's lane predicate
aligned_end + simd_lid * values_per_thread < Kand ml-explore#4516'spartial_rowsrow clamp. The fork'saligned_end + (simd_lid + 1) * values_per_thread <= Kis 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_1at 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 wasgen >= (arch == 'p' ? 18 : 17).Audit method:
git grep -n -iE nax mlx/backend/metal, coveringkernels/*.h,jit/and the JIT and no-JIT kernel getters. Every NAX dispatch decision asks that single predicate:matmul.cpp:964steel_matmul_axpbyuse_nax: thin ml-explore#4654:995, NAX split-K:1050, fused:1087is_nax_available()matmul.cpp:2964gather_mm_rhs_naxis_nax_available()matmul.cpp:3043use_nax(:3060,:3103)is_nax_available()quantized.cpp:1243qmm→qmm_nax(qmm_t_nax/qmm_n_nax, including the_ibandaffine_symvariants)is_nax_available()quantized.cpp:1522gather_qmm→gather_qmm_naxis_nax_available()quantized.cpp:1917gather_qmm_rhs→gather_qmm_rhs_naxis_nax_available()quantized.cpp:2468gather_qqmmK alignment (ml-explore#4481), which then dispatches via:1917is_nax_available()scaled_dot_product_attention.cpp:29!is_nax_available()scaled_dot_product_attention.cpp:412sdpa_full_self_attention_naxis_nax_available()scaled_dot_product_attention.cpp:436is_nax_available()scaled_dot_product_attention.cpp:964is_nax_available()scaled_dot_product_attention.cpp:1436is_nax_available()scaled_dot_product_attention.cpp:1445is_nax_available()scaled_dot_product_attention.cpp:1452!is_nax_available()gated_delta_update.cpp:18,:24gated_delta_*_naxkernels at:133,:335,:369,:416(ml-explore#4020, ml-explore#4565)is_nax_available()Nothing else can select a NAX kernel:
get_*_nax_kernelgetters are called only from the sites above.get_sdpa_vjp_nax_kernelhas no caller.tensor_opscode exists only in*_nax.h/*_nax.metal.mlx/backend/metalmentions NAX.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 asmetal::precise::exp. It does not touch the compile options.The fork's rule is wider. On macOS 15 / iOS 18+,
mathFloatingPointFunctionsis set to Precise for every Safe or Relaxed runtime compile. Other runtime-compiled ops still use unqualifiedmetal::transcendentals, for example:LogAddExp's exp / log1p;Power's pow;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)pytest python/tests/test_quantized.py python/tests/test_fast.py -qpytest 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"-DMLX_METAL_JIT=ON)build/)./build/tests/testsbuild-lib/)cmake -B build-lib -DMLX_BUILD_TESTS=ON -DMLX_METAL_JIT=OFF -DCMAKE_BUILD_TYPE=Release, then./build-lib/tests/testsThe 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
[1,0]where CPU returns[2,2], and a broadcast draft read past the stored elements.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_spec_decode_verify_strided_inputs. RED: 4 failed (every Metal draft-view subtest). GREEN:-k spec2 passed, 86 subtests.minover an empty axis and threw ([min] Cannot min reduce over axis 1 with size 0), while Metal returns a result.52d3ef54b. The CPU composition now matches Metal's measured K = 0 result: n_accepted = 0 and committed[B, 1]= the argmax bonus token.test_spec_decode_verify_empty_draft. RED: 2 failed (both CPU subtests). GREEN:-k spec3 passed, 90 subtests.[argmax] Cannot argmax reduce zero size array.). The Metal validation layer reported nothing else.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_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 spec4 passed, 96 subtests.MTL_DEBUG_LAYER=1 MTL_SHADER_VALIDATION=1aborts withmissing Buffer binding at index 0 for draft[0].78bdb20eb. K = 0 now takes the composition on every backend.test_spec_decode_verify_empty_draft, run under the validation layer. RED: abort (rc 134). GREEN: the whole-k specsuite passes under the validation layer, 4 passed, 96 subtests.PR review round
device.cpp). Fix:87d00a604condenses it to two lines. The measurements stay in0bf8503e0and73db231d.astype(int32)silently turned a floating draft1.9into token1, and accepted bool. Fix:2fc35a344rejects non-integerdraft_tokenswithValueError; 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 spec5 passed, 107 subtests.