Tokenwerk · The LLM textbook

Chapter 26 · V · Modern language models · 5 minutes

GQA, FlashAttention and the KV cache

Three different levers: fewer key/value heads, less memory traffic and fewer repeated computations.

Not every speedup changes the architecture

GQA changes the number of shared key/value representations. FlashAttention computes a dense attention operation with a more memory-efficient algorithm. A KV cache stores already computed intermediate values for later generation steps. These techniques solve related but different bottlenecks.

So when someone says "faster attention", first ask: less theoretical compute, fewer stored intermediate values, less memory traffic or fewer redundant computations? A benchmark on one piece of hardware does not automatically answer all these questions.

MHA, MQA and GQA

In multi-head attention, every query head has its own key and value head. In multi-query attention, all query heads share a single key/value pair. Grouped-query attention lies in between: groups of query heads each share one key head and one value head.

If Hq=8H_q=8 and Hkv=2H_{kv}=2, every four query heads use the same K/V representations. The Q projection still produces 8d8d components, K and V only 2d2d each. This reduces the K/V projection weights and especially the cache requirement. Original GQA paper.

GQA in easy-to-read code

python
repeat = cfg.heads // cfg.kv_heads
k_for_attention = k.repeat_interleave(repeat, dim=1)
v_for_attention = v.repeat_interleave(repeat, dim=1)
out = F.scaled_dot_product_attention(q, k_for_attention, v_for_attention, is_causal=True)

This explicit repetition is easy to follow, but it materializes additional tensors. An optimized GQA implementation can avoid that. Our course project saves projection parameters, but doesn't claim that this naive repetition already achieves the optimal kernel memory footprint.

Prefill and decode

During prefill, the model reads the complete prompt. Many query positions are processed in parallel. During decode, one new token is added at a time. Without a cache, you recompute all K/V representations of the previous context again and again.

With a cache, each layer stores the keys and values of the previous positions. The new token only produces its new Q/K/V vectors. Its query reads the stored keys and values plus the new entry. The old hidden states don't need to be recomputed, as long as the context and model rules continue unchanged.

Computing the cache requirement

For layer count L, batch B, stored length T, K/V head count Hkv, head dimension d and bytes per element e, you get approximately:

MKV=2LBTHkvde.M_{KV}=2LBTH_{kv}de.

M_KV is the required memory in bytes, that is, units of storage. L counts the layers; B the examples processed at the same time; T the stored positions; H_kv the distinct key/value heads; d the components per head; e the bytes per stored number. All factors are multiplied. The two stands for the two separate fields K and V. With L=6, B=1, T=1024, Hkv=2, d=64 and two bytes per element, you get exactly 2 × 6 × 1 × 1024 × 2 × 64 × 2 = 3,145,728 bytes, that is, three MiB. One MiB (mebibyte) is 1024 × 1024 = 1,048,576 bytes. On top of that come weights, temporary tensors and implementation metadata.

Batch size and context length enter this cache formula linearly. That explains why a model with small quantized weights can still need considerable memory for many parallel long conversations.

The important mask bug with a cache

In a normal complete forward pass, query and key lengths are equal. During decode with a cache, the query length can be one and the key length a thousand. A blind is_causal=True with a triangular mask aligned to the top left can then allow the wrong keys. The new query should see all previous keys, not just the first.

For a single new token with a cache that lies completely in the past, it can be correct to apply no additional causal restriction. For several new tokens, the mask needs a correct position offset. Likewise, RoPE must use the actual new position. Our project deliberately doesn't have a persistent KV cache yet; adding a cache is an advanced task, with a comparison test against the uncached decoder.

FlashAttention

Naive attention stores the full score or weight matrix. FlashAttention processes blocks and uses a numerically stable online softmax to compute the dense result without materializing the entire matrix in the large GPU memory. The dense pairwise work basically remains; the method is not automatically linear sparse attention. FlashAttention.

PyTorch SDPA can choose a suitable kernel depending on device, dtype, masks and shape. Calling this function doesn't prove that FlashAttention was actually chosen on your Mac or on every CUDA configuration. Measure and check the backend instead of reading a function name as a performance guarantee.

Two more widely used methods: multi-head latent attention (MLA) stores a compressed latent vector per token instead of K and V, shrinking the cache more than GQA does (DeepSeek-V2). Speculative decoding lets a small draft model (or extra prediction heads) propose several tokens, which the large model checks in a single forward pass. With a correct acceptance rule, the output distribution stays unchanged (Leviathan et al., 2022).

Your extension plan

First implement a cache for the classic decoder without context shifting. Compare the uncached and cached logits at every position. Then extend RoPE with an offset and repeat the comparison. Only after that, examine GQA caches and long contexts. A cache that is faster but predicts different tokens may simply be wrong.