Part 2 · 1 chapters · ~8 min
Attention and Transformers
Scaled dot-product attention with queries, keys and values, the causal mask, multi-head attention, grouped-query attention, the transformer block (attention, feed-forward, residual connections, normalisation), stacking layers, why attention is quadratic in context length, and mixture-of-experts layers.
3
Attention in a few lines
code
import torch, math
def attention(x, Wq, Wk, Wv): # x: (T, d) tokens of one sequence
Q, K, V = x @ Wq, x @ Wk, x @ Wv
scores = Q @ K.T / math.sqrt(Q.shape[-1]) # (T, T): every token against every token
mask = torch.triu(torch.ones_like(scores), diagonal=1).bool()
scores = scores.masked_fill(mask, float("-inf")) # causal: no attention to the future
return torch.softmax(scores, dim=-1) @ V # weighted sum of values
# one transformer block (pre-norm, as in Llama)
h = x + attention(norm1(x)) # residual connection
out = h + feed_forward(norm2(h)) # an MLP applied to each token independently| idea | why |
|---|---|
| T × T score matrix | compute grows quadratically with context; FlashAttention computes it in tiles without storing it |
| grouped-query attention (GQA) | many query heads share fewer key/value heads: Llama 3.1 8B has 32 query heads and 8 KV heads, shrinking the KV cache 4× |
| residual connections and normalisation | let gradients flow through dozens of layers (ML maths P5) |
| mixture of experts | route each token to a few of many feed-forward experts: more parameters, similar compute per token |
ONE ATTENTION HEAD
each token gathers information from earlier tokens
swipe the figure sideways, or tap expand for full screen
1/4
queries, keys, values
Each token is projected three ways: a query (what am I looking for?), a key (what do I contain?) and a value (what do I pass on?).
three projectionslearned matrices