“Walk me through self-attention” opens almost every interview that touches modern deep learning. Interviewers are not checking whether you can recite the formula. They want to hear where queries, keys and values come from, why there is a square root in the denominator, what a decoder adds, and what attention costs in compute and memory. Computing a small example by hand and then estimating a real model’s KV cache shows you understand the mechanism.
Before you start
You need matrix multiplication, the softmax and a rough idea of gradient descent. The examples are pure Python 3.12+ using only math and random, so every number in the comments can be reproduced. The few PyTorch 2.x lines are illustrative. This article starts where tokens have already become vectors; for tokenization see LLM tokens, temperature and top-p sampling.
The short answer
Three learned matrices project each token into a query, a key and a value. Scores are query-key dot products divided by sqrt(d_k) to keep their variance near 1, and a softmax turns each row into weights that mix the values: softmax(QK^T / sqrt(d_k)) V. Decoders add a causal mask so token i sees only tokens 0 to i, multi-head attention runs several smaller attentions in parallel, and each block wraps attention and an MLP in residual connections with layer normalization.
How it works
Start with an input X of shape (n, d_model): one row per token, already combined with position information. Learned matrices W_q, W_k and W_v project it into Q = X W_q, K = X W_k and V = X W_v. A query describes what a token is looking for, a key what it offers to be matched on, and a value the information it hands over. All three come from the same sequence, hence self-attention; in cross-attention the queries come from a different sequence.
Entry (i, j) of the (n, n) score matrix QK^T measures how strongly token i attends to token j. These helpers compute everything:
import math
def matmul(A, B):
return [[sum(a * b for a, b in zip(row, col)) for col in zip(*B)] for row in A]
def transpose(A):
return [list(col) for col in zip(*A)]
def softmax(row):
m = max(row) # subtract the max for numerical stability
exps = [math.exp(x - m) for x in row]
total = sum(exps)
return [e / total for e in exps]
def attention(Q, K, V, causal=False):
d_k = len(K[0])
scores = matmul(Q, transpose(K)) # (n, n): query i against key j
scores = [[s / math.sqrt(d_k) for s in row] for row in scores]
if causal: # token i may only look at tokens 0..i
scores = [[s if j <= i else -math.inf for j, s in enumerate(row)]
for i, row in enumerate(scores)]
weights = [softmax(row) for row in scores]
return matmul(weights, V), weights
def show(M, nd=3):
for row in M:
print([round(x, nd) for x in row])Multi-head attention splits d_model into h heads of size d_model / h. Each head attends over its own slice, the outputs are concatenated back to width d_model, and an output projection W_o mixes them. Heads can specialise, for example one tracking the previous token. The split adds no parameters, because the four projections stay d_model by d_model:
def multi_head(Q, K, V, h, causal=False):
head_dim = len(Q[0]) // h
heads = []
for i in range(h): # head i owns columns i*head_dim .. (i+1)*head_dim
cols = slice(i * head_dim, (i + 1) * head_dim)
q = [row[cols] for row in Q]
k = [row[cols] for row in K]
v = [row[cols] for row in V]
heads.append(attention(q, k, v, causal)[0])
return [sum((hd[t] for hd in heads), []) for t in range(len(Q))] # concat per token
def mha_params(d_model, h):
head_dim = d_model // h
qkv = 3 * d_model * (h * head_dim) # W_q, W_k, W_v
out = (h * head_dim) * d_model # W_o
return qkv + out
print(mha_params(512, 1), mha_params(512, 8)) # 1048576 1048576Positional information must be added, because attention treats its input as a set: shuffle the tokens and the outputs shuffle the same way. The original transformer added fixed sinusoids, BERT and GPT-2 learn a position table, and most current LLMs use rotary embeddings (RoPE), which rotate queries and keys by position-dependent angles so their dot product depends on relative distance.
def sinusoidal(pos, d_model):
pe = []
for i in range(0, d_model, 2):
angle = pos / 10000 ** (i / d_model)
pe += [math.sin(angle), math.cos(angle)]
return pe
for pos in range(3):
print(pos, [round(x, 3) for x in sinusoidal(pos, 4)])
# 0 [0.0, 1.0, 0.0, 1.0]
# 1 [0.841, 0.54, 0.01, 1.0]
# 2 [0.909, -0.416, 0.02, 1.0]A transformer block is attention followed by a position-wise MLP (classically four times wider than d_model), each wrapped in a residual connection and a LayerNorm. The 2017 paper normalised after the residual add (post-LN). GPT-2 and most later models use pre-LN, x = x + attn(norm(x)) then x = x + mlp(norm(x)), which keeps a clean residual path and trains more stably in deep stacks. Many LLMs swap LayerNorm for the cheaper RMSNorm.
Step-by-step walkthrough
Step 1: Project three tokens into Q, K and V
Use three tokens, d_model = 4 and small integer weights so the arithmetic is checkable. Real weights are learned floats.
X = [[1, 0, 1, 0], # "the"
[0, 2, 0, 2], # "cat"
[1, 1, 1, 1]] # "sat"
W_q = [[1, 0, 0, 0], [0, 1, 0, 0], [1, 0, 0, 0], [0, 0, 0, 1]]
W_k = [[0, 1, 0, 0], [1, 0, 0, 0], [0, 0, 1, 0], [0, 1, 0, 1]]
W_v = [[1, 0, 0, 0], [0, 0, 1, 0], [0, 1, 0, 0], [0, 0, 0, 1]]
Q, K, V = matmul(X, W_q), matmul(X, W_k), matmul(X, W_v)
print(Q) # [[2, 0, 0, 0], [0, 2, 0, 2], [2, 1, 0, 1]]
print(K) # [[0, 1, 1, 0], [2, 2, 0, 2], [1, 2, 1, 1]]
print(V) # [[1, 1, 0, 0], [0, 0, 2, 2], [1, 1, 1, 1]]Each token gets three vectors because it plays three roles. Training tunes W_q and W_k so tokens that should exchange information produce large dot products.
Step 2: Compute scaled dot-product attention by hand
show(matmul(Q, transpose(K)), 1)
# [0, 4, 2]
# [2, 8, 6]
# [1, 8, 5]
out, w = attention(Q, K, V)
show(w)
# [0.09, 0.665, 0.245]
# [0.035, 0.705, 0.259]
# [0.024, 0.798, 0.178]
show(out)
# [0.335, 0.335, 1.575, 1.575]
# [0.295, 0.295, 1.67, 1.67]
# [0.202, 0.202, 1.774, 1.774]Check the first row yourself: scores [0, 4, 2] divided by sqrt(4) = 2 give [0, 2, 1], and the softmax of that is [0.09, 0.665, 0.245]. Each row sums to 1, so each output is a weighted average of value rows. “cat” matches every query best, so all outputs lean toward its value [0, 0, 2, 2].
Step 3: See why the scores are divided by sqrt(d_k)
If query and key components are independent with mean 0 and variance 1, their dot product sums d_k terms of variance 1, so its variance is d_k. A seeded simulation confirms it:
import random
rng = random.Random(42)
def dot_variance(d_k, trials=2000, scale=False):
dots = []
for _ in range(trials):
q = [rng.gauss(0, 1) for _ in range(d_k)]
k = [rng.gauss(0, 1) for _ in range(d_k)]
s = sum(a * b for a, b in zip(q, k))
dots.append(s / math.sqrt(d_k) if scale else s)
mean = sum(dots) / trials
return sum((x - mean) ** 2 for x in dots) / trials
for d_k in (4, 64, 512):
print(d_k, round(dot_variance(d_k), 1), round(dot_variance(d_k, scale=True), 2))
# 4 3.8 1.0
# 64 63.9 1.0
# 512 482.3 1.01Raw scores with a standard deviation around 22 push the softmax into saturation:
rng = random.Random(0)
d_k = 512
q = [rng.gauss(0, 1) for _ in range(d_k)]
keys = [[rng.gauss(0, 1) for _ in range(d_k)] for _ in range(6)]
raw = [sum(a * b for a, b in zip(q, k)) for k in keys]
print([round(s, 1) for s in raw]) # [-23.6, -17.1, 21.6, -43.7, -22.4, -15.7]
print([round(p, 3) for p in softmax(raw)])
# [0.0, 0.0, 1.0, 0.0, 0.0, 0.0]
print([round(p, 3) for p in softmax([s / math.sqrt(d_k) for s in raw])])
# [0.08, 0.106, 0.585, 0.033, 0.084, 0.113]A nearly one-hot softmax hurts training. Its gradient involves p(1 - p), which is close to zero when p is near 0 or 1, so the query and key projections get almost no learning signal. Dividing by sqrt(d_k) restores variance 1 for any head size.
Step 4: Add the causal mask
A decoder generates left to right, so in training position i must not see later tokens. Scores set to negative infinity get softmax weight exactly zero:
out_c, w_c = attention(Q, K, V, causal=True)
show(w_c)
# [1.0, 0.0, 0.0]
# [0.047, 0.953, 0.0]
# [0.024, 0.798, 0.178]
show(out_c)
# [1.0, 1.0, 0.0, 0.0]
# [0.047, 0.047, 1.905, 1.905]
# [0.202, 0.202, 1.774, 1.774]The last row is unchanged because the last token already saw everything. The mask lets a decoder train on every position in one parallel pass while each prediction uses only the past. Encoders such as BERT omit it; padding masks are a separate mechanism.
Step 5: Size the KV cache
Each head computes an n by n score matrix, so attention time and memory grow with n squared. FlashAttention computes the same result in tiles without storing that matrix, making memory linear, but compute stays quadratic.
During generation the model keeps every previous token’s keys and values per layer, so each step projects only the new token and attends over the cache. The cache costs 2 (K and V) times layers, KV heads, head dimension, sequence length, batch and bytes per element. This configuration is illustrative, roughly an 8B-class model with grouped-query attention, in fp16:
def kv_cache_bytes(layers, kv_heads, head_dim, seq_len, batch, bytes_per_elem):
return 2 * layers * kv_heads * head_dim * seq_len * batch * bytes_per_elem # 2 = K and V
GiB = 1024 ** 3
cfg = dict(layers=32, kv_heads=8, head_dim=128, bytes_per_elem=2) # illustrative, fp16
print(kv_cache_bytes(seq_len=1, batch=1, **cfg) // 1024, "KiB per token") # 128 KiB per token
for seq_len, batch in [(4_096, 1), (32_768, 1), (32_768, 16), (131_072, 8)]:
print(seq_len, batch, kv_cache_bytes(seq_len=seq_len, batch=batch, **cfg) / GiB, "GiB")
# 4096 1 0.5 GiB
# 32768 1 4.0 GiB
# 32768 16 64.0 GiB
# 131072 8 128.0 GiB
mha = dict(cfg, kv_heads=32) # same model without grouped-query attention
print(kv_cache_bytes(seq_len=32_768, batch=16, **mha) / GiB, "GiB without GQA")
# 256.0 GiB without GQAGrouped-query attention (GQA) lets several query heads share one key/value head; here 32 query heads share 8, cutting the cache fourfold. Multi-query attention uses a single KV head.
Worked scenario
A team writes its own attention layer for a small decoder trained on support tickets. Training loss falls far faster than in earlier runs and validation perplexity approaches 1, yet generated text is repetitive and incoherent. The illustrative PyTorch 2.x call:
import torch.nn.functional as F
# q, k, v: (batch, heads, seq_len, head_dim)
out = F.scaled_dot_product_attention(q, k, v) # bug: no causal mask
out = F.scaled_dot_product_attention(q, k, v, is_causal=True) # fixIn next-token training the target at position t is the input at position t + 1. Without a mask, position t attends to t + 1 and copies it, learning a lookup instead of language modelling. Validation used the same leaky forward pass, so it looked equally good. At generation time the future token does not exist, so the shortcut has nothing to copy. The fix is the causal mask, retraining, and a regression test that changing a later token never changes earlier outputs (see the next section).
Common mistake
- “Q, K and V are the same vectors.” They are three different learned projections of the same input.
- “The sqrt(d_k) is cosmetic.” It controls score variance so the softmax keeps useful gradients.
- “More heads means more parameters.” With head_dim = d_model / h the projections stay at 4 d_model squared parameters, plus biases.
- “Attention knows word order.” Without positional encoding it is permutation-equivariant.
- “The KV cache stores attention weights.” It stores keys and values; weights for the new token are recomputed each step.
- “FlashAttention makes attention linear.” It cuts memory use and traffic, not the quadratic operation count.
Verify the behavior
Append these tests to the same file as the snippets above and run it with python3:
def layer(X, causal):
return attention(matmul(X, W_q), matmul(X, W_k), matmul(X, W_v), causal)
def test_rows_are_distributions():
_, w = layer(X, causal=False)
assert all(abs(sum(row) - 1) < 1e-9 for row in w)
def test_causal_outputs_ignore_the_future():
changed = X[:2] + [[5, -3, 2, 7]] # replace only the last token
before, _ = layer(X, causal=True)
after, _ = layer(changed, causal=True)
assert before[:2] == after[:2] # earlier positions are untouched
def test_unmasked_outputs_leak_the_future():
changed = X[:2] + [[5, -3, 2, 7]]
assert layer(X, causal=False)[0][0] != layer(changed, causal=False)[0][0]
def test_attention_is_order_blind_without_positions():
perm = [2, 0, 1]
out, _ = layer(X, causal=False)
out_perm, _ = layer([X[i] for i in perm], causal=False)
assert all(abs(a - b) < 1e-9 for i, p in enumerate(perm) for a, b in zip(out_perm[i], out[p]))
def test_head_count_does_not_change_parameters():
assert mha_params(512, 1) == mha_params(512, 8) == 4 * 512 * 512
for name, fn in list(globals().items()):
if name.startswith("test_"):
fn()
print("ok") # okFollow-up questions
Why does pre-LN train more stably than post-LN? The residual stream passes from input to output through additions only, so gradients reach early layers without crossing a normalisation in every block. Post-LN models are more sensitive to learning rate and usually need warmup.
How does decoding cost change with a KV cache? Without it, step t recomputes keys and values for all t tokens. With it, each step projects only the new token and does attention work linear in t, paying in memory that grows with context.
What is the difference between an encoder and a decoder? An encoder uses unmasked, bidirectional attention for understanding tasks. A decoder uses causal self-attention to generate. Encoder-decoder models add cross-attention from decoder queries to encoder keys and values.
Interview exercise
You serve the illustrative model above (32 layers, 8 KV heads, head dimension 128, fp16 cache) on an 80 GiB GPU. The weights take about 15 GiB. Batch 16 at 8,192 tokens runs fine, but raising the context limit to 32,768 tokens causes out-of-memory errors. Explain why and propose fixes.
Answer and reasoning
Each token costs 128 KiB of cache. Batch 16 at 8,192 tokens is 131,072 tokens, or 16 GiB, which fits beside the weights. At 32,768 tokens it is 524,288 tokens, or 64 GiB, and with 15 GiB of weights that is 79 GiB of an 80 GiB card before activations, the CUDA context and fragmentation. Serving engines also reserve only part of memory (vLLM defaults to 90 percent). The cache is linear in context and batch, so quadrupling the context quadruples it. Fixes: limit batches by total cached tokens rather than request count, use paged KV allocation so short requests do not reserve the maximum length, quantize the cache to 8 bits to halve it, or shard across GPUs. Show the arithmetic before the fixes.