Course Content
Attention Mechanisms and Transformers
4 sections · 11 lessons
Scaled Dot-Product and Multi-Head Attention
A team implements attention from the formula, sets dk=64 because that is what the paper used, and forgets one division. The loss starts at 10.4, drops to 9.8 in the first fifty steps, and then stops. Flat. Not slowly improving — flat to four decimal places for the next ten thousand steps.
They check the learning rate, the initialisation, the data loader. Everything is correct. Then someone prints the attention weights and finds rows like [1.0, 0.0, 0.0, ..., 0.0] — not approximately, but to the limit of float32. Every query has locked onto exactly one key and will never move, because the gradient through a saturated softmax is on the order of 10−11.
The missing division was by 64=8. That is the difference between a model that trains and a model that does not. This lesson takes the attention formula apart term by term — why the dot product, why the square root, why softmax specifically, and why the whole thing gets run eight times in parallel instead of once.
Why the dot product
It measures alignment, cheaply
The dot product of two vectors relates to the angle between them:
So it is large when the vectors point in the same direction, zero when perpendicular, negative when opposed — and it also scales with magnitude, which turns out to be a feature rather than a bug. A key with a large norm can advertise itself more loudly, and models do learn to use this.
Concretely, with q=[1,2,0]:
| Key | Vector | q⋅k | Interpretation |
|---|---|---|---|
| k1 | [1,2,0] | 1+4+0=5 | same direction, strong match |
| k2 | [2,4,0] | 2+8+0=10 | same direction, twice the norm — twice the score |
| k3 | [0,0,3] | 0 | orthogonal, no match |
| k4 | [−1,−2,0] | −5 | opposed, actively suppressed |
What was the alternative
The original neural attention mechanism did not use a dot product. It used additive attention, a small feed-forward network scoring each query-key pair:
This works, and on paper the asymptotic cost is the same O(T2d). In practice it is far slower, and the reason is worth understanding because it explains a lot about why transformer design looks the way it does.
| Additive attention | Dot-product attention | |
|---|---|---|
| Score for one pair | Two matrix-vector products, a tanh over d elements, a dot product | One dot product over d elements |
| All T2 pairs | Cannot be fused into a single matrix multiply — the nonlinearity sits between the pair-specific sum and the score | Exactly one GEMM: QK⊤ |
| Hardware mapping | Many small ops, memory-bandwidth bound | Maps onto tensor cores at near peak throughput |
| Extra parameters | Wq,Wk,va on top of the projections | None beyond the projections |
| Quality at small dk | Comparable | Comparable |
| Quality at large dk without scaling | Fine — the tanh bounds the pre-score | Degrades badly (this is the reason dk exists) |
Other scoring functions exist and are occasionally useful. Cosine similarity divides by the norms, which throws away magnitude and bounds scores to [−1,1]; that bound makes softmax too flat unless you add a learned temperature. Bilinear scoring, q⊤Wk, is strictly more general than the dot product — and strictly redundant, because WQ and WK are already learned, so W can be absorbed into either of them.
The dot product won not because it scores better, but because it is a single matrix multiply — and everything the transformer does is arranged so that the hardware sees big matrix multiplies.
Why divide by dk
The variance argument
Assume the components of q and k are independent, mean 0, variance 1 — approximately what you get from standard initialisation followed by layer normalisation. Then for each term of the sum:
Independent variances add, so summing dk of them:
| dk | sd of scores | Typical spread (±3 sd) | Largest logit gap |
|---|---|---|---|
| 4 | 2 | −6 to +6 | 12 |
| 64 | 8 | −24 to +24 | 48 |
| 512 | 22.6 | −68 to +68 | 136 |
Dividing every score by dk rescales the standard deviation back to 1 regardless of dk, which puts softmax in a consistent operating range for any head size.
The gradient argument, with numbers
Why does a large logit gap matter? Because softmax's derivative collapses. For a distribution α, the diagonal of its Jacobian is
Take a two-key example and vary the gap:
| Logit gap Δ | αmax=1+eΔeΔ | Gradient α(1−α) |
|---|---|---|
| 0 | 0.500 | 0.250 |
| 3 | 0.953 | 0.045 |
| 6 | 0.9975 | 2.5×10−3 |
| 12 | 0.9999939 | 6.1×10−6 |
| 24 | 1−3.8×10−11 | 3.8×10−11 |
At dk=64 unscaled, gaps of 24 are ordinary. A gradient of 3.8×10−11 multiplied through the rest of the backward pass is indistinguishable from zero in float32, whose epsilon is about 1.2×10−7. Scale by 8 and that gap becomes 3, with gradient 0.045 — nine orders of magnitude larger.
Note that this is a problem at initialisation, before any learning has happened. The model does not saturate because it became confident; it starts saturated and therefore never becomes anything. That is why the symptom is a loss that flatlines almost immediately rather than one that degrades over time.
The square root is not a tuning constant. It is the exact factor that keeps score variance at 1 as dk changes, and softmax has a usable gradient only near variance 1.
Softmax, and the two things it guarantees
Two properties do the work.
The weights are non-negative and sum to 1. This makes the output a convex combination of the value vectors — it lies inside their convex hull, so it cannot blow up in magnitude no matter what the scores are. If instead you normalised by dividing by the sum of raw scores, a negative score would produce a negative weight and the output could leave that hull entirely.
It is shift-invariant. Adding a constant c to every score changes nothing:
Every serious implementation exploits this for numerical stability by subtracting the row maximum first. It matters: with scores [100,101,102], computing e102≈2×1044 overflows float32 (max ≈3.4×1038) and you get nan. Subtract the max to get [−2,−1,0]:
Same answer, no overflow. PyTorch's F.softmax does this internally, which is one good reason not to hand-roll it.
Weighting the values
The final multiply, αV, is where the actual information moves. Three tokens, dv=2, values v1=[2,0], v2=[0,4], v3=[1,1], weights [0.074, 0.620, 0.306]:
Because the weights are a probability distribution, z is the expected value vector under that distribution. If the weights were uniform (1/3 each) you would get the plain mean [1,1.667]. Attention is a learned, input-dependent departure from averaging.
Multi-head attention
The problem one head cannot solve
A single attention head produces one probability distribution per query. One distribution can emphasise one thing. But a token frequently needs several unrelated things at once.
Take The tired cat sat and the query token sat. It needs its subject (cat) to conjugate, and it may also need the modifier (tired) for semantics. Suppose the value vectors are vcat=[4,0] and vtired=[0,4].
With one head, the best a single distribution can do is split: αcat=αtired=0.5, giving
That output is identical to what you would get from a single token with value [2,2]. The two facts have been averaged into one and cannot be separated afterwards. This is destructive interference, and it is the concrete cost of one head.
Now run two heads in separate 2-dimensional subspaces. Head 1 puts weight 1.0 on cat and outputs [4,0]. Head 2 puts weight 1.0 on tired and outputs [0,4]. Concatenate:
Both facts survive, in known coordinates, and the output projection WO can route each to wherever downstream layers want it.
The mechanics
The standard configuration splits rather than adds capacity: with dmodel=512 and h=8, each head uses dk=dv=dmodel/h=64. The parameter count is therefore the same as a single head of full width:
| Matrix | Shape | Parameters |
|---|---|---|
| WQ (all heads, one fused matrix) | 512×512 | 262,144 |
| WK | 512×512 | 262,144 |
| WV | 512×512 | 262,144 |
| WO | 512×512 | 262,144 |
| Total per attention block | 1,048,576 |
Eight heads cost exactly what one head of size 512 would. You are not buying capacity — you are buying independence, eight separate softmax distributions instead of one. That is the entire trade.
What heads actually learn, in trained models, is reasonably consistent: some track the previous or next token positionally, some track syntactic relations such as verb-to-subject or noun-to-determiner, some track rarer patterns like matching brackets or coreference. Many heads in large models learn nothing useful and can be pruned with little loss — head redundancy is well documented.
Implementation
1import torch2import torch.nn as nn3import torch.nn.functional as F4import math56class MultiHeadAttention(nn.Module):7 def __init__(self, d_model, num_heads, dropout=0.1):8 super().__init__()9 assert d_model % num_heads == 0, "d_model must be divisible by num_heads"10 self.d_model = d_model11 self.h = num_heads12 self.d_k = d_model // num_heads1314 # One fused matrix per role; the head split happens by reshaping.15 self.W_q = nn.Linear(d_model, d_model)16 self.W_k = nn.Linear(d_model, d_model)17 self.W_v = nn.Linear(d_model, d_model)18 self.W_o = nn.Linear(d_model, d_model)19 self.dropout = nn.Dropout(dropout)2021 def _split(self, x):22 B, T, _ = x.shape23 # (B, T, d_model) -> (B, h, T, d_k)24 return x.view(B, T, self.h, self.d_k).transpose(1, 2)2526 def forward(self, query, key, value, mask=None):27 B, T_q, _ = query.shape2829 Q = self._split(self.W_q(query)) # (B, h, T_q, d_k)30 K = self._split(self.W_k(key)) # (B, h, T_k, d_k)31 V = self._split(self.W_v(value)) # (B, h, T_k, d_k)3233 scores = Q @ K.transpose(-2, -1) / math.sqrt(self.d_k) # (B, h, T_q, T_k)3435 if mask is not None:36 # mask broadcasts over the head axis:37 # padding mask -> (B, 1, 1, T_k)38 # causal mask -> (1, 1, T_q, T_k)39 scores = scores.masked_fill(mask == 0, float('-inf'))4041 attn = F.softmax(scores, dim=-1)42 attn = self.dropout(attn)4344 out = attn @ V # (B, h, T_q, d_k)45 out = out.transpose(1, 2).contiguous().view(B, T_q, self.d_model)46 return self.W_o(out), attnPassing query, key and value separately is what lets the same class serve self-attention (pass x, x, x) and cross-attention (pass decoder_x, encoder_out, encoder_out).
The .contiguous() before .view() is not optional. After transpose the tensor's memory layout no longer matches its shape, and view will raise. Using reshape instead hides the copy but does the same work.
Practical details that bite
Dropout on attention weights
Dropout is applied to α after softmax, which means during training the rows no longer sum to 1. With p=0.1, PyTorch zeroes 10% of entries and scales the survivors by 1/(1−0.1)=1.111, so the row sums to 1 in expectation but not in any individual sample. This is intentional: it forces the model not to depend on one specific attention edge. It also means your "rows sum to 1" assertion must run under model.eval() or with dropout disabled, or it will fail spuriously.
Mask shapes and broadcasting
| Mask type | Shape | What it blocks | Varies with |
|---|---|---|---|
| Padding | (B,1,1,Tk) | Columns corresponding to [PAD] tokens | Each example in the batch |
| Causal | (1,1,Tq,Tk) | Upper triangle — every future position | Nothing; build once, cache it |
| Combined | (B,1,Tq,Tk) | Both, via logical AND | Batch and position |
The failure mode here can be silent. If you build a padding mask of shape (B,Tk) and pass it directly, broadcasting against (B,h,Tq,Tk) aligns it as (1,1,B,Tk): the batch axis lands on the query axis. When B differs from Tq you get a shape error, which is the lucky case. When they happen to match — a batch of 32 with 32-token queries — no error is raised, each example is masked with some other example's padding, and the model trains to a mediocre plateau. Always expand masks to four dimensions explicitly.
When dv differs from dk
Nothing in the formula requires dv=dk. Query and key must share a dimension so the dot product is defined; the value dimension only sets the output width per head. Real architectures exploit this. Multi-query attention keeps h separate query heads but a single shared key and value head, which shrinks the inference-time key-value cache by a factor of h — the difference between fitting a long context in memory and not. Grouped-query attention splits the difference, sharing K and V across groups of, say, 8 query heads.
What it costs
| Stage | Multiply-accumulates | Memory for activations |
|---|---|---|
| Q, K, V, O projections | 4Td2 | O(Td) |
| QK⊤ | T2d | O(hT2) — the score matrix |
| softmax | O(hT2) | O(hT2) — kept for the backward pass |
| αV | T2d | O(Td) |
The projections cost 4Td2 and the attention itself 2T2d. Set them equal:
With d=512 the crossover is T=1024. Below that, the projections dominate and attention is nearly free. Above it, the quadratic term takes over and grows without bound.
Memory is the harder wall, because the T×T matrix must be stored per head for the backward pass. Batch 32, 8 heads, T=1024, fp16:
Per layer. A 12-layer model needs 6 GiB for attention probabilities alone. Double T to 2048 and it becomes 2 GiB per layer, 24 GiB total — you are out of memory on most hardware before you are out of compute.
This is precisely what FlashAttention addresses. It never materialises the full matrix: it processes the score matrix in tiles that fit in on-chip SRAM, computing softmax with a running maximum and running sum, and recomputes the needed tiles during the backward pass rather than storing them. The FLOP count is unchanged — slightly higher, in fact — but memory drops from O(T2) to O(T) and wall-clock time improves because the operation stops being bandwidth-bound.
What this means when you build something
Use F.scaled_dot_product_attention in production rather than the hand-written version above. It dispatches to a fused kernel, handles the scaling and numerical stability correctly, and gives you memory-efficient attention for free:
1out = F.scaled_dot_product_attention(2 Q, K, V, # each (B, h, T, d_k)3 attn_mask=None,4 dropout_p=0.1 if self.training else 0.0,5 is_causal=True, # cheaper than materialising a triangular mask6)Write the explicit loop version anyway, keep it in your test suite, and assert the two agree to within 10−5. When something is wrong you will need the version whose intermediates you can print. Note the one real cost of the fused kernel: it never produces the attention matrix, so if you need weights for analysis you must fall back to the explicit path.
Three failure modes are worth committing to memory, because each one produces a model that runs without complaint:
| Mistake | Symptom | How to detect it |
|---|---|---|
| Missing dk | Loss drops briefly then flatlines from step ~50 | Print max(α) per row; if it is above 0.999 at initialisation, scaling is wrong |
| Softmax over the wrong axis | Trains, converges to a mediocre plateau, never errors | assert torch.allclose(attn.sum(-1), torch.ones_like(attn.sum(-1))) in eval mode |
| Mask broadcast against the wrong axis | Padding leaks in, or causal leaks the future; validation loss much worse than training | Change a token at position j and confirm the output at every i<j is unchanged |
On head count: more heads means more independent distributions but a smaller dk each, and dk below about 32 starts to hurt because each head has too little room to represent a useful similarity function. For dmodel=512, eight heads at dk=64 is the standard choice for a reason. If you increase heads, increase dmodel with them.