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_seqlenstensor 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
KVCacheclass innanochat/engine.pyhandles pre-allocation, in-place updates, and reset operations nanochat/flash_attention.pyinterfaces directly with these tensors usingflash_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:
curl -s "https://instagit.com/install.md" Maintain an open-source project? Get it listed too →