# How the Paged KV Cache Works in MTPLX: Architecture and Implementation

> Explore the paged KV cache in MTPLX. Learn about its 4D tensor layout, lazy allocation, dynamic growth, and quantization support for efficient vLLM-Metal kernel integration.

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

---

**MTPLX implements a paged key-value cache using a 4-D tensor layout compatible with vLLM-Metal kernels, supporting lazy allocation, dynamic growth, and optional quantization paths via the `VllmMetalPagedKVCache` class in [`mtplx/cache_state.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/cache_state.py).**

The MTPLX inference engine, available at the `youssofal/MTPLX` repository, optimizes transformer attention through a high-performance paged KV cache implementation. This system mirrors the memory layout expected by vLLM-Metal paged-attention kernels while maintaining compatibility with MLX operations. At its core, the `VllmMetalPagedKVCache` class (defined in [`mtplx/cache_state.py`](https://github.com/youssofal/MTPLX/blob/main/mtplx/cache_state.py) at lines 793–825) manages key-value storage through a block-based paging mechanism that minimizes memory fragmentation and supports advanced quantization strategies.

## Physical Layout and Memory Organization

### 4-D Tensor Structure

The paged KV cache stores tensors in a four-dimensional array with shape `[num_blocks, block_size, num_kv_heads, head_dim]`. By default, `block_size` is set to **16 tokens per page**, while `num_blocks` defaults to **1024 maximum pages**. This layout aligns precisely with the expectations of Metal paged-attention kernels exported from [`vllm_metal/metal/__init__.py`](https://github.com/youssofal/MTPLX/blob/main/vllm_metal/metal/__init__.py), enabling efficient GPU memory access patterns during attention computation.

### Contiguous Addressing Scheme

Tokens are addressed contiguously from position zero, creating a deterministic mapping between token indices and their physical block locations. Each token resides at a specific block and offset pair, eliminating the need for complex indirection tables while maintaining the memory efficiency benefits of paged storage.

## Lazy Allocation and Dynamic Growth

### On-Demand Initialization

Allocation follows a lazy pattern to conserve memory. The `key_cache` and `value_cache` tensors initialize as `None`, with actual memory allocation triggered only upon the first write operation. The `_load_contiguous_state` method performs the initial tensor creation based on the configured `block_size` and `num_blocks` parameters.

### Capacity Expansion

When requests exceed current capacity, the `_grow_to_capacity` method expands the cache dynamically, contingent on the `MTPLX_DYNAMIC_PAGED_KV` environment flag. Growth respects the serving context window defined by `MTPLX_CONTEXT_WINDOW_TOKENS`, ensuring the cache never allocates more blocks than addressable by the model's maximum sequence length.

## Quantization Support: TurboQuant and KV-Quant

### TurboQuant Fast Path

When `turboquant_config` is enabled, the cache writes directly into packed representations stored in `_turboquant_v_centroids` and related buffers. This configuration activates the `paged_attention_v2_online` kernel (defined in [`vllm_metal/metal/__init__.py`](https://github.com/youssofal/MTPLX/blob/main/vllm_metal/metal/__init__.py) at lines 102–115), providing accelerated quantized attention without intermediate dequantization steps.

### KV-Quant Mirrored Storage

KV-Quant mode (`kv_quant_config`) maintains dual storage: a mirrored **bfloat16** representation in `_dequant_memo` alongside the packed-quant bank in `_quant_bank`. The inference path latches on the first attention call for each request, guaranteeing consistent quantization math throughout the generation phase.

## Metal Kernel Integration and Fallback

The native Metal operations loader (`_load_vllm_metal_ops_optional`) attempts to JIT-compile optimized kernels via [`vllm_metal/metal/build.py`](https://github.com/youssofal/MTPLX/blob/main/vllm_metal/metal/build.py). If loading fails, the system falls back to an in-tree implementation using standard MLX vector or fast-SDPA kernels. The fallback preserves the same `[num_blocks, block_size, num_kv_heads, head_dim]` physical layout while sacrificing Metal-specific performance optimizations. A one-time warning indicates when fallback mode activates.

## Usage in Model Forward Pass

Integration occurs through the `cache` argument in attention layer calls. The `install_vllm_metal_paged_attention_kv_cache` function (tested in [`tests/test_turboquant_fallback.py`](https://github.com/youssofal/MTPLX/blob/main/tests/test_turboquant_fallback.py)) replaces the default `KVCache` with the paged implementation when the `--paged-kv-quantization` flag is present. This seamless integration works with MTPLX's graphbank, runtime scheduler, and snapshot systems.

```python
from mtplx.cache_state import VllmMetalPagedKVCache

# Initialize with default block size (16) and num_blocks (1024)

paged_cache = VllmMetalPagedKVCache()

# Enable quantization paths

turbo_cfg = {"k_quant": "q8_0", "v_quant": "q3_0"}
kv_cfg = {"dtype": "bf16"}
paged_cache = VllmMetalPagedKVCache(
    turboquant_config=turbo_cfg,
    kv_quant_config=kv_cfg
)

# Inject into model forward pass

outputs = model(input_ids, attention_mask=mask, cache=paged_cache)

# Retrieve performance statistics

stats = paged_cache.paged_stats()
print(f"Allocated blocks: {stats['allocated_blocks']}")
print(f"Paged attention calls: {stats['paged_attention_calls']}")

```

## Performance Monitoring and Statistics

The cache tracks comprehensive metrics through counters including `paged_attention_calls`, `turboquant_attention_calls`, and `kv_quant_attention_calls`. Additional per-phase bailout reasons populate `paged_attention_bailouts_by_phase_reason`. The `paged_stats()` method exposes these metrics for profiling tools, as utilized in [`tests/test_context_degradation_profiles.py`](https://github.com/youssofal/MTPLX/blob/main/tests/test_context_degradation_profiles.py) for performance validation.

## Summary

- MTPLX implements a **paged KV cache** compatible with vLLM-Metal kernels using a 4-D tensor shape `[num_blocks, block_size, num_kv_heads, head_dim]`.
- **Lazy allocation** via `_load_contiguous_state` and dynamic growth via `_grow_to_capacity` optimize memory usage while respecting `MTPLX_CONTEXT_WINDOW_TOKENS`.
- **TurboQuant** and **KV-Quant** provide optimized quantization paths, with the former using `paged_attention_v2_online` kernels and the latter maintaining mirrored bfloat16 storage.
- A **fallback mechanism** ensures operation continues using MLX kernels when Metal operations are unavailable.
- Comprehensive statistics via `paged_stats()` enable detailed performance monitoring and debugging.

## Frequently Asked Questions

### What is the default configuration for block size and number of blocks?

The `VllmMetalPagedKVCache` class defaults to a `block_size` of 16 tokens per block and `num_blocks` of 1024 maximum blocks, resulting in a physical tensor shape of `[1024, 16, num_kv_heads, head_dim]`. These defaults balance memory efficiency with kernel performance characteristics.

### How does MTPLX handle quantization in the paged KV cache?

MTPLX supports two quantization modes: **TurboQuant**, which writes directly to packed centroids and uses the `paged_attention_v2_online` kernel for speed, and **KV-Quant**, which stores both packed quantized values in `_quant_bank` and mirrored bfloat16 values in `_dequant_memo` for accuracy. The system latches onto one path per request to ensure consistent computation.

### What happens if the Metal kernels fail to load at runtime?

If `_load_vllm_metal_ops_optional` fails to load native Metal operations from [`vllm_metal/metal/__init__.py`](https://github.com/youssofal/MTPLX/blob/main/vllm_metal/metal/__init__.py), MTPLX falls back to MLX-based vector or fast-SDPA kernels. This fallback maintains the same paged memory layout and correctness but without Metal-specific acceleration, issuing a one-time warning to indicate the degraded performance path.

### How can I monitor paged KV cache performance during inference?

Call the `paged_stats()` method on your cache instance to retrieve counters for `paged_attention_calls`, `turboquant_attention_calls`, and allocated block counts. These statistics help identify whether quantization paths or fallback mechanisms are activating during generation, as validated in [`tests/test_turboquant_fallback.py`](https://github.com/youssofal/MTPLX/blob/main/tests/test_turboquant_fallback.py).