2. Self-attention from scratch¶
Intermediate · 14 min read
Self-attention answers one question for every token: "which other tokens should I take information from, and how much?" It is built from three ingredients you already know: dot products (similarity), softmax (turn scores into weights), and a weighted average.
2.1 The intuition: a soft dictionary lookup¶
A Python dict does a hard lookup: one key matches exactly, you get its value. Attention does a soft lookup: the query is compared with every key, and you get a blend of all values, weighted by how well each key matches.
| Role | Each token produces… | Think of it as |
|---|---|---|
Query q |
what I'm looking for | "I'm the pronoun it — I need a noun that can be tired" |
Key k |
what I contain, for matching | "I'm animal: a noun, living thing" |
Value v |
what I hand over if chosen | the information that gets copied into it |
2.2 Step by step on a tiny example¶
Four tokens, each already an embedding of size 4:
import numpy as np
np.set_printoptions(precision=2, suppress=True)
tokens = ["the", "cat", "sat", "down"]
X = np.array([
[0.1, 0.0, 0.2, 0.0], # the
[0.9, 0.8, 0.1, 0.3], # cat
[0.2, 0.9, 0.7, 0.1], # sat
[0.1, 0.3, 0.9, 0.4], # down
])
rng = np.random.default_rng(42)
d_model, d_k = 4, 3
W_q = rng.normal(size=(d_model, d_k)) # learned during training
W_k = rng.normal(size=(d_model, d_k))
W_v = rng.normal(size=(d_model, d_k))
Q, K, V = X @ W_q, X @ W_k, X @ W_v # every token gets a query, key and value (4 x 3 each)
print(Q.shape, K.shape, V.shape)
Step 1 — scores. Compare every query with every key using a dot product. One matrix multiply gives a 4 × 4 table: row i = how much token i is interested in each token.
[[-0.01 -0.35 -0.19 -0.13]
[-0.09 -4.31 -1.5 -1.24]
[-0.06 -3.79 -0.67 -0.71]
[-0.05 -1.1 -0.4 -0.38]]
Step 2 — scale. Divide by √d_k. Without this, dot products grow with the vector size, softmax becomes extremely peaked, and training stalls (tiny gradients).
Step 3 — softmax each row into weights that sum to 1:
def softmax(x, axis=-1):
x = x - x.max(axis=axis, keepdims=True) # numerical stability
e = np.exp(x)
return e / e.sum(axis=axis, keepdims=True)
weights = softmax(scaled)
print(weights)
print(weights.sum(axis=1))
[[0.27 0.22 0.25 0.25]
[0.49 0.04 0.22 0.25]
[0.4 0.05 0.28 0.27]
[0.31 0.17 0.26 0.26]]
[1. 1. 1. 1.]
Step 4 — weighted sum of values. Each token's output is a blend of all tokens' values:
out = weights @ V
print(out.shape)
for t, w in zip(tokens, weights):
print(f"{t:5} attends most to {tokens[w.argmax()]!r} ({w.max():.0%})")
(4, 3)
the attends most to 'the' (27%)
cat attends most to 'the' (49%)
sat attends most to 'the' (40%)
down attends most to 'the' (31%)
With random weights the pattern means nothing — in a trained model these weights are what let "it" pick out "animal". The whole thing is one formula:
def attention(Q, K, V, mask=None):
scores = Q @ K.swapaxes(-1, -2) / np.sqrt(Q.shape[-1])
if mask is not None:
scores = np.where(mask, scores, -np.inf) # blocked positions get weight 0 after softmax
return softmax(scores) @ V
print(np.allclose(attention(Q, K, V), out))
2.3 Why divide by √d_k?¶
for d in (4, 64, 1024):
q, k = rng.normal(size=(1000, d)), rng.normal(size=(1000, d))
dots = (q * k).sum(axis=1)
print(f"d={d:5} std of q·k = {dots.std():6.1f} after /√d = {(dots / np.sqrt(d)).std():.2f}")
d= 4 std of q·k = 2.1 after /√d = 1.05
d= 64 std of q·k = 8.2 after /√d = 1.03
d= 1024 std of q·k = 30.4 after /√d = 0.95
Raw dot products grow like √d. Scaling keeps them around 1 whatever the size, so softmax stays smooth and trainable.
2.4 The causal mask — no peeking at the future¶
A GPT-style model is trained to predict the next token. If token 2 could look at token 3, it would just copy the answer. The causal mask lets each token attend only to itself and earlier tokens:
n = len(tokens)
causal = np.tril(np.ones((n, n), dtype=bool)) # lower triangle = allowed
print(causal.astype(int))
masked_weights = softmax(np.where(causal, scaled, -np.inf))
print(masked_weights)
[[1 0 0 0]
[1 1 0 0]
[1 1 1 0]
[1 1 1 1]]
[[1. 0. 0. 0. ]
[0.92 0.08 0. 0. ]
[0.55 0.06 0.39 0. ]
[0.31 0.17 0.26 0.26]]
"the" can only see itself (weight 1.0); "down" sees everything. Encoder models like BERT skip this mask — every token sees the whole text, which is why they're good at understanding but can't generate.
Training is parallel, generation is not
With the mask, one forward pass trains on every position at once: predict token 2 from 1, token 3 from 1–2, … all in parallel. Generation, however, must still produce one token at a time — see How LLMs generate text.
2.5 Multi-head attention¶
One attention pattern can't capture everything at once — grammar, coreference, topic, position. Multi-head attention runs several attentions in parallel, each with its own smaller W_q, W_k, W_v, then concatenates the results and mixes them with an output matrix W_o.
def multi_head_attention(X, n_heads, rng, causal=True):
n, d_model = X.shape
d_head = d_model // n_heads
W_q, W_k, W_v, W_o = (rng.normal(scale=d_model ** -0.5, size=(d_model, d_model)) for _ in range(4))
def split(M): # (n, d_model) → (heads, n, d_head)
return M.reshape(n, n_heads, d_head).transpose(1, 0, 2)
Q, K, V = split(X @ W_q), split(X @ W_k), split(X @ W_v)
mask = np.tril(np.ones((n, n), dtype=bool)) if causal else None
heads = attention(Q, K, V, mask) # all heads at once: (heads, n, d_head)
concat = heads.transpose(1, 0, 2).reshape(n, d_model)
return concat @ W_o
X8 = rng.normal(size=(5, 8)) # 5 tokens, d_model = 8
Y = multi_head_attention(X8, n_heads=2, rng=rng)
print(X8.shape, "→", Y.shape)
Output shape = input shape. That is what lets blocks be stacked. Real sizes for reference:
| Model | d_model | Heads | d_head | Layers |
|---|---|---|---|---|
| BERT-base | 768 | 12 | 64 | 12 |
| GPT-2 small | 768 | 12 | 64 | 12 |
| Llama 3 8B | 4096 | 32 | 128 | 32 |
| Llama 3 70B | 8192 | 64 | 128 | 80 |
Modern LLMs often use grouped-query attention (GQA): many query heads share a few key/value heads. Same quality, much smaller KV cache — that's what makes long contexts affordable at serving time.
2.6 Cross-attention¶
In encoder–decoder models (T5, translation) the decoder also runs cross-attention: queries come from
the decoder's tokens, keys and values from the encoder's output — "while writing the Hindi word, look at
the relevant English words". Same attention() function, different inputs.
Interview questions¶
Explain self-attention in one minute.
Each token is projected into a query, key and value. The query of each token is dot-producted with the keys of all tokens to get relevance scores, scaled by √d_k, softmaxed into weights, and used to take a weighted average of the values. The result is a new representation of each token that mixes in information from the tokens most relevant to it. It's all matrix multiplies, so it runs in parallel.
Why scale by √d_k?
The variance of a dot product of random vectors grows with dimension. Large scores push softmax into a near one-hot regime where gradients vanish. Dividing by √d_k keeps scores at unit scale.
Why multiple heads?
Each head can learn a different relationship (syntax, coreference, position, topic) in its own subspace. A single head averages them into one pattern. Heads cost about the same as one big head because each works in d_model / n_heads dimensions.
What does the causal mask do and which models use it?
It sets attention scores to future positions to −∞, so each token attends only to itself and earlier tokens. Decoder-only (GPT-style) models use it so training on next-token prediction can't cheat and so generation works left to right.
Practice¶
- Change
d_kto 16 and remove the √d_k scaling — how peaked do the softmax rows become? - Write a padding mask: for a batch where the last two tokens are
<pad>, make sure no token attends to them.
Next: The transformer block — positions, normalisation, feed-forward, and stacking.