How MTPLX Selects Between the m6 NAX Tile and SIMD Fallback for Verify Kernels
MTPLX selects the m6 NAX tile exclusively when the batch size m equals 6, the output dimension exceeds 100,000, and strict hardware eligibility constraints pass; otherwise, it cascades through split-K NAX variants or falls back to the stock MLX SIMD matmul.
The MTPLX library extends MLX quantized linear operations with specialized verify-shape kernels that optimize small-batch matmuls. When the environment variable MTPLX_NAX_VERIFY is enabled, the dispatcher in mtplx/nax_verify.py intercepts 8-bit quantized linear calls and routes them through either a wide-message NAX tile or standard SIMD implementations based on tensor geometry and hardware capabilities.
When Verify Kernels Activate
Verify kernels engage only when MTPLX_NAX_VERIFY is set to a truthy value and the input tensor produces a batch size m (the product of all non-final dimensions) in the range 4 ≤ m ≤ 6. The patched QuantizedLinear.__call__ method in mtplx/nax_verify.py evaluates this condition before deciding between NAX-accelerated paths and the baseline implementation.
The Selection Decision Tree
For eligible shapes, the dispatcher evaluates three mutually exclusive paths in strict priority order:
1. Wide NAX Tile (vk_qmm_m6)
The wide-message geometry kernel (vk_qmm_m6) is selected only when all four conditions in mtplx/nax_verify.py (lines 78‑84) are satisfied:
- m == 6 exactly
- n ≥ 100,000 (huge output dimension)
- The lane is not disabled:
not lane_disabled("qmm_m6_wide") - Hardware eligibility returns
Truefromvk_eligible_m6(m, k, n, bits, group_size, x.dtype)
This path is designed for high-throughput scenarios where the output matrix width dominates computation.
2. Split-K NAX Tile (vk_qmm_m6_ksplit)
If the wide-tile conditions fail but vk_eligible_ksplit returns True, MTPLX selects the split-K variant vk_qmm_m6_ksplit. This kernel parallelizes the reduction dimension across SIMD groups and serves as the secondary NAX option for shapes that meet eligibility but do not qualify for the wide-message optimization.
3. SIMD Fallback
When neither NAX path is viable, the dispatcher increments the fallback counter _count_qlinear_fallback and executes original(self, x), invoking the stock MLX quantized linear implementation. This fallback runs on standard Apple Silicon SIMD units without NAX/G17 gating and serves as the universal compatibility layer.
Hardware Eligibility Constraints
The vk_eligible_m6 function in mtplx/verify_kernels.py (line 226) enforces strict architectural requirements for the wide NAX tile:
def vk_eligible_m6(m: int, K: int, N: int, bits: int,
group_size: int, dtype) -> bool:
return (
int(bits) in (4, 8)
and int(group_size) in (32, 64, 128)
and dtype in (mx.bfloat16, mx.float16)
and 5 <= int(m) <= 6
and int(K) % 64 == 0
and int(N) % (4 * _m6_nsg()) == 0
)
Key constraints include:
- Bit-width: Only 4-bit or 8-bit affine quantized layouts
- Group size: Must be 32, 64, or 128
- Data types: Restricted to
bfloat16orfloat16 - Dimension alignment: K must be divisible by 64; N must be divisible by 4 ×
_m6_nsg()
The helper _m6_nsg() (lines 65‑67 of verify_kernels.py) reads the runtime knob MTPLX_VK_M6_NSG (default 4) to determine SIMD-groups per threadgroup, allowing fine-tuning of the tile width without recompilation.
Code Examples
Forcing the Wide NAX Tile
import os
import mlx.core as mx
import mlx.nn as nn
# Enable verify kernels and set optional implementation
os.environ["MTPLX_NAX_VERIFY"] = "1"
os.environ["MTPLX_VK_M6_NSG"] = "4"
# Create layer with huge output dimension to trigger wide tile
layer = nn.QuantizedLinear(
in_features=1024,
out_features=120000, # n >= 100,000 required
bits=8
)
# Input yielding m=6 (2*3)
x = mx.random.uniform(shape=(2, 3, 1024))
y = layer(x) # Routes through vk_qmm_m6 if eligible
Monitoring Fallback Behavior
from mtplx.nax_verify import _count_qlinear_fallback
# Reset diagnostic counters
_count_qlinear_fallback.clear()
# Small N dimension prevents wide-tile selection
layer_small_n = nn.QuantizedLinear(1024, 512, bits=8)
x = mx.random.uniform(shape=(2, 3, 1024)) # m=6 but n=512
y = layer_small_n(x)
# Check if SIMD fallback was invoked
print("Fallback count:", _count_qlinear_fallback.get("exact_t0", 0))
Summary
- Activation: Verify kernels require
MTPLX_NAX_VERIFY=1and batch sizes 4 ≤ m ≤ 6. - Wide Tile: The m6 NAX tile (
vk_qmm_m6) activates only when m=6, n≥100,000, the lane is enabled, andvk_eligible_m6passes. - Eligibility: Strict constraints govern bit-width, group size, dtype, and dimension alignment (K%64, N%(4×nsg)).
- Fallback Chain: Wide NAX → Split-K NAX → Stock MLX SIMD, with
_count_qlinear_fallbacktracking SIMD invocations. - Tuning:
MTPLX_VK_M6_NSGcontrols SIMD-group count for the wide tile.
Frequently Asked Questions
What does the m6 NAX tile optimize specifically?
The m6 NAX tile optimizes 8-bit quantized matrix multiplication for verify shapes where the batch dimension m equals 6 and the output dimension N is extremely large (≥100,000). It uses a wide-message geometry that maximizes SIMD-group utilization across the output width, reducing memory bandwidth pressure compared to standard tiled approaches.
Why does MTPLX fall back to SIMD even when NAX is enabled?
MTPLX falls back to the stock MLX SIMD implementation when any eligibility constraint fails—such as unsupported group sizes, misaligned dimensions, disabled kernel lanes, or output dimensions below the 100,000 threshold for the wide tile. This ensures functional correctness across all hardware configurations while allowing NAX acceleration only for verified-safe shapes.
How can I force the split-K kernel instead of the wide tile?
You cannot directly force the split-K kernel when the wide tile is eligible without modifying the source code in mtplx/nax_verify.py. However, you can disable the wide tile lane specifically by setting the environment variable that triggers lane_disabled("qmm_m6_wide") to True, which causes the dispatcher to evaluate the split-K path next.
What is the performance impact of the SIMD fallback?
The SIMD fallback invokes the standard MLX QuantizedLinear implementation, which runs on baseline Apple Silicon SIMD units without NAX-specific micro-architectural optimizations. For verify shapes (small m), this typically results in lower compute utilization and higher latency compared to the NAX paths, which is why MTPLX tracks fallback frequency via _count_qlinear_fallback for profiling purposes.
Have a question about this repo?
These articles cover the highlights, but your codebase questions are specific. Give your agent direct access to the source. Share this with your agent to get started:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →