Course Content
Attention Mechanisms and Transformers
4 sections · 11 lessons
Self-Attention Mechanism Explained
Take the sentence The animal didn't cross the street because it was too tired. Now change one word: The animal didn't cross the street because it was too wide.
In the first sentence it means the animal. In the second it means the street. Nothing about the token it changed — same spelling, same embedding, same position. What changed is which other word in the sentence it should be looking at, and that depends on a word (tired versus wide) that appears three positions after the pronoun.
Any representation that assigns it a fixed vector is stuck. What you need is a mechanism where the representation of a token is built from the other tokens, with the mixing weights decided by content rather than by position. That mechanism is self-attention, and the whole of it is four steps of linear algebra you can do by hand.
Query, key, value — the retrieval analogy
Self-attention borrows its structure from database lookup. Imagine a filing system where each document has a short index card describing it, and you search by writing down what you want. Three objects are involved:
- Query (q) — what this token is looking for. For
it, something like "I am a pronoun; find me a noun I could refer to." - Key (k) — the index card each token advertises. For
animal, something like "I am an animate singular noun." - Value (v) — the actual content a token contributes if it is selected. Distinct from the key, because how you are found and what you deliver are different jobs.
The difference between a real database and attention is that the lookup is soft. A database returns one row. Attention returns a weighted blend of every row, with weights that sum to 1. That softness is what makes it differentiable, and therefore learnable.
All three come from the same input by three separate learned projections:
where X∈RT×dmodel holds one row per token. The word self in self-attention means exactly this: Q, K and V all derive from the same X. Every token queries every token, itself included.
Keys and values are separate matrices for a reason. A token can advertise "I am a plural noun" as its key while contributing rich semantic content as its value — the address is not the contents.
The four steps, with real numbers
The full operation is one formula:
We will compute it by hand for a three-token sequence with dk=2, so every number is checkable.
Take the tokens the, cat, sat and suppose the projections have produced:
| Token | Key kj | Value vj |
|---|---|---|
| the | [1, 0] | [2, 0] |
| cat | [0, 2] | [0, 4] |
| sat | [1, 1] | [1, 1] |
and the query for sat is q3=[1, 2]. We compute the output for position 3.
Step 1 — relevance scores by dot product
Score every key against the query:
The dot product is a similarity measure: it is large and positive when two vectors point the same way, near zero when they are orthogonal, negative when they oppose. cat scores highest, which is what we would want a verb's query to do.
Step 2 — divide by dk
With dk=2, dk=1.4142. The scaled scores are:
Why this particular constant is load-bearing gets its own section below — it is the step people most often assume is cosmetic.
Step 3 — softmax
Softmax turns scores into a probability distribution:
Exponentiate:
Sum: 2.0281+16.918+8.3418=27.288. Divide:
| Attends to | Scaled score | es | Weight α |
|---|---|---|---|
| the | 0.7071 | 2.0281 | 0.074 |
| cat | 2.8284 | 16.918 | 0.620 |
| sat | 2.1213 | 8.3418 | 0.306 |
| Total | 1.000 |
Check the sum: 0.074+0.620+0.306=1.000. It must, by construction — if your implementation produces a row that does not sum to 1, you have masked incorrectly or softmaxed along the wrong axis, which is the single most common bug in a from-scratch implementation.
Step 4 — weighted sum of values
Componentwise: first coordinate 0.074×2+0.620×0+0.306×1=0.148+0+0.306=0.454. Second coordinate 0+0.620×4+0.306×1=2.480+0.306=2.786. So
Compare that to the value of sat alone, [1,1]. The output has been pulled strongly towards v2=[0,4], the value of cat. The representation of sat now carries information about its subject. That is the entire trick, and it happens for every token in parallel — QK⊤ computes all T2 scores in one matrix multiply.
Why dk is not cosmetic
Run the same example without the scaling. Raw scores 1, 4, 3:
Weights: 0.035, 0.705, 0.259. Compare side by side:
| Attends to | Unscaled weight | Scaled weight |
|---|---|---|
| the | 0.035 | 0.074 |
| cat | 0.705 | 0.620 |
| sat | 0.259 | 0.306 |
At dk=2 the difference is mild. The problem is that it grows with dk, and real models use dk=64 or 128.
Here is the argument. Suppose the components of q and k are independent with mean 0 and variance 1 — roughly what standard initialisation gives you. Then
Each term qiki has mean E[qi]E[ki]=0 and variance E[qi2]E[ki2]=1. Summing dk independent terms adds the variances:
With dk=64 the standard deviation is 8. Scores across a sequence will routinely span three standard deviations either side of zero — a range of roughly 48.
Now feed a gap of 24 into softmax. Take two scores, 0 and 24:
The distribution is a one-hot vector for all practical purposes. That would merely be over-confident if it were not for the gradient. The derivative of softmax with respect to its own logit is
With α≈1−3.8×10−11, that derivative is about 3.8×10−11. Multiply it into the backward pass and the gradient reaching WQ and WK through this attention weight is numerically zero. The layer stops learning — not slowly, but completely, and from step one.
Dividing by dk=8 turns the gap of 24 into a gap of 3:
Still confident, but α(1−α)=0.047×0.953=0.045 — a gradient nine orders of magnitude healthier. The scaling restores unit variance to the scores, which puts softmax back in the range where it has a usable derivative.
Without dk, attention with dk=64 does not train badly — it does not train at all, because softmax saturates before the first gradient step.
Self, cross, and causal attention
The formula never changes. What changes is where Q, K and V come from and whether any scores are suppressed.
| Self-attention | Cross-attention | Causal (masked) self-attention | |
|---|---|---|---|
| Q from | sequence A | sequence A (e.g. decoder) | sequence A |
| K, V from | sequence A | sequence B (e.g. encoder) | sequence A |
| Attention matrix shape | TA×TA | TA×TB | TA×TA, lower-triangular |
| Each token sees | all positions | all of the other sequence | itself and everything before it |
| Typical use | BERT-style encoders | translation, image captioning | GPT-style generation |
Causal masking deserves its arithmetic too, because "set the future to zero" is the wrong instruction and produces a subtly broken model. You must zero the weights after softmax normalisation has excluded them, which means suppressing the scores before softmax by setting them to −∞.
Take position 3 of a four-token sequence with raw scaled scores [2.0, 1.0, 3.0, 5.0]. Position 4 is in the future and must be masked:
Exponentiate: e2=7.389, e1=2.718, e3=20.086, e−∞=0. Sum =30.193. Weights:
which sums to 1.000 over the three legal positions. Now the unmasked version, for contrast: e5=148.413, total 178.606, giving [0.041, 0.015, 0.112, 0.831]. Eighty-three percent of the attention mass would have gone to the token the model is supposed to be predicting. A causal model trained without this mask reaches near-zero training loss and produces gibberish at inference, because at generation time the future column simply does not exist.
The reason −∞ (in practice a large negative number like -1e9, or float('-inf')) is correct while post-softmax zeroing is not: zeroing after softmax leaves the remaining weights summing to 1−0.831=0.169 instead of 1, silently scaling the whole output down by a factor that varies per row.
What this buys over recurrence
| Structural problem in recurrent models | How self-attention removes it |
|---|---|
| Position t cannot be computed until t−1 finishes, so the GPU idles | QK⊤ is one matrix multiply covering all T2 pairs; the sequential depth is O(1) regardless of T |
| Gradient between positions i and j passes through ∣i−j∣ Jacobians and decays exponentially | Positions i and j are connected by a single weight αij — a one-hop path at any distance |
| All history compressed into one fixed-size state | All T value vectors stay available; the summary is rebuilt per query |
The cost is the T×T score matrix. At T=512 that is 262,144 entries per head per layer; at T=2048 it is 4.2 million. The compute grows quadratically and so does the activation memory.
Implementation
1import torch2import torch.nn as nn3import torch.nn.functional as F4import math56class SelfAttention(nn.Module):7 def __init__(self, d_model, d_k=None, dropout=0.1):8 super().__init__()9 self.d_k = d_k or d_model10 self.W_q = nn.Linear(d_model, self.d_k, bias=False)11 self.W_k = nn.Linear(d_model, self.d_k, bias=False)12 self.W_v = nn.Linear(d_model, self.d_k, bias=False)13 self.dropout = nn.Dropout(dropout)1415 def forward(self, x, mask=None):16 # x: (batch, seq_len, d_model)17 Q = self.W_q(x) # (B, T, d_k)18 K = self.W_k(x)19 V = self.W_v(x)2021 # (B, T, d_k) @ (B, d_k, T) -> (B, T, T)22 scores = Q @ K.transpose(-2, -1) / math.sqrt(self.d_k)2324 if mask is not None:25 # mask: True where the position is allowed26 scores = scores.masked_fill(mask == 0, float('-inf'))2728 attn = F.softmax(scores, dim=-1) # normalise over KEYS, the last axis29 attn = self.dropout(attn)30 return attn @ V, attn # (B, T, d_k), (B, T, T)Two lines carry most of the bugs people hit.
dim=-1 in the softmax. Each row of the score matrix is one query against all keys, so normalisation runs along the key axis. Using dim=-2 normalises down columns instead — the code runs, the shapes are identical, the loss decreases a bit, and the model is nonsense. Assert it in a test: attn.sum(-1) must be all ones.
masked_fill before softmax, not after. The arithmetic above shows why.
A causal mask is built once and reused:
1def causal_mask(T, device=None):2 # lower-triangular ones: position i may attend to j <= i3 return torch.tril(torch.ones(T, T, dtype=torch.bool, device=device))45m = causal_mask(4)6# tensor([[ True, False, False, False],7# [ True, True, False, False],8# [ True, True, True, False],9# [ True, True, True, True]])Reading attention weights
The returned attn tensor is the model's own account of what it looked at, and it is worth plotting.
1import matplotlib.pyplot as plt23def plot_attention(weights, tokens):4 # weights: (T, T) for one head, one example5 w = weights.detach().cpu().numpy()6 fig, ax = plt.subplots(figsize=(6, 5))7 im = ax.imshow(w, cmap='viridis', vmin=0, vmax=1)8 ax.set_xticks(range(len(tokens)), tokens, rotation=45, ha='right')9 ax.set_yticks(range(len(tokens)), tokens)10 ax.set_xlabel('attending to (keys)')11 ax.set_ylabel('query token')12 fig.colorbar(im)13 plt.tight_layout()14 return figRead it row by row: row i is where token i looked, and it sums to 1. Three patterns show up constantly and each means something:
| Pattern in the heatmap | What it usually means |
|---|---|
| A bright diagonal | Tokens attending mostly to themselves — common in early layers, and normal |
| A bright vertical stripe on one column | An attention sink: every token dumps leftover probability on one position, often the first token or a separator. Because softmax rows must sum to 1, a head with nothing useful to do has to put its mass somewhere |
| Near-uniform grey, every weight around 1/T | The head has collapsed and is averaging. If this appears everywhere at initialisation it is expected; if it persists after training, that head is dead capacity |
One warning about interpretation, because this is where people overreach. A high attention weight means information flowed along that edge; it does not prove the model's decision depended on it. The value vector could be near zero, or a later layer could discard the contribution. Attention maps are a hypothesis-generating tool, not an explanation.
What this means when you build something
Three consequences follow directly from the arithmetic above, and they shape design decisions long before you write any code.
Sequence length is your dominant cost, not model width. Doubling dmodel makes the projections more expensive but leaves the T×T attention matrix the same size. Doubling T quadruples that matrix — both the compute and the memory holding it. Before scaling context, work out the activation memory: batch × heads ×T2× bytes, per layer. At batch 32, 8 heads, T=1024, fp16, that is 32×8×10242×2=512 MiB for the weights of a single layer.
Your masking is a correctness invariant, so test it. Two assertions catch nearly every masking bug: rows of the attention matrix sum to 1 within tolerance, and a causal model's output at position i is bit-identical when you change any token at position j>i. The second test is cheap and catches leakage that the loss curve will happily hide from you.
Self-attention alone is order-blind. Notice that nothing in the four steps referenced position. Permute the input rows and the output rows permute identically — the operation is permutation-equivariant. The cat sat and Sat cat the produce the same set of representations. Position information must be injected separately into the embeddings; without it your model is a very expensive bag of words, and the symptom is a model that gets word identity right and word order consistently wrong.