# How the NAX Verify Kernel Path Improves Decode Performance in MTPLX

> Discover how the NAX verify kernel path boosts MTPLX decode performance by 60% on Apple M4/M5 GPUs. Optimized Metal kernels slash latency by removing threadgroup barriers and enhancing SIMD cooperation.

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

---

**The NAX verify kernel path reduces decode latency by approximately 60% by replacing the stock MLX quantized matrix multiplication with a specialized Metal kernel optimized for Apple M4/M5-class GPUs, eliminating threadgroup barriers and optimizing SIMD-group cooperation for small batch verification shapes.**

The NAX verify kernel path is a low-latency optimization integrated into MTPLX that specifically targets speculative decoding verification workloads. When activated via the `MTPLX_NAX_VERIFY` environment variable, MTPLX bypasses the generic MLX `qmm` implementation in favor of hand-tuned Metal kernels (`vk_qmm_m4` and `vk_qmm_m6`) that align threadgroup geometry with the NAX hardware's compute characteristics, cutting per-call latency from ~3.5ms to ~1.38ms.

## Architectural Differences: Stock MLX vs. NAX Verify

The NAX verify kernel path diverges significantly from the stock MLX quantized matrix multiplication (`qmm`) implementation across four critical dimensions:

**Target Matrix Shapes**
- **Stock MLX qmm**: Optimized for `M = 1` (single-token decode) or very large `M` (prefill), using a one-size-fits-all approach that becomes inefficient for intermediate batch sizes.
- **NAX Verify Kernel**: Purpose-built for `M = 4-6` (the "verify" shapes D3-D5 encountered when validating draft tokens), with fixed padding to `M = 4` on M4 GPUs and `M = 6` on M6 GPUs to enable kernel caching without recompilation.

**Threadgroup Layout**
- **Stock MLX qmm**: Distributes work across many SIMD-groups with barrier-based K-splits, causing scheduler thrashing when processing wide matrices (e.g., language model heads with large `N` dimensions).
- **NAX Verify Kernel**: Consolidates work into a single heavy threadgroup containing `NSG` SIMD-groups (8 for M4, 4 for M6) that operate without K-splits or barriers, eliminating synchronization overhead.

**Memory Access Patterns**
- **Stock MLX qmm**: Uses pack-interleaved 32-bit weight loads where SIMD-groups share column tiles, leading to tiny-tile thrashing on large-N layers.
- **NAX Verify Kernel**: Assigns each SIMD-group exclusive ownership of a `BN = 4` column tile, loading it once into registers and retaining it for the entire computation, which maximizes data reuse.

**Register Utilization**
- **Stock MLX qmm**: Uses approximately 24 accumulators per thread but spreads them across competing SIMD-groups, creating occupancy pressure and scheduling overhead.
- **NAX Verify Kernel**: Maintains exactly 24 accumulators per thread—the proven ceiling for this architecture—but coordinates all SIMD-groups to cooperate without barriers, maintaining high occupancy while minimizing per-call dispatch overhead.

## How the NAX Verify Kernel Path Eliminates Bottlenecks

The performance improvement stems from a *msg-geometry* design implemented in [`mtplx/verify_kernels.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/verify_kernels.py) using `mx.fast.metal_kernel`. This design exploits the NAX hardware's higher compute-to-memory ratio through three architectural choices:

**Single SIMD-Group Column Ownership**
Each SIMD-group processes exactly one `BN = 4` column tile independently. This eliminates the column-sharing contention that causes thrashing in the stock implementation, as noted in the source comments of [`verify_kernels.py`](https://github.com/youssofal/MTPLX/blob/main/verify_kernels.py) (lines 15-16).

**Lane-Strided K-Reduction Without Barriers**
The kernel performs the full K-dimension reduction lane-strided within each SIMD-group. By keeping the reduction entirely within SIMD-group boundaries, the NAX verify kernel path avoids threadgroup-wide barriers entirely, removing a major source of latency in the standard decode path.

**Hardware-Aligned Threadgroup Geometry**
The kernel exposes exactly `NSG` SIMD-groups per threadgroup—matching the NAX GPU's physical SIMD lane count (8 for M4, 4 for M6). This alignment keeps the execution units saturated while the single threadgroup design ensures the Metal scheduler does not interleave other work, preventing context-switch overhead.

## Implementation and Configuration

The NAX verify kernel path is gated behind automatic hardware detection and explicit environment configuration. According to the source code analysis of [`mtplx/server/openai.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/server/openai.py) (lines 1052-1062), the runtime inspects `MTPLX_NAX_VERIFY` at request initialization and injects it into profile overrides:

```python

# -------------------------------------------------

# Enable the NAX verify kernel path

# -------------------------------------------------

import os
os.environ["MTPLX_NAX_VERIFY"] = "1"   # Force NAX verify kernel selection

```

When enabled, the profile logic in [`mtplx/profiles.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/profiles.py) forces the dispatcher to select `vk_qmm_m4` or `vk_qmm_m6` instead of the generic `qmm` path. The kernels are compiled and cached using `mx.fast.metal_kernel` with fixed M-dimension padding to avoid JIT compilation overhead during inference.

To use the verify kernel explicitly in custom code:

```python
import mlx.core as mx
from mtplx.verify_kernels import vk_qmm_m4, vk_qmm_m6

# Activation matrix with verify shape (M=4 for M4-class GPUs)

x2 = mx.random.normal([4, 4096], dtype=mx.bfloat16)

# Quantized weight placeholders

w_q = mx.random.uniform([4096, 4096], dtype=mx.uint32)
scales = mx.random.uniform([4, 4096], dtype=mx.bfloat16)
biases = mx.random.uniform([4, 4096], dtype=mx.bfloat16)

# Execute via NAX verify path

y = vk_qmm_m4(x2, w_q, scales, biases)
print(y.shape)  # (4, 4096)

```

To revert to standard MLX behavior, clear the environment variable or set it to `"0"`:

```python
os.environ["MTPLX_NAX_VERIFY"] = "0"  # Fallback to stock MLX qmm

```

## Performance Characteristics

The NAX verify kernel path delivers measurable gains on NAX-capable Apple Silicon (M4/M5-class GPUs). Benchmarks in [`tests/test_sdpa_nax_flash.py`](https://github.com/youssofal/MTPLX/blob/main/tests/test_sdpa_nax_flash.py) and [`tests/test_sdpa_nax_tile.py`](https://github.com/youssofal/MTPLX/blob/main/tests/test_sdpa_nax_tile.py) demonstrate:

- **Baseline Latency**: ~3.5ms per verify call using stock MLX qmm on large-N matrices (approximately 1.18× the weight-stream bandwidth floor).
- **Optimized Latency**: ~1.38ms per call after applying the NAX verify patch.
- **Net Improvement**: Approximately 60% speed-up, closing the gap toward the theoretical weight-stream bandwidth limit.

The [`mtplx/verify_qmv.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/verify_qmv.py) wrapper handles automatic fallback logic, ensuring that if `MTPLX_NAX_VERIFY` is unset or the hardware is incompatible, the system transparently reverts to stock kernels without runtime errors.

## Summary

The NAX verify kernel path in MTPLX delivers substantial decode performance improvements through specialized Metal kernel architecture:

- **Eliminates threadgroup barriers** by confining K-reductions to individual SIMD-groups, removing synchronization overhead.
- **Optimizes memory access** via exclusive SIMD-group ownership of `BN = 4` column tiles, preventing tiny-tile thrashing on wide matrices.
- **Aligns with NAX hardware** using threadgroup geometries matching the GPU's SIMD lane count (8 for M4, 4 for M6).
- **Reduces latency by ~60%** on verification-shaped matrices (M=4-6), cutting per-call time from 3.5ms to 1.38ms.
- **Activates via environment variable** `MTPLX_NAX_VERIFY=1`, with automatic hardware detection in [`mtplx/profiles.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/profiles.py) and request-time configuration in [`mtplx/server/openai.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/server/openai.py).

## Frequently Asked Questions

### What is the NAX verify kernel path in MTPLX?

The NAX verify kernel path is a specialized matrix multiplication implementation in MTPLX designed for Apple M4/M5-class GPUs (NAX architecture). Located in [`mtplx/verify_kernels.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/verify_kernels.py), it provides optimized Metal kernels (`vk_qmm_m4` and `vk_qmm_m6`) that replace the stock MLX quantized matmul during speculative decoding verification, specifically targeting small batch sizes of 4-6 tokens.

### How do I enable the NAX verify kernel path?

Set the environment variable `MTPLX_NAX_VERIFY=1` before starting your MTPLX server or script. The runtime checks this variable in [`mtplx/server/openai.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/server/openai.py) (lines 1052-1062) and updates the compute profile to select the verify kernels. You can also import directly from `mtplx.verify_kernels` for programmatic use in custom inference code.

### Which hardware supports the NAX verify kernel path?

The NAX verify kernel path requires Apple Silicon with NAX-class GPUs, specifically the M4 and M5 series chips. The kernel automatically detects the specific variant (M4 vs. M6) and selects the appropriate implementation (`vk_qmm_m4` using 8 SIMD-groups or `vk_qmm_m6` using 4 SIMD-groups) based on the hardware capabilities exposed via Metal.

### Why does the NAX verify kernel path focus on M=4-6 matrix shapes?

These shapes correspond to the "verify" batch sizes (D3-D5) encountered when validating draft tokens in speculative decoding. Unlike standard decode (M=1) or prefill (large M), these intermediate sizes suffer from scheduling overhead in generic implementations. The NAX verify kernel path pads and optimizes specifically for these dimensions, caching the compiled kernel to avoid recompilation overhead while maximizing register utilization per SIMD-group.