LectionesPart VII
Position, memory and the cost of attention
Attention is quadratic and the cache does not fit. Five papers about that.
Nothing in this part changes what attention computes. Every one of these papers returns the same numbers a naive implementation would, or very nearly, and the contribution is entirely in where those numbers live and how often they move.
That is a genuinely different kind of research and it is worth reading a run of it together. The constraint being optimised is not FLOPs — it is memory bandwidth, cache size, and the number of bytes you must keep per token per layer for the whole of a conversation.
RoPE is the exception: it changes how position enters the model, and it is here because everything about long context depends on it.
The reading
- RoPE
RoFormer: Enhanced Transformer with Rotary Position Embedding
Su et al. · Neurocomputing · 2021
- Claim
- Rotating query and key vectors by an angle proportional to position makes attention depend on relative position, with no learned position embeddings at all.
- Why
- In nearly every open model you will run. The rotation is also what explains why extending context by rescaling frequencies works — you cannot reason about that trick without this paper.
- Read
- Section 3.2. Do the two-dimensional case first; the general case is that, repeated.
- FlashAttention
FlashAttention: Fast and Memory-Efficient Exact Attention with IO-Awareness
Dao, Fu, Ermon, Rudra & Re · NeurIPS · 2022
- Claim
- Tiling the attention computation so it stays in on-chip SRAM makes it two to four times faster and linear in memory, returning exactly the same result.
- Why
- The paper that made reading the hardware a respectable research method. Nothing about the mathematics changes; the entire gain is in avoiding round trips to high-bandwidth memory.
- Read
- Section 3.1 and the memory-hierarchy figure.
- MQA
Fast Transformer Decoding: One Write-Head is All You Need
Shazeer · arXiv · 2019
- Claim
- Sharing a single key-value head across all query heads collapses the decoding cache and makes generation an order of magnitude faster.
- Why
- Nine pages, one author, and the reason your inference server fits in memory. Read it to understand what the key-value cache actually is before anyone optimises it further.
- Read
- All of it.
- GQA
GQA: Training Generalized Multi-Query Transformer Models from Multi-Head Checkpoints
Ainslie et al. · EMNLP · 2023
- Claim
- Grouping query heads over a small number of key-value heads recovers most of multi-query speed without its quality loss, and existing checkpoints can be converted rather than retrained.
- Why
- The compromise that shipped. Llama, Mistral and most of what came after use it, and the conversion recipe is why adoption was immediate.
- Ring Attention
Ring Attention with Blockwise Transformers for Near-Infinite Context
Liu, Zaharia & Abbeel · ICLR 2024 · 2023
- Claim
- Distributing blocks across devices arranged in a ring, and overlapping communication with computation, makes context length scale with the number of devices.
- Why
- Read it directly after FlashAttention. It is the same tiling argument, moved from one GPU's memory hierarchy to a cluster's.