Attention Mechanisms and Transformers

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.

One word changed, and "it" points somewhere else0.510.030.090.210.160.080.020.620.170.11animalcrossstreetitadjit — too tiredit — too wideSame sentence, same weights; only the final adjective differs, and the row redistributes.
The query comes from "it", the keys from every word, and the softmax row decides what "it" is made of.

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 (qq) — what this token is looking for. For it, something like "I am a pronoun; find me a noun I could refer to."
  • Key (kk) — the index card each token advertises. For animal, something like "I am an animate singular noun."
  • Value (vv) — 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:

Q=XWQ,K=XWK,V=XWVQ = XW^Q, \qquad K = XW^K, \qquad V = XW^V

where X∈RT×dmodelX \in \mathbb{R}^{T \times d_{model}} holds one row per token. The word self in self-attention means exactly this: QQ, KK and VV all derive from the same XX. 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:

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

We will compute it by hand for a three-token sequence with dk=2d_k = 2, so every number is checkable.

Take the tokens the, cat, sat and suppose the projections have produced:

TokenKey kjk_jValue vjv_j
the[1, 0][1,\ 0][2, 0][2,\ 0]
cat[0, 2][0,\ 2][0, 4][0,\ 4]
sat[1, 1][1,\ 1][1, 1][1,\ 1]

and the query for sat is q3=[1, 2]q_3 = [1,\ 2]. We compute the output for position 3.

Step 1 — relevance scores by dot product

Score every key against the query:

q3⋅k1=(1)(1)+(2)(0)=1q_3 \cdot k_1 = (1)(1) + (2)(0) = 1

q3⋅k2=(1)(0)+(2)(2)=4q_3 \cdot k_2 = (1)(0) + (2)(2) = 4
q3⋅k3=(1)(1)+(2)(1)=3q_3 \cdot k_3 = (1)(1) + (2)(1) = 3

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\sqrt{d_k}

With dk=2d_k = 2, dk=1.4142\sqrt{d_k} = 1.4142. The scaled scores are:

11.4142=0.7071,41.4142=2.8284,31.4142=2.1213\frac{1}{1.4142} = 0.7071, \qquad \frac{4}{1.4142} = 2.8284, \qquad \frac{3}{1.4142} = 2.1213

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:

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

Exponentiate:

e0.7071=2.0281,e2.8284=16.918,e2.1213=8.3418e^{0.7071} = 2.0281, \qquad e^{2.8284} = 16.918, \qquad e^{2.1213} = 8.3418

Sum: 2.0281+16.918+8.3418=27.2882.0281 + 16.918 + 8.3418 = 27.288. Divide:

Attends toScaled scoreese^{s}Weight α\alpha
the0.70712.02810.074
cat2.828416.9180.620
sat2.12138.34180.306
Total1.000

Check the sum: 0.074+0.620+0.306=1.0000.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

z3=0.074 [2,0]+0.620 [0,4]+0.306 [1,1]z_3 = 0.074\,[2,0] + 0.620\,[0,4] + 0.306\,[1,1]

Componentwise: first coordinate 0.074×2+0.620×0+0.306×1=0.148+0+0.306=0.4540.074 \times 2 + 0.620 \times 0 + 0.306 \times 1 = 0.148 + 0 + 0.306 = 0.454. Second coordinate 0+0.620×4+0.306×1=2.480+0.306=2.7860 + 0.620 \times 4 + 0.306 \times 1 = 2.480 + 0.306 = 2.786. So

z3=[0.454, 2.786]z_3 = [0.454,\ 2.786]

Compare that to the value of sat alone, [1,1][1, 1]. The output has been pulled strongly towards v2=[0,4]v_2 = [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⊤Q K^{\top} computes all T2T^2 scores in one matrix multiply.

Why dk\sqrt{d_k} is not cosmetic

Run the same example without the scaling. Raw scores 1, 4, 3:

e1=2.718,e4=54.598,e3=20.086,sum=77.402e^1 = 2.718, \quad e^4 = 54.598, \quad e^3 = 20.086, \quad \text{sum} = 77.402

Weights: 0.035, 0.705, 0.259. Compare side by side:

Attends toUnscaled weightScaled weight
the0.0350.074
cat0.7050.620
sat0.2590.306

At dk=2d_k = 2 the difference is mild. The problem is that it grows with dkd_k, and real models use dk=64d_k = 64 or 128.

Here is the argument. Suppose the components of qq and kk are independent with mean 0 and variance 1 — roughly what standard initialisation gives you. Then

q⋅k=∑i=1dkqikiq \cdot k = \sum_{i=1}^{d_k} q_i k_i

Each term qikiq_i k_i has mean E[qi]E[ki]=0\mathbb{E}[q_i]\mathbb{E}[k_i] = 0 and variance E[qi2]E[ki2]=1\mathbb{E}[q_i^2]\mathbb{E}[k_i^2] = 1. Summing dkd_k independent terms adds the variances:

Var(q⋅k)=dk,sd(q⋅k)=dk\mathrm{Var}(q \cdot k) = d_k, \qquad \mathrm{sd}(q \cdot k) = \sqrt{d_k}

With dk=64d_k = 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:

αsmall=e0e0+e24=11+2.649×1010=3.8×10−11\alpha_{\text{small}} = \frac{e^{0}}{e^{0} + e^{24}} = \frac{1}{1 + 2.649 \times 10^{10}} = 3.8 \times 10^{-11}

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

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

With α≈1−3.8×10−11\alpha \approx 1 - 3.8\times10^{-11}, that derivative is about 3.8×10−113.8 \times 10^{-11}. Multiply it into the backward pass and the gradient reaching WQW^Q and WKW^K through this attention weight is numerically zero. The layer stops learning — not slowly, but completely, and from step one.

Dividing by dk=8\sqrt{d_k} = 8 turns the gap of 24 into a gap of 3:

α=[e0e0+e3, e3e0+e3]=[121.086, 20.08621.086]=[0.047, 0.953]\alpha = \left[\frac{e^0}{e^0 + e^3},\ \frac{e^3}{e^0+e^3}\right] = \left[\frac{1}{21.086},\ \frac{20.086}{21.086}\right] = [0.047,\ 0.953]

Still confident, but α(1−α)=0.047×0.953=0.045\alpha(1-\alpha) = 0.047 \times 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\sqrt{d_k}, attention with dk=64d_k = 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 QQ, KK and VV come from and whether any scores are suppressed.

Self-attentionCross-attentionCausal (masked) self-attention
QQ fromsequence Asequence A (e.g. decoder)sequence A
KK, VV fromsequence Asequence B (e.g. encoder)sequence A
Attention matrix shapeTA×TAT_A \times T_ATA×TBT_A \times T_BTA×TAT_A \times T_A, lower-triangular
Each token seesall positionsall of the other sequenceitself and everything before it
Typical useBERT-style encoderstranslation, image captioningGPT-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 −∞-\infty.

Take position 3 of a four-token sequence with raw scaled scores [2.0, 1.0, 3.0, 5.0][2.0,\ 1.0,\ 3.0,\ 5.0]. Position 4 is in the future and must be masked:

s=[2.0, 1.0, 3.0, −∞]s = [2.0,\ 1.0,\ 3.0,\ -\infty]

Exponentiate: e2=7.389e^{2}=7.389, e1=2.718e^{1}=2.718, e3=20.086e^{3}=20.086, e−∞=0e^{-\infty}=0. Sum =30.193= 30.193. Weights:

[0.245, 0.090, 0.665, 0][0.245,\ 0.090,\ 0.665,\ 0]

which sums to 1.000 over the three legal positions. Now the unmasked version, for contrast: e5=148.413e^{5} = 148.413, total 178.606178.606, giving [0.041, 0.015, 0.112, 0.831][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 −∞-\infty (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.1691 - 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 modelsHow self-attention removes it
Position tt cannot be computed until t−1t-1 finishes, so the GPU idlesQK⊤QK^{\top} is one matrix multiply covering all T2T^2 pairs; the sequential depth is O(1)O(1) regardless of TT
Gradient between positions ii and jj passes through ∣i−j∣|i-j| Jacobians and decays exponentiallyPositions ii and jj are connected by a single weight αij\alpha_{ij} — a one-hop path at any distance
All history compressed into one fixed-size stateAll TT value vectors stay available; the summary is rebuilt per query

The cost is the T×TT \times T score matrix. At T=512T = 512 that is 262,144 entries per head per layer; at T=2048T = 2048 it is 4.2 million. The compute grows quadratically and so does the activation memory.

Implementation

Python
import torchimport torch.nn as nnimport torch.nn.functional as Fimport mathclass SelfAttention(nn.Module):    def __init__(self, d_model, d_k=None, dropout=0.1):        super().__init__()        self.d_k = d_k or d_model        self.W_q = nn.Linear(d_model, self.d_k, bias=False)        self.W_k = nn.Linear(d_model, self.d_k, bias=False)        self.W_v = nn.Linear(d_model, self.d_k, bias=False)        self.dropout = nn.Dropout(dropout)    def forward(self, x, mask=None):        # x: (batch, seq_len, d_model)        Q = self.W_q(x)                       # (B, T, d_k)        K = self.W_k(x)        V = self.W_v(x)        # (B, T, d_k) @ (B, d_k, T) -> (B, T, T)        scores = Q @ K.transpose(-2, -1) / math.sqrt(self.d_k)        if mask is not None:            # mask: True where the position is allowed            scores = scores.masked_fill(mask == 0, float('-inf'))        attn = F.softmax(scores, dim=-1)      # normalise over KEYS, the last axis        attn = self.dropout(attn)        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:

Python
def causal_mask(T, device=None):    # lower-triangular ones: position i may attend to j <= i    return torch.tril(torch.ones(T, T, dtype=torch.bool, device=device))m = causal_mask(4)# tensor([[ True, False, False, False],#         [ True,  True, False, False],#         [ True,  True,  True, False],#         [ 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.

Python
import matplotlib.pyplot as pltdef plot_attention(weights, tokens):    # weights: (T, T) for one head, one example    w = weights.detach().cpu().numpy()    fig, ax = plt.subplots(figsize=(6, 5))    im = ax.imshow(w, cmap='viridis', vmin=0, vmax=1)    ax.set_xticks(range(len(tokens)), tokens, rotation=45, ha='right')    ax.set_yticks(range(len(tokens)), tokens)    ax.set_xlabel('attending to (keys)')    ax.set_ylabel('query token')    fig.colorbar(im)    plt.tight_layout()    return fig

Read it row by row: row ii is where token ii looked, and it sums to 1. Three patterns show up constantly and each means something:

Pattern in the heatmapWhat it usually means
A bright diagonalTokens attending mostly to themselves — common in early layers, and normal
A bright vertical stripe on one columnAn 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/T1/TThe 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 dmodeld_{model} makes the projections more expensive but leaves the T×TT \times T attention matrix the same size. Doubling TT quadruples that matrix — both the compute and the memory holding it. Before scaling context, work out the activation memory: batch ×\times heads ×T2×\times T^2 \times bytes, per layer. At batch 32, 8 heads, T=1024T=1024, fp16, that is 32×8×10242×2=51232 \times 8 \times 1024^2 \times 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 ii is bit-identical when you change any token at position j>ij > 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.