Understanding KV Cache and Flash Attention Optimization Internals
KV caching reduces LLM inference complexity from O(N²) to O(N) by reusing previously computed key/value vectors, while Flash Attention employs tiled softmax with incremental normalization to slash memory usage without numerical approximation.
Modern large language models (LLMs) rely on two critical optimizations to handle long contexts efficiently: the KV cache for computational reuse and Flash Attention-style tiling for memory efficiency. This article examines the exact implementations found in the rohitg00/ai-engineering-from-scratch repository, demonstrating how these techniques transform the economics of transformer inference through direct source code analysis.
The KV Cache: Eliminating Quadratic Recomputation
During autoregressive generation, a transformer decoder must compute attention between the current token's query and all previous tokens. Without optimization, this creates a quadratic explosion in compute and memory bandwidth.
The Naive O(N²) Bottleneck
In a standard implementation, every generation step recomputes attention over the entire sequence history. For a sequence of length N, this requires computing an N×N attention matrix, resulting in O(N²) operations per step. As sequences grow to 100k+ tokens, this becomes prohibitively expensive.
The repository demonstrates this inefficiency in phases/07-transformers-deep-dive/12-kv-cache-flash-attention/code/main.py, contrasting it against the optimized cached approach.
Minimal KV Cache Implementation
The repository implements a pure-Python KV cache using simple list accumulation. The KVCache class maintains running lists of key and value vectors:
class KVCache:
def __init__(self):
self.K = [] # list of past key vectors
self.V = [] # list of past value vectors
def append(self, k, v):
self.K.append(k)
self.V.append(v)
The cached decoder, decode_cached, leverages this structure to reduce per-step complexity to O(N). Instead of recomputing over the full history, it appends new (k, v) pairs to the cache and attends only against the accumulated storage:
def decode_cached(all_K, all_V, all_queries):
cache = KVCache()
outputs = []
ops = 0
for q, k, v in zip(all_queries, all_K, all_V):
cache.append(k, v) # grow KV-cache
out = attention_full(q, cache.K, cache.V)
ops += len(cache) # O(N) ops per step
outputs.append(out)
return outputs, ops
This approach cuts the total operation count from 55 to 10 for a sequence of 10 tokens in the repository's demo, scaling linearly rather than quadratically.
Memory Footprint and Configuration
While the KV cache saves compute, it consumes significant memory. The repository provides kv_cache_bytes to calculate exact storage requirements:
def kv_cache_bytes(N, n_layers, n_heads_kv, d_head, dtype=2):
"""Total KV cache bytes. dtype=2 for fp16/bf16, 1 for int8, 4 for fp32."""
return 2 * N * n_layers * n_heads_kv * d_head * dtype
For Llama-3-70B at 128k context length, this function reveals the cache alone exceeds 10 GB in fp16 precision. This explains why modern architectures employ Grouped-Query Attention (GQA) and other compression techniques to make long-context inference economically viable.
Flash Attention: Tiled Softmax Without Approximation
Standard attention materializes the full softmax matrix in memory, creating bandwidth bottlenecks. Flash Attention solves this by processing the sequence in tiles while maintaining exact numerical equivalence through incremental statistics.
The Running-Max Numerical Stability Trick
The core challenge with tiled softmax is maintaining numerical stability across partial computations. The repository's tiled_softmax_dot function implements the running-maximum technique to prevent overflow/underflow:
def tiled_softmax_dot(q, Ks, Vs, tile=4):
"""Flash-attention-style incremental softmax(QKᵀ)V with tile size `tile`."""
d_head = len(Vs[0])
scale = 1.0 / math.sqrt(len(q))
m = float("-inf") # running max for numerical stability
s = 0.0 # running sum of exponentials
out = [0.0] * d_head
for start in range(0, len(Ks), tile):
k_block = Ks[start:start + tile]
v_block = Vs[start:start + tile]
scores = [dot(q, k) * scale for k in k_block]
new_m = max(m, *scores) # updated max
exp_old = math.exp(m - new_m) if m != float("-inf") else 0.0
exp_new = [math.exp(sc - new_m) for sc in scores]
s = s * exp_old + sum(exp_new) # updated denominator
for j in range(d_head):
out[j] = out[j] * exp_old + sum(e * v[j] for e, v in zip(exp_new, v_block))
m = new_m
return [o / s for o in out] # final normalized output
The algorithm tracks m (running maximum) and s (running sum) across tiles. When processing a new tile, it rescales previous accumulations using exp(m - new_m), ensuring the final result matches a full-pass softmax exactly.
Tile-Based Memory Efficiency
By processing keys and values in blocks (default tile=4), the implementation bounds memory usage to O(tile · d_head) rather than O(N²). The repository demonstrates this with multiple tile sizes (1, 2, 4, 8, 32), showing zero numerical deviation from standard softmax in all cases.
This tiling strategy eliminates the need to materialize large intermediate attention matrices, drastically reducing memory bandwidth requirements during the attention computation.
Verification and Benchmarking
The main.py file includes a comprehensive demonstration validating both optimizations:
- Correctness: Naive vs. cached decoding produces identical outputs (max absolute difference 0.00e+00) while reducing operations from 55 to 10 for N=10
- Flash Attention fidelity: Tiled softmax matches standard softmax exactly across all tile sizes, confirming the mathematical correctness of the incremental approach
- Memory projections: The KV cache size table quantifies storage requirements across different model scales, from Llama-3.2-3B (0.13 GB at 128k context) to Llama-3-70B (10.24 GB)
Running python3 main.py executes these benchmarks, providing empirical proof that these optimizations preserve model behavior while delivering massive efficiency gains.
Summary
- KV caching transforms transformer inference from O(N²) to O(N) complexity by storing and reusing key/value vectors across generation steps, implemented in the repository via the
KVCacheclass inmain.py. - Flash Attention uses tiled softmax with running maximum and sum statistics to process attention in constant-memory blocks without numerical approximation, as demonstrated in the
tiled_softmax_dotfunction. - Memory trade-offs are significant: at 128k context, a 70B parameter model requires over 10 GB for KV storage alone, necessitating architectural innovations like GQA.
- Exact correctness is maintained throughout both optimizations—the repository verifies zero numerical deviation from naive implementations.
Frequently Asked Questions
What is the primary benefit of KV caching in transformer inference?
KV caching eliminates redundant computation during autoregressive generation. By storing the key and value vectors for each generated token, the model avoids recomputing attention over the entire sequence history at every step. This reduces per-step complexity from quadratic O(N²) to linear O(N), enabling practical inference for long contexts exceeding 100,000 tokens.
How does Flash Attention maintain numerical accuracy while tiling?
Flash Attention uses a running-maximum algorithm to track the global softmax maximum across tiles incrementally. When processing each tile, it rescales previous accumulations by exp(old_max - new_max), ensuring the exponential calculations remain numerically stable and the final output matches a full-matrix softmax exactly. The repository's implementation demonstrates zero approximation error across all tested tile sizes.
Why does KV cache memory grow linearly with sequence length?
The cache stores two vectors (key and value) for every token in the context, every layer, and every KV head. As shown in the repository's kv_cache_bytes function, total storage equals 2 × N × n_layers × n_heads_kv × d_head × dtype_size. For long contexts (N=131072), this linear scaling creates multi-gigabyte memory requirements, which is why modern LLMs employ grouped-query attention (GQA) to reduce n_heads_kv relative to the total head count.
Can Flash Attention work without KV caching?
Yes, Flash Attention and KV caching are orthogonal optimizations. Flash Attention reduces memory bandwidth during the attention computation itself by avoiding materialization of the full attention matrix, while KV caching reduces redundant computation across generation steps. Production systems typically combine both: Flash Attention efficiently computes attention within each step, while the KV cache avoids recomputing keys and values for prior tokens.
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 →