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

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:

// 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:

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. This C++ bridge uses nanobind to expose Metal kernel dispatch to Python while eliminating the PyTorch MPS overhead:

// 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 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, 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.

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:

Share the following with your agent to get started:
curl -s "https://instagit.com/install.md"

Works with
Claude Codex Cursor VS Code OpenClaw Any MCP Client

Maintain an open-source project? Get it listed too →