Skip to content

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)
Output
(4, 3) (4, 3) (4, 3)

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.

scores = Q @ K.T
print(scores)
Output
[[-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).

scaled = scores / np.sqrt(d_k)

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))
Output
[[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%})")
Output
(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:

Attention(Q, K, V) = softmax( Q·Kᵀ / √d_k ) · V
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))
Output
True

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}")
Output
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)
Output
[[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
(5, 8) → (5, 8)

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_k to 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.