Attention Mechanisms and Transformers

Scaled Dot-Product and Multi-Head Attention


A team implements attention from the formula, sets dk=64d_k = 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−1110^{-11}.

The missing division was by 64=8\sqrt{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.

Attention(Q,K,V)=softmax ⁣(QK⊤dk)V\mathrm{Attention}(Q,K,V) = \mathrm{softmax}\!\left(\frac{QK^{\top}}{\sqrt{d_k}}\right)V

Eight heads: split, attend apart, put back togetherOne 512-dimvector per tokenProject to 8sets of64-dim Q, K, VEach headscores andsoftmaxes aloneConcatenatethe 8 outputsback to 512One outputprojectionmixes themTotal work is unchanged: 8 heads of 64 dimensions cost the same as 1 head of 512.
One softmax row can only put its mass in one place, so heads exist to let a token attend several ways at once.

Why the dot product

It measures alignment, cheaply

The dot product of two vectors relates to the angle between them:

q⋅k=∥q∥ ∥k∥cos⁡θq \cdot k = \|q\|\,\|k\|\cos\theta

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]q = [1, 2, 0]:

KeyVectorq⋅kq \cdot kInterpretation
k1k_1[1,2,0][1, 2, 0]1+4+0=51 + 4 + 0 = 5same direction, strong match
k2k_2[2,4,0][2, 4, 0]2+8+0=102 + 8 + 0 = 10same direction, twice the norm — twice the score
k3k_3[0,0,3][0, 0, 3]00orthogonal, no match
k4k_4[−1,−2,0][-1, -2, 0]−5-5opposed, 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:

eij=va⊤tanh⁡(Wqqi+Wkkj)e_{ij} = v_a^{\top}\tanh(W_q q_i + W_k k_j)

This works, and on paper the asymptotic cost is the same O(T2d)O(T^2 d). 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 attentionDot-product attention
Score for one pairTwo matrix-vector products, a tanh⁡\tanh over dd elements, a dot productOne dot product over dd elements
All T2T^2 pairsCannot be fused into a single matrix multiply — the nonlinearity sits between the pair-specific sum and the scoreExactly one GEMM: QK⊤Q K^{\top}
Hardware mappingMany small ops, memory-bandwidth boundMaps onto tensor cores at near peak throughput
Extra parametersWq,Wk,vaW_q, W_k, v_a on top of the projectionsNone beyond the projections
Quality at small dkd_kComparableComparable
Quality at large dkd_k without scalingFine — the tanh⁡\tanh bounds the pre-scoreDegrades badly (this is the reason dk\sqrt{d_k} 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][-1, 1]; that bound makes softmax too flat unless you add a learned temperature. Bilinear scoring, q⊤Wkq^{\top} W k, is strictly more general than the dot product — and strictly redundant, because WQW^Q and WKW^K are already learned, so WW 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\sqrt{d_k}

The variance argument

Assume the components of qq and kk are independent, mean 0, variance 1 — approximately what you get from standard initialisation followed by layer normalisation. Then for each term of the sum:

E[qiki]=E[qi] E[ki]=0,Var(qiki)=E[qi2] E[ki2]=1\mathbb{E}[q_i k_i] = \mathbb{E}[q_i]\,\mathbb{E}[k_i] = 0, \qquad \mathrm{Var}(q_i k_i) = \mathbb{E}[q_i^2]\,\mathbb{E}[k_i^2] = 1

Independent variances add, so summing dkd_k of them:

Var(q⋅k)=dk⟹sd(q⋅k)=dk\mathrm{Var}(q \cdot k) = d_k \quad\Longrightarrow\quad \mathrm{sd}(q \cdot k) = \sqrt{d_k}
dkd_ksd of scoresTypical spread (±3\pm 3 sd)Largest logit gap
42−6-6 to +6+612
648−24-24 to +24+2448
51222.6−68-68 to +68+68136

Dividing every score by dk\sqrt{d_k} rescales the standard deviation back to 1 regardless of dkd_k, 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 α\alpha, the diagonal of its Jacobian is

∂αi∂si=αi(1−αi)\frac{\partial \alpha_i}{\partial s_i} = \alpha_i (1 - \alpha_i)

Take a two-key example and vary the gap:

Logit gap Δ\Deltaαmax⁡=eΔ1+eΔ\alpha_{\max} = \dfrac{e^{\Delta}}{1+e^{\Delta}}Gradient α(1−α)\alpha(1-\alpha)
00.5000.250
30.9530.045
60.99752.5×10−32.5 \times 10^{-3}
120.99999396.1×10−66.1 \times 10^{-6}
241−3.8×10−111 - 3.8\times10^{-11}3.8×10−113.8 \times 10^{-11}

At dk=64d_k = 64 unscaled, gaps of 24 are ordinary. A gradient of 3.8×10−113.8 \times 10^{-11} multiplied through the rest of the backward pass is indistinguishable from zero in float32, whose epsilon is about 1.2×10−71.2 \times 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 dkd_k changes, and softmax has a usable gradient only near variance 1.

Softmax, and the two things it guarantees

αj=esj∑m=1Tesm\alpha_j = \frac{e^{s_j}}{\sum_{m=1}^{T} e^{s_m}}

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 cc to every score changes nothing:

esj+c∑mesm+c=ecesjec∑mesm=esj∑mesm\frac{e^{s_j + c}}{\sum_m e^{s_m + c}} = \frac{e^{c} e^{s_j}}{e^{c}\sum_m e^{s_m}} = \frac{e^{s_j}}{\sum_m e^{s_m}}

Every serious implementation exploits this for numerical stability by subtracting the row maximum first. It matters: with scores [100,101,102][100, 101, 102], computing e102≈2×1044e^{102} \approx 2 \times 10^{44} overflows float32 (max ≈3.4×1038\approx 3.4 \times 10^{38}) and you get nan. Subtract the max to get [−2,−1,0][-2, -1, 0]:

e−2=0.1353,e−1=0.3679,e0=1.0000,sum=1.5032e^{-2} = 0.1353,\quad e^{-1} = 0.3679,\quad e^{0} = 1.0000,\quad \text{sum} = 1.5032
α=[0.0900, 0.2447, 0.6652]\alpha = [0.0900,\ 0.2447,\ 0.6652]

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\alpha V, is where the actual information moves. Three tokens, dv=2d_v = 2, values v1=[2,0]v_1 = [2,0], v2=[0,4]v_2 = [0,4], v3=[1,1]v_3 = [1,1], weights [0.074, 0.620, 0.306][0.074,\ 0.620,\ 0.306]:

z=0.074[2,0]+0.620[0,4]+0.306[1,1]=[0.148+0.306,  2.480+0.306]=[0.454, 2.786]z = 0.074[2,0] + 0.620[0,4] + 0.306[1,1] = [0.148 + 0.306,\ \ 2.480 + 0.306] = [0.454,\ 2.786]

Because the weights are a probability distribution, zz is the expected value vector under that distribution. If the weights were uniform (1/31/3 each) you would get the plain mean [1,1.667][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]v_{\text{cat}} = [4, 0] and vtired=[0,4]v_{\text{tired}} = [0, 4].

With one head, the best a single distribution can do is split: αcat=αtired=0.5\alpha_{\text{cat}} = \alpha_{\text{tired}} = 0.5, giving

z=0.5[4,0]+0.5[0,4]=[2,2]z = 0.5[4,0] + 0.5[0,4] = [2, 2]

That output is identical to what you would get from a single token with value [2,2][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][4, 0]. Head 2 puts weight 1.0 on tired and outputs [0,4][0, 4]. Concatenate:

z=[4, 0, 0, 4]z = [4,\ 0,\ 0,\ 4]

Both facts survive, in known coordinates, and the output projection WOW^O can route each to wherever downstream layers want it.

The mechanics

MultiHead(X)=Concat(head1,…,headh) WO\mathrm{MultiHead}(X) = \mathrm{Concat}(\mathrm{head}_1, \dots, \mathrm{head}_h)\,W^O
headi=Attention(XWiQ, XWiK, XWiV)\mathrm{head}_i = \mathrm{Attention}(XW_i^Q,\ XW_i^K,\ XW_i^V)

The standard configuration splits rather than adds capacity: with dmodel=512d_{model} = 512 and h=8h = 8, each head uses dk=dv=dmodel/h=64d_k = d_v = d_{model}/h = 64. The parameter count is therefore the same as a single head of full width:

MatrixShapeParameters
WQW^Q (all heads, one fused matrix)512×512512 \times 512262,144
WKW^K512×512512 \times 512262,144
WVW^V512×512512 \times 512262,144
WOW^O512×512512 \times 512262,144
Total per attention block1,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

Python
import torchimport torch.nn as nnimport torch.nn.functional as Fimport mathclass MultiHeadAttention(nn.Module):    def __init__(self, d_model, num_heads, dropout=0.1):        super().__init__()        assert d_model % num_heads == 0, "d_model must be divisible by num_heads"        self.d_model = d_model        self.h = num_heads        self.d_k = d_model // num_heads        # One fused matrix per role; the head split happens by reshaping.        self.W_q = nn.Linear(d_model, d_model)        self.W_k = nn.Linear(d_model, d_model)        self.W_v = nn.Linear(d_model, d_model)        self.W_o = nn.Linear(d_model, d_model)        self.dropout = nn.Dropout(dropout)    def _split(self, x):        B, T, _ = x.shape        # (B, T, d_model) -> (B, h, T, d_k)        return x.view(B, T, self.h, self.d_k).transpose(1, 2)    def forward(self, query, key, value, mask=None):        B, T_q, _ = query.shape        Q = self._split(self.W_q(query))   # (B, h, T_q, d_k)        K = self._split(self.W_k(key))     # (B, h, T_k, d_k)        V = self._split(self.W_v(value))   # (B, h, T_k, d_k)        scores = Q @ K.transpose(-2, -1) / math.sqrt(self.d_k)   # (B, h, T_q, T_k)        if mask is not None:            # mask broadcasts over the head axis:            #   padding mask  -> (B, 1, 1,  T_k)            #   causal mask   -> (1, 1, T_q, T_k)            scores = scores.masked_fill(mask == 0, float('-inf'))        attn = F.softmax(scores, dim=-1)        attn = self.dropout(attn)        out = attn @ V                                    # (B, h, T_q, d_k)        out = out.transpose(1, 2).contiguous().view(B, T_q, self.d_model)        return self.W_o(out), attn

Passing 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 α\alpha after softmax, which means during training the rows no longer sum to 1. With p=0.1p = 0.1, PyTorch zeroes 10% of entries and scales the survivors by 1/(1−0.1)=1.1111/(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 typeShapeWhat it blocksVaries with
Padding(B,1,1,Tk)(B, 1, 1, T_k)Columns corresponding to [PAD] tokensEach example in the batch
Causal(1,1,Tq,Tk)(1, 1, T_q, T_k)Upper triangle — every future positionNothing; build once, cache it
Combined(B,1,Tq,Tk)(B, 1, T_q, T_k)Both, via logical ANDBatch and position

The failure mode here can be silent. If you build a padding mask of shape (B,Tk)(B, T_k) and pass it directly, broadcasting against (B,h,Tq,Tk)(B, h, T_q, T_k) aligns it as (1,1,B,Tk)(1, 1, B, T_k): the batch axis lands on the query axis. When BB differs from TqT_q 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 dvd_v differs from dkd_k

Nothing in the formula requires dv=dkd_v = d_k. 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 hh separate query heads but a single shared key and value head, which shrinks the inference-time key-value cache by a factor of hh — 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

StageMultiply-accumulatesMemory for activations
Q, K, V, O projections4Td24 T d^2O(Td)O(Td)
QK⊤QK^{\top}T2dT^2 dO(hT2)O(h T^2) — the score matrix
softmaxO(hT2)O(hT^2)O(hT2)O(hT^2) — kept for the backward pass
αV\alpha VT2dT^2 dO(Td)O(Td)

The projections cost 4Td24Td^2 and the attention itself 2T2d2T^2d. Set them equal:

2T2d=4Td2⟹T=2d2T^2 d = 4 T d^2 \quad\Longrightarrow\quad T = 2d

With d=512d = 512 the crossover is T=1024T = 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×TT \times T matrix must be stored per head for the backward pass. Batch 32, 8 heads, T=1024T = 1024, fp16:

32×8×1024×1024×2 bytes=536,870,912 B=512 MiB32 \times 8 \times 1024 \times 1024 \times 2 \text{ bytes} = 536{,}870{,}912 \text{ B} = 512 \text{ MiB}

Per layer. A 12-layer model needs 6 GiB for attention probabilities alone. Double TT 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)O(T^2) to O(T)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:

Python
out = F.scaled_dot_product_attention(    Q, K, V,                 # each (B, h, T, d_k)    attn_mask=None,    dropout_p=0.1 if self.training else 0.0,    is_causal=True,          # cheaper than materialising a triangular mask)

Write the explicit loop version anyway, keep it in your test suite, and assert the two agree to within 10−510^{-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:

MistakeSymptomHow to detect it
Missing dk\sqrt{d_k}Loss drops briefly then flatlines from step ~50Print max⁡(α)\max(\alpha) per row; if it is above 0.999 at initialisation, scaling is wrong
Softmax over the wrong axisTrains, converges to a mediocre plateau, never errorsassert torch.allclose(attn.sum(-1), torch.ones_like(attn.sum(-1))) in eval mode
Mask broadcast against the wrong axisPadding leaks in, or causal leaks the future; validation loss much worse than trainingChange a token at position jj and confirm the output at every i<ji < j is unchanged

On head count: more heads means more independent distributions but a smaller dkd_k each, and dkd_k below about 32 starts to hurt because each head has too little room to represent a useful similarity function. For dmodel=512d_{model} = 512, eight heads at dk=64d_k = 64 is the standard choice for a reason. If you increase heads, increase dmodeld_{model} with them.