In the previous lecture we assembled the Transformer. Now we zoom into its core operation. Self-attention is deceptively simple — three matrix multiplications and a softmax — but understanding its behaviour, costs and modern variants is essential for anyone building or deploying transformer models.
Self-attention as learned, dynamic routing#
For an input sequence $\mathbf{X} \in \mathbb{R}^{T \times d}$:
Contrast with other layers:
- A fully connected layer mixes features with fixed weights.
- A convolution mixes neighbouring positions with fixed weights.
- Self-attention mixes all positions with weights $\mathbf{A}$ that are computed from the input itself — the connectivity pattern changes for every sentence.
Row $t$ of $\mathbf{A}$ is a probability distribution saying how much token $t$ reads from each other token.
Why separate queries, keys and values?#
Using the raw embeddings for everything would make attention symmetric ($\mathbf{x}_i^\top\mathbf{x}_j = \mathbf{x}_j^\top\mathbf{x}_i$) and force each token to attend most to itself. Separate projections let a token ask for one kind of information (query) while advertising another (key) and providing a third (value). A verb's query might seek its subject; a noun's key might advertise "I am a plausible subject".
Why scale by $\sqrt{d_k}$?#
If query and key components are independent with zero mean and unit variance, $\mathbf{q}^\top\mathbf{k}$ has variance $d_k$. Large scores push the softmax into saturated regions where gradients vanish. Dividing by $\sqrt{d_k}$ restores unit variance.
What do attention heads learn?#
Analyses of trained models (e.g. Clark et al., 2019 on BERT; mechanistic studies of GPT-style models) found heads that:
- attend to the previous or next token;
- track syntactic relations (a verb attending to its direct object, determiners to their nouns);
- resolve coreference (pronouns to their antecedents);
- attend heavily to special tokens like [SEP] or the first token — often acting as a "no-op" or attention sink;
- form induction heads in decoder models — circuits that look for a previous occurrence of the current token and copy what followed it, a mechanism linked to in-context learning.
Many heads are redundant; studies have pruned a large fraction of heads with little loss.
Complexity#
For sequence length $T$ and dimension $d$:
| Operation | Time | Memory |
|---|---|---|
| Projections ($\mathbf{Q}, \mathbf{K}, \mathbf{V}$) | $O(Td^2)$ | $O(Td)$ |
| Scores $\mathbf{Q}\mathbf{K}^\top$ | $O(T^2d)$ | $O(T^2)$ per head |
| Weighted sum $\mathbf{A}\mathbf{V}$ | $O(T^2d)$ | — |
For short sequences the $Td^2$ projections dominate; for long sequences the quadratic $T^2$ term dominates. Doubling context length quadruples attention cost — the main obstacle to long-context models. FlashAttention computes exact attention in tiles without materialising the $T \times T$ matrix in slow GPU memory, greatly reducing memory and wall-clock time (next lectures cover efficient transformers).
Causal attention and the KV cache#
In decoder-only models, token $t$ attends only to positions $\le t$. During generation, we produce one token at a time. Recomputing keys and values for the whole prefix at each step would be wasteful, since they do not change. The KV cache stores each layer's keys and values for previous tokens; each new step computes $\mathbf{q}, \mathbf{k}, \mathbf{v}$ only for the new token and attends over the cache.
The cache grows linearly with context length, and its memory can dominate inference:
(factor 2 for keys and values, $L$ layers, $n_{\text{kv}}$ key/value heads).
Multi-query and grouped-query attention#
To shrink the KV cache:
- Multi-Query Attention (MQA) (Shazeer, 2019): all query heads share a single key and value head — the cache shrinks by a factor of $h$, speeding decoding, with some quality loss.
- Grouped-Query Attention (GQA) (Ainslie et al., 2023): query heads are divided into $g$ groups, each sharing one key/value head — a middle ground adopted by many modern LLMs.
- Multi-head Latent Attention (MLA): compresses keys and values into a low-dimensional latent vector that is cached, further reducing memory.
import torch
def kv_cache_gb(layers, context, kv_heads, head_dim, batch=1, bytes_per=2):
return 2 * layers * context * kv_heads * head_dim * batch * bytes_per / 1e9
# A hypothetical 32-layer model, head_dim 128, 32 query heads, 32K context, fp16
print("MHA (32 KV heads):", round(kv_cache_gb(32, 32_768, 32, 128), 1), "GB per sequence")
print("GQA (8 KV heads): ", round(kv_cache_gb(32, 32_768, 8, 128), 1), "GB per sequence")
print("MQA (1 KV head): ", round(kv_cache_gb(32, 32_768, 1, 128), 2), "GB per sequence")Cross-attention#
In encoder–decoder models (and in multimodal models that attend to image features), queries come from one sequence and keys/values from another. The same equation, a different source of $\mathbf{K}$ and $\mathbf{V}$.