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

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 with this specific arrangement:

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:

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


# 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 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.

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 handles pre-allocation, in-place updates, and reset operations
  • 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) 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 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, where the cached tensors are passed to flash_attn_with_kvcache along with the cache_seqlens metadata.

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 →