# KV Cache Layout Used by NanoChat's Engine for Flash Attention 3 (FA3)

> Discover the KV cache layout NanoChat uses for Flash Attention 3 (FA3). Understand the `(n_layers, B, T, H, D)` tensor structure for efficient transformer inference.

- Repository: [Andrej/nanochat](https://github.com/karpathy/nanochat)
- Tags: internals
- Published: 2026-03-10

---

**NanoChat stores key-value caches in a `(n_layers, B, T, H, D)` tensor layout to match Flash Attention 3's `flash_attn_with_kvcache` API requirements.**

The `karpathy/nanochat` inference engine implements a specialized **KV cache layout** designed specifically for Flash Attention 3 (FA3) compatibility. Unlike traditional layouts used in Flash Attention 2, this structure reorders the sequence and head dimensions to enable efficient in-place updates during autoregressive generation.

## Understanding the FA3 KV Cache Tensor Structure

Flash Attention 3 expects a specific memory arrangement that differs from previous versions. The engine pre-allocates both key (`k_cache`) and value (`v_cache`) tensors using a five-dimensional shape that interleaves layer, batch, sequence, head, and dimension data.

### Dimension Order and Meaning

The cache layout follows the exact ordering required by FA3's kernel expectations:

| Dimension | Symbol | Description |
|-----------|--------|-------------|
| **n_layers** | — | Number of transformer layers (leading dimension) |
| **B** | Batch | Number of concurrent sequences being generated |
| **T** | seq_len | Maximum context length the cache supports |
| **H** | num_heads | Number of attention heads per layer |
| **D** | head_dim | Dimensionality of each individual head |

Both tensors are instantiated in [`nanochat/engine.py`](https://github.com/karpathy/nanochat/blob/main/nanochat/engine.py) with this specific arrangement:

```python
self.k_cache = torch.zeros(
    num_layers, batch_size, seq_len, num_heads, head_dim,
    device=device, dtype=dtype
)
self.v_cache = torch.zeros(
    num_layers, batch_size, seq_len, num_heads, head_dim,
    device=device, dtype=dtype
)

```

### Sequence Length Tracking with `cache_seqlens`

FA3 requires explicit tracking of the current sequence position for each batch element using a separate `int32` tensor. This allows the kernels to perform in-place appends without full cache copies:

```python
self.cache_seqlens = torch.zeros(batch_size, dtype=torch.int32, device=device)

```

The `KVCache` class updates this tensor in-place during generation to indicate how many tokens are currently valid in the cache slots.

## Key Differences from Flash Attention 2

The **KV cache layout** changes significantly between FA2 and FA3. Understanding this distinction is critical when migrating models or debugging shape mismatches.

- **FA2 Layout**: Expects caches shaped `(B, H, T, D)` — batch, heads, sequence, then dimension
- **FA3 Layout**: Requires caches shaped `(B, T, H, D)` — batch, sequence, heads, then dimension

NanoChat's implementation wraps this FA3 structure with an additional leading dimension for layers, resulting in the final `(n_layers, B, T, H, D)` shape. This permutation eliminates expensive transposition operations inside the attention kernel and enables direct memory appending during token generation.

## Implementing the Cache in NanoChat's Engine

The `KVCache` class defined in [`nanochat/engine.py`](https://github.com/karpathy/nanochat/blob/main/nanochat/engine.py) (lines 81-102) abstracts these tensors and provides helper methods for manipulation. The constructor handles device placement and dtype selection, ensuring CUDA tensor cores can operate on the cache efficiently.

### Core Implementation Details

The source code shows the cache is initialized once and reused across generation steps:

```python

# From nanochat/engine.py lines 81-102

class KVCache:
    def __init__(self, batch_size, num_heads, seq_len, head_dim, 
                 num_layers, device, dtype):
        self.num_layers = num_layers
        self.batch_size = batch_size
        self.seq_len = seq_len
        # Pre-allocate continuous GPU memory for all layers

        self.k_cache = torch.zeros(
            num_layers, batch_size, seq_len, num_heads, head_dim,
            device=device, dtype=dtype
        )
        self.v_cache = torch.zeros(
            num_layers, batch_size, seq_len, num_heads, head_dim,
            device=device, dtype=dtype
        )
        self.cache_seqlens = torch.zeros(
            batch_size, dtype=torch.int32, device=device
        )

```

The [`flash_attention.py`](https://github.com/karpathy/nanochat/blob/main/flash_attention.py) file contains the wrapper that passes these tensors directly to `flash_attn_with_kvcache`, ensuring zero-copy memory access between the cache and the attention computation.

## Practical Usage Example

Instantiating and manipulating the KV cache requires specifying the model's hyperparameters upfront. The cache supports both prefill operations and incremental decoding phases.

```python
import torch
from nanochat.engine import KVCache

# Model hyper-parameters (example)

batch_size   = 4
num_heads    = 12
seq_len      = 2048          # maximum context length

head_dim     = 64
num_layers   = 24
device       = torch.device('cuda')
dtype        = torch.bfloat16

# Create a KV cache ready for Flash Attention 3

kv_cache = KVCache(
    batch_size=batch_size,
    num_heads=num_heads,
    seq_len=seq_len,
    head_dim=head_dim,
    num_layers=num_layers,
    device=device,
    dtype=dtype,
)

# Access layer-specific caches (used by the model's forward pass)

layer_idx = 0
k_layer, v_layer = kv_cache.get_layer_cache(layer_idx)

# Advance the position after generating `n` new tokens

kv_cache.advance(num_tokens=1)

# Reset the cache (e.g., when starting a new prompt)

kv_cache.reset()

```

The `advance()` method increments `cache_seqlens` appropriately, while `get_layer_cache()` returns views into the pre-allocated tensors for specific transformer layers without copying memory.

## Summary

- NanoChat's engine stores KV caches in a **(n_layers, B, T, H, D)** layout specifically for FA3 compatibility
- The `cache_seqlens` tensor tracks valid sequence lengths as **int32** for each batch element
- This layout differs from FA2 by placing the sequence dimension **before** the head dimension
- The `KVCache` class in [`nanochat/engine.py`](https://github.com/karpathy/nanochat/blob/main/nanochat/engine.py) handles pre-allocation, in-place updates, and reset operations
- [`nanochat/flash_attention.py`](https://github.com/karpathy/nanochat/blob/main/nanochat/flash_attention.py) interfaces directly with these tensors using `flash_attn_with_kvcache`

## Frequently Asked Questions

### What is the exact tensor shape of NanoChat's KV cache for FA3?

The cache uses a five-dimensional tensor with shape **(n_layers, B, T, H, D)**. This represents (number of layers, batch size, maximum sequence length, number of attention heads, head dimension). Both the key and value caches share this identical structure, with the layer dimension serving as the outermost index to facilitate easy layer-wise access during the forward pass.

### How does the FA3 KV cache layout differ from FA2 in NanoChat?

Flash Attention 2 expects caches shaped **(B, H, T, D)**, placing heads before the sequence dimension. FA3 requires **(B, T, H, D)**, moving the sequence dimension earlier to support efficient in-place appending. NanoChat wraps this in an additional layer dimension, resulting in the final **(n_layers, B, T, H, D)** structure that matches FA3's memory access patterns.

### Why does NanoChat use a separate `cache_seqlens` tensor?

The `cache_seqlens` tensor (defined in [`nanochat/engine.py`](https://github.com/karpathy/nanochat/blob/main/nanochat/engine.py)) stores the current valid length for each batch element as **int32** values. FA3's `flash_attn_with_kvcache` function requires this metadata to distinguish between pre-filled cache positions and empty slots, enabling variable-length sequences within the same pre-allocated buffer without padding or reallocation overhead.

### Where is the KV cache logic implemented in the NanoChat repository?

The primary implementation resides in [`nanochat/engine.py`](https://github.com/karpathy/nanochat/blob/main/nanochat/engine.py) at lines 81-102, which defines the `KVCache` class including initialization, `reset()`, `advance()`, and `get_layer_cache()` methods. The integration with Flash Attention 3 kernels occurs in [`nanochat/flash_attention.py`](https://github.com/karpathy/nanochat/blob/main/nanochat/flash_attention.py), where the cached tensors are passed to `flash_attn_with_kvcache` along with the `cache_seqlens` metadata.