<!-- llms-explorer concept facts · https://llms-explorer.com/tree/mlx-nax-kernel-gating-and-quantized-matmul/ · pack 2026-10-05 · ~2191 tokens -->

# MLX NAX kernel gating and quantized matmul

> Quantized `qmm` dispatch: `has_nax_kernel = is_nax_available() && (transpose || mode == "affine")`; `nax_aligned = (K % 64 == 0) && (transpose || N % 64 == 0)`; NAX runs when both hold and `(enable_tf32() || x.dtype() != float32)`. Otherwise the classic steel qmm kernel runs. Transposed weights q...

Parent: [Mac local LLMs: MLX kernels, numerics and internals](https://llms-explorer.com/tree/mac-local-llms-mlx-kernels-numerics-and-internals/) · 1 facets · 34 facts · page: https://llms-explorer.com/tree/mlx-nax-kernel-gating-and-quantized-matmul/

## Facts

- Quantized `qmm` dispatch: `has_nax_kernel = is_nax_available() && (transpose || mode == "affine")`; `nax_aligned = (K % 64 == 0) && (transpose || N % 64 == 0)`; NAX runs when both hold and `(enable_tf32() || x.dtype() != float32)`. Otherwise the classic steel qmm kernel runs. Transposed weights qualify in every quantization mode (affine, nvfp4, mxfp4, mxfp8); non-transposed qualifies only for affine and needs N % 64 too. — source: `asserted`
- `qmm_nax` tile: 2x2 simdgroups, bn = bk = 64, bm = 32 when transposed and M <= 32 else 64; only `qmm_t_nax` has a 32-row instantiation. Kernel names are `<mode>_qmm_t_nax_...` and `<mode>_qmm_n_nax_...` with an aligned-N variant. — source: `asserted`
- Gathered quantized matmul (MoE): `gather_qmm_nax` needs transposed weights, K % 64 == 0 and the TF32 clause; `gather_qmm_rhs_nax` needs transposed weights and the TF32 clause. The gather NAX kernels are instantiated with BK = 64 only; another BK makes the kernel-name lookup fail. — source: `asserted`
- The activation-quantized (qqmm) gather path uses matrix kernels when K % 64 == 0 with NAX available (K % 32 otherwise), bf16 activations, 3-D weights and global scales. — source: `asserted`
- No NAX variant exists, in the fetched source, for `qmv`, `qmv_quad`, `qmv_wide`, `qvm`, `qvm_split_k` or `qmm_splitk`. NAX therefore touches quantized matmul only when M is at or above the qmv batch limit and `qmm_splitk` does not intercept. — source: `asserted`
- Interaction with split-K: for transposed B == 1 and M >= limit, `qmm_splitk` is called first. It uses 32x32 tiles (non-NAX `qmm_t_splitk`) with `split_k = max(1, 512 / (ceil(N/32) * ceil(M/32)))`, and falls back to `qmm` (and so NAX) only when split_k computes to 1. split_k >= 2 requires at most 256 tiles, so for M <= 32 any N above 8192 goes to NAX qmm, and for M of 33-64 any N above 4096. Below those sizes, on an M5 the small-M quantized matmul runs on the classic non-NAX split-K kernel. This is derived from the source, not measured. — source: `asserted`
- Dense matmul: `use_nax = is_nax_available() && !complex && (tf32 || a.dtype != float32)`. SIMD split-K runs only when `!use_nax`. NAX split-K runs for batch 1 when `K >= 3 * max(M, N)` or `(max(M, N) <= 1024 && K > 2 * max(M, N))`. NAX split-K tiles default to bm = bn = 128, bk = 512, 4x4 simdgroups with partition size 4096 (halved to 2048 / 1024 / K/2 as K drops below 4096 / 2048 / 1024; 64x64x256 2x2 when K <= 4096 or (M+N)/2 < 512). Otherwise the regular `steel_gemm_fused_nax` kernel runs. — source: `asserted`
- `gemv_wide` precedes the split-K and NAX choice in dense matmul and addmm, so bf16/fp16 `x @ W.T` with M = 2..15 on gen >= 15 never reaches NAX. — source: `asserted`
- NAX also serves gather_mm (`steel_gather_mm_rhs_nax_n`, for sorted right-hand-side gathers when TF32 or non-fp32) and segmented matmul (`steel_gemm_segmented_nax`). — source: `asserted`
- PR 2772 (2025-11) initial NAX matmul, attention and QMM; the gather and split-K NAX kernels and the `bm = 32` quantized variant were added later (kernel comments only; dates not in the sources read). — source: `asserted`
- Main differs from v0.32.0 in SDPA head-dim handling (v0.32.0 gate cited in issue 3897: `q.shape(3) != 80` and head dims (64, 80, 128); main: NAX for (64, 96, 128, 256, 512) plus 72/80 padded to 96). — source: `asserted`
- K not a multiple of 64 silently drops quantized matmul off NAX to the slower classic qmm; group size does not matter for the gate, only K. Common hidden sizes (2048, 2880, 4096, 5120) are multiples of 64; a K such as 2400 is not. — source: `asserted`
- For fp32 activations the TF32 clause means `MLX_ENABLE_TF32=0` pushes quantized matmul with fp32 `x` back to classic kernels; fp16/bf16 activations are unaffected. — source: `asserted`
- The 64-element K alignment also appears in the split-K comment: nvfp4 group size 16 would otherwise over-read, so split K is aligned to max(group_size, 32). — source: `asserted`
- None. Differences from older dossiers are scope-only: they describe NAX as "the prefill path"; the source shows it is one of several large-M branches with a split-K override. — source: `asserted`
- Whether `qmm_splitk` will gain a NAX variant, which would put small-M quantized matmul on NAX for M5; none in sources read. — source: `asserted`
- Measured NAX versus classic qmm throughput by quantization mode (fp modes vs affine) on M5; no source. — source: `asserted`
- Whether non-multiple-of-64 K occurs in popular model shapes enough to matter; not surveyed. — source: `asserted`
- `qmm` dispatch computes `has_nax_kernel = is_nax_available() && (transpose || mode == "affine")` and requires `K % 64 == 0` (and `N % 64 == 0` when not transposed) plus the TF32 clause before calling `qmm_nax`. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/quantized.cpp)
- `qmm_nax` uses wm = wn = 2, bn = bk = 64 and bm = 32 when the weights are transposed and M <= 32, else bm = 64; only `qmm_t_nax` has a 32-row instantiation. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/quantized.cpp)
- `gather_qmm_nax` requires transposed weights and K % 64 == 0 and uses BK = 64 only because other values fail the kernel-name lookup. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/quantized.cpp)
- `gather_qmm_rhs_nax` requires NAX availability, transposed weights and the TF32 clause. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/quantized.cpp)
- The qqmm gather path requires K % 64 == 0 with NAX available (K % 32 otherwise) for its matrix kernels. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/quantized.cpp)
- `QuantizedMatmul::eval_gpu` routes M >= vector_limit with transposed weights and B == 1 to `qmm_splitk`; no NAX branch appears in `qmm_splitk`. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/quantized.cpp)
- `qmm_splitk` uses bm = bn = 32 and `split_k = max(1, 512 / (n_tiles * m_tiles))`, falling back to `qmm` when split_k is 1 or less. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/quantized.cpp)
- `qmm_splitk` aligns K partitions to max(group_size, 32) and decrements split_k until K divides evenly. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/quantized.cpp)
- Dense `use_nax` requires NAX availability, a non-complex dtype and (TF32 on or non-fp32 input). — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/matmul.cpp)
- SIMD split-K is chosen only when NAX is not in use, for batch 1, (ceil(M/16) * ceil(N/16)) at most 2048 (1024 on non-'s'/'d' parts), ceil(K/16) at least 8 and K at least max(M, N). — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/matmul.cpp)
- NAX split-K is chosen for batch 1 when K >= 3 * max(M, N) or (max(M, N) <= 1024 and K > 2 * max(M, N)). — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/matmul.cpp)
- NAX split-K defaults to bm = bn = 128, bk = 512, wm = wn = 4 and partition size 4096, shrinking to 64x64x256 with 2x2 simdgroups when (M+N)/2 < 512 or K <= 4096. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/matmul.cpp)
- The gemv_wide route is called before the split-K and NAX selection in matmul and addmm. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/matmul.cpp)
- Gather matmul with sorted rhs uses `gather_mm_rhs_nax` when NAX is available and (TF32 on or non-fp32), and segmented matmul has a NAX kernel. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/matmul.cpp)
- Under `MLX_METAL_GPU_ARCH=applegpu_g16s` every `is_nax_available()` gate is false because the generation is parsed from the overridden name. — [source](https://raw.githubusercontent.com/ml-explore/mlx/main/mlx/backend/metal/device.cpp)
- On M5, quantized matmul at small M (below the qmv batch limit, or intercepted by split-K) does not run on NAX. — source: `asserted`
