# How MTPLX Implements NAX Kernels Using Metal 4 Tensor-Ops for Accelerated Attention

> Discover how MTPLX accelerates NAX kernels with Metal 4 tensor-ops. Explore custom N-lane vector types and JIT-compiled shaders for enhanced performance. Learn more now.

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

---

**MTPLX implements NAX (N‑axis) kernels by leveraging Metal 4's tensor operations through custom N‑lane vector types and JIT-compiled compute shaders dispatched via MLX's native Metal command encoder.**

MTPLX is an open-source inference engine that accelerates transformer attention on Apple Silicon by bypassing the PyTorch MPS bridge. The library implements **NAX kernels**—operations optimized along the batch/lane (N) axis—using Metal 4's native tensor-ops to execute matrix multiplications via SIMD-group instructions. This implementation delivers low-latency paged attention by compiling Metal shaders at runtime and dispatching them through MLX's C++ API.

## NAX Kernel Architecture and Metal 4 Tensor Operations

MTPLX refers to its attention kernels as **NAX kernels** because they map the logical N‑axis (batch dimension) onto Metal's SIMD-group lanes using custom vector types. In `vllm_metal/metal/kernels_v2/pagedattention.metal`, the code defines N‑lane vector primitives that align with Metal 4's tensor-ops hardware:

```metal
// Custom vector types for N‑lane operations
typedef ushort2 uchar2x8;  // 8‑lane uchar vector
typedef ushort4 uchar4x8;  // 16‑lane vector

```

These types enable the kernel to load and process query, key, and value tensors in vectorized chunks that match the SIMD-group width. The actual attention computation relies on **Metal 4 tensor-ops**—specifically `simdgroup_matrix_multiply` instructions generated by the compiler when operating on these vector types. The kernel performs scaled dot-product attention through a sequence of element-wise tensor operations:

```metal
kernel void paged_attention_kernel(
    device const float* query [[buffer(0)]],
    device const float* key   [[buffer(1)]],
    device const float* value [[buffer(2)]],
    device float* output      [[buffer(3)]],
    uint2 gid [[thread_position_in_grid]],
    uint2 tid [[thread_position_in_threadgroup]],
    threadgroup float shared_mem[/*...*/]) {
    
    // Vectorized loads utilizing N‑lane layout
    float4 q = *(device float4*)(query + gid.x * 4);
    float4 k = *(device float4*)(key + gid.x * 4);
    float4 v = *(device float4*)(value + gid.x * 4);
    
    // Tensor‑ops: element‑wise multiply, exp, and combine
    float4 dot = q * k;      // SIMD-group vector multiply
    float4 attn = exp(dot);  // Element-wise exponential
    float4 out = attn * v;   // Final combination
    
    *(device float4*)(output + gid.x * 4) = out;
}

```

## JIT Compilation via the C++ Nanobind Bridge

Rather than pre-compiling Metal libraries, MTPLX uses runtime JIT compilation through [`vllm_metal/metal/paged_ops.cpp`](https://github.com/youssofal/MTPLX/blob/main/vllm_metal/metal/paged_ops.cpp). This C++ bridge uses **nanobind** to expose Metal kernel dispatch to Python while eliminating the PyTorch MPS overhead:

```cpp
// Core function: paged_attention
static py::object paged_attention(py::list inputs) {
    // Evaluate inputs, ensure they are on Metal device
    auto ops = get_ops();  // Retrieve cached compiled kernels
    
    // Dispatch kernel with MLX Metal command encoder
    mlx::metal::CommandEncoder encoder;
    encoder.set_compute_pipeline_state(ops->paged_attention);
    encoder.set_buffer(/* query buffer */);
    encoder.set_buffer(/* key buffer */);
    encoder.set_buffer(/* value buffer */);
    encoder.dispatch_threads(grid_size, threadgroup_size);
    encoder.end_encoding();
    
    return /* mlx::array result */;
}

```

The `get_ops()` function implements a compilation cache: on first invocation, it reads the `.metal` source files from `vllm_metal/metal/kernels_v2/`, inlines `#include` directives, and invokes MLX's Metal compiler to generate `MTLComputePipelineState` objects. Subsequent calls reuse these compiled pipelines, eliminating shader compilation overhead during inference.

## Thread-Group Tiling and Shared Memory Layout

To maximize tensor-op utilization, MTPLX kernels implement explicit **thread-group tiling** strategies. The `paged_attention_kernel` allocates `threadgroup` shared memory tiles sized to match Metal 4's matrix-multiply units (typically 16×16 for FP16 operations). Each SIMD-group processes a tile of the attention matrix, with thread positions mapped via `thread_position_in_grid` and `thread_position_in_threadgroup` attributes.

This tiling approach ensures that:
- **Memory coalescing**: Vectorized loads (`float4`, `half4`) saturate memory bandwidth
- **Tensor-op occupancy**: Tile dimensions align with SIMD-group matrix-multiply capabilities
- **Register pressure reduction**: Shared memory caches KV-cache blocks to avoid redundant global memory loads

## Summary

- **NAX kernels** in `vllm_metal/metal/kernels_v2/pagedattention.metal` utilize custom `uchar2x8` and `ushort4` vector types to map the N‑axis onto Metal SIMD lanes.
- **Metal 4 tensor-ops** accelerate attention through compiler-generated `simdgroup_matrix_multiply` instructions operating on these N‑lane vectors.
- **JIT compilation** in [`vllm_metal/metal/paged_ops.cpp`](https://github.com/youssofal/MTPLX/blob/main/vllm_metal/metal/paged_ops.cpp) caches compiled `MTLComputePipelineState` objects via MLX's runtime, avoiding pre-compilation dependencies.
- **MLX integration** bypasses the PyTorch MPS bridge by using `mlx::metal::CommandEncoder` for direct kernel dispatch.
- **Thread-group tiling** aligns computation with GPU matrix units via shared memory tiles and vectorized memory access patterns.

## Frequently Asked Questions

### What does NAX mean in MTPLX's kernel naming convention?

**NAX refers to "N‑axis" kernels** that optimize operations along the batch/lane dimension. In MTPLX, this specifically describes attention kernels that use custom N‑lane vector types (e.g., `uchar2x8`) to align data with SIMD-group lanes, enabling efficient tensor-op execution on the N-axis of query/key/value tensors.

### Why does MTPLX JIT-compile Metal shaders instead of using pre-compiled binaries?

**JIT compilation ensures Metal 4 feature availability** across different macOS versions and GPU generations. By compiling `pagedattention.metal` at runtime via [`paged_ops.cpp`](https://github.com/youssofal/MTPLX/blob/main/paged_ops.cpp), MTPLX can detect Metal 4 capabilities and optimize tensor-op usage dynamically, while maintaining compatibility with older hardware through fallback paths.

### How do Metal 4 tensor-ops improve attention performance compared to standard SIMD?

**Metal 4 tensor-ops provide dedicated matrix-multiply acceleration** through instructions like `simdgroup_matrix_multiply`. While standard SIMD performs element-wise operations, tensor-ops execute fused multiply-accumulate on 16×16 tiles, delivering 2×–3× higher throughput for the Q×K^T and attn×V matrix operations in attention layers.

### Can MTPLX NAX kernels run on Metal 3 or older GPUs?

**No, the tensor-op implementation requires Metal 4**. The kernels rely on SIMD-group matrix-multiply instructions and vector types introduced in Metal 4; running on older Metal versions would fall back to less efficient scalar or basic SIMD implementations, though the current codebase does not expose these fallbacks explicitly.