32 Transformer ★
Context
By 2017 the best sequence models were recurrent (LSTM, #11) with an attention mechanism (#22). Recurrence is the bottleneck: the computation is strictly sequential in time (token t waits for t−1), so it parallelizes badly, it is slow on long inputs, and long-range dependencies still come hard. The authors ask an audacious question: is recurrence needed at all, if attention already connects any two positions directly?
The idea and the mechanism
Every token is projected into three vectors: Query (what I am looking for), Key (what I am), Value (what I carry). Attention is a retrieval of values, weighted by how well the query matches the keys:
The matrix QK⊤ gives a score for "how relevant token i is to token j"; softmax turns the scores into weights; multiplying by V collects a weighted blend of values. Every token aggregates information from all the others in a single step.
Multi-head. Instead of one attention, several parallel "heads", each with its own projections; one catches syntax, another coreference, and so on. The results are concatenated:
Positional encodings. Attention is permutation-invariant (it knows nothing of order), so positional signals are added to the embeddings (sinusoids of different frequencies). A transformer block: attention → residual + LayerNorm → position-wise FFN → residual + LayerNorm. The original is an encoder-decoder (for translation); later the decoder-only (GPT) and encoder-only (BERT) lines split apart.
probability Why we divide by √dk
Suppose the components of Q and K are independent random variables with mean 0 and variance 1 (roughly what you get after normalization). An attention score is the dot product of a query and a key of dimension dk:
The mean of such a sum is zero, and its variance is the sum of dk independent terms, each of which has Var(qi ki) = E[qi2]·E[ki2] = 1:
So as the dimension grows the scores swell like √ dk . Logits that are large in magnitude drive the softmax into "saturation" — it becomes almost one-hot, and its gradient at the inactive positions goes to zero (training stalls). Dividing by √ dk brings the variance of the scores back to 1, keeping the softmax in its sensitive range with live gradients. A tiny constant — but without it deep transformers would train poorly.
PyTorch Implementation: self-attention in ~12 lines
import torch
import torch.nn.functional as F
def attention(Q, K, V):
# Q, K, V: (..., seq, d_k)
d_k = Q.size(-1)
scores = Q @ K.transpose(-2, -1) / d_k ** 0.5 # scores QKᵀ / √d_k
weights = F.softmax(scores, dim=-1) # rows sum to 1
return weights @ V # weighted sum of values
def multi_head(x, Wq, Wk, Wv, Wo, h):
Q, K, V = x @ Wq, x @ Wk, x @ Wv # projections (seq, d_model)
split = lambda t: t.view(t.size(0), h, -1).transpose(0, 1)
out = attention(split(Q), split(K), split(V)) # h heads in parallel
out = out.transpose(0, 1).reshape(x.size(0), -1)
return out @ Wo # concatenate the heads
Why it matters
The Transformer is the foundational architecture of almost all modern AI: GPT, BERT, Claude, the Vision Transformer (#39), multimodal models. Its main engineering property is that it parallelizes across positions: there is no recurrent dependency, so the model sits beautifully on GPUs and TPUs and scales to long contexts and enormous data. "Transformer + scaling laws (#37)" is the whole LLM revolution.
The price is that attention costs O(N2) in the length of the sequence (everything against everything), which FlashAttention (#48) and the efficient variants will later set about fixing.
Connections
Attention was not invented here — Bahdanau (2014) added it to seq2seq as a way for the decoder to "look at" the source words it needed. The Transformer pushed the idea to its limit: it removed recurrence altogether and made self-attention the only mechanism connecting tokens. Self-attention = a sequence attending to itself.
The LSTM connected distant elements through a recurrent state — sequentially, and with difficulty over long distances. The Transformer connects any two positions directly, in one step, in parallel. The same class of problems (sequences), but the main limitation of recurrence is gone — speed and long-range dependencies.
Two great branches grow straight out of this: encoder-only BERT (bidirectional understanding) and decoder-only GPT (autoregressive generation). The whole modern LLM lineup is a scaled-up Transformer plus recipes for pre-training and alignment.
The quadratic cost of attention O(N2) and its appetite for memory are the bottleneck for long contexts. FlashAttention computes exactly the same attention, but without ever materialising the full N×N matrix, which is what makes long contexts practical.
Questions worth asking
If attention already connects all the tokens, what are the FFNs and all those layers for — isn't attention alone enough?
Attention only mixes and routes information between positions — in essence it takes linear combinations of values. All the per-token nonlinear processing, and most of the model's capacity, sit in the FFN (a two-layer MLP applied to each token separately).
And the stack of layers builds a hierarchy: early layers catch local, syntactic relations, later layers abstract ones. A single attention layer with no FFN is helpless: it has nothing with which to transform content nonlinearly, only to average what the other positions hold.
Is O(N²) fundamental, or can attention be made cheaper?
Full attention compares every pair of tokens — hence the N². But that is not sacred: there are linear and sparse approximations (Linformer, Performer, sparse attention) that give up accuracy for speed.
FlashAttention (#48) goes the other way — it keeps attention exact but removes the quadratic memory through tiling. For honest full attention the N² compute is unavoidable; the only question is whether you are willing to approximate.
How does the model know the order of the words, if attention is permutation-invariant?
On its own it does not; with no positional signals a transformer sees a "bag of tokens". Order is injected by positional encodings (sinusoids in the original).
This turned out to be a subtle spot: how you encode position affects generalization to long sequences — hence the evolution towards learned and, especially, rotary embeddings (RoPE), which encode relative positions and stretch further into a long context.
Why are several heads better than one big one of the same total dimension?
Different heads look into different subspaces at once and catch different kinds of relation: one agreement, another coreference, a third positional patterns. One big head would average it all into a single similarity mechanism.
Empirically the heads do specialize — although some of them can afterwards be cut away without harm (they are not all equally useful), which spawned a whole line of work on head pruning.
What to read in the original
Read it in full and re-read the sections on (multi-head) attention and positional encoding — this is the single most important write-up in the entire canon. It helps enormously to implement attention by hand once: it is literally a few matrix operations (projections → QKᵀ → scale → softmax → ×V).