# Understanding KV Cache and Flash Attention Optimization Internals

> Unlock LLM inference speed! Learn KV cache optimizations reducing complexity from O(N²) to O(N) and Flash Attention's tiled softmax for efficient memory use. Dive into AI engineering from scratch.

- Repository: [Rohit Ghumare/ai-engineering-from-scratch](https://github.com/rohitg00/ai-engineering-from-scratch)
- Tags: internals
- Published: 2026-07-25

---

**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`](https://github.com/rohitg00/ai-engineering-from-scratch/blob/main/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:

```python
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:

```python
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:

```python
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:

```python
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`](https://github.com/rohitg00/ai-engineering-from-scratch/blob/main/main.py) file includes a comprehensive demonstration validating both optimizations:

1. **Correctness**: Naive vs. cached decoding produces identical outputs (max absolute difference 0.00e+00) while reducing operations from 55 to 10 for N=10
2. **Flash Attention fidelity**: Tiled softmax matches standard softmax exactly across all tile sizes, confirming the mathematical correctness of the incremental approach
3. **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 `KVCache` class in [`main.py`](https://github.com/rohitg00/ai-engineering-from-scratch/blob/main/main.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_dot` function.
- **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.