# How MTPLX Selects Between the m6 NAX Tile and SIMD Fallback for Verify Kernels

> Discover how MTPLX intelligently chooses between m6 NAX tile and SIMD fallback for verify kernels. Learn the specific conditions and fallback logic for optimal performance.

- Repository: [Youssof Altoukhi/MTPLX](https://github.com/youssofal/MTPLX)
- Tags: internals
- Published: 2026-09-05

---

**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`](https://github.com/youssofal/MTPLX/blob/main/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`](https://github.com/youssofal/MTPLX/blob/main/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`](https://github.com/youssofal/MTPLX/blob/main/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 `True` from `vk_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`](https://github.com/youssofal/MTPLX/blob/main/mtplx/verify_kernels.py) (line 226) enforces strict architectural requirements for the wide NAX tile:

```python
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 `bfloat16` or `float16`
- **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`](https://github.com/youssofal/MTPLX/blob/main/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

```python
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

```python
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=1` and 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, and `vk_eligible_m6` passes.
- **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_fallback` tracking SIMD invocations.
- **Tuning:** `MTPLX_VK_M6_NSG` controls 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`](https://github.com/youssofal/MTPLX/blob/main/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.