
Book 14 of 50 · Free
Transformers Explained for Beginners
27,659 words · 17 chapters · illustrated

Book 14 of 50 · Free
27,659 words · 17 chapters · illustrated
Book 14 of 50 — AstolixGen Learning Series For researcher and publication students

The transformer is the single most important architecture in modern AI. It powers large language models, machine translation systems, image classifiers, speech recognizers, and protein-folding models. If you plan to do research in AI in the 2020s and beyond, you need to understand it — not as a magic box, but as a machine you could build yourself.
This book explains transformers from zero. We start with the problem they were invented to solve, build up the idea of attention with everyday analogies, walk through the mathematics one tiny step at a time (with a hand-worked numeric example you can check with a pocket calculator), and then assemble the full architecture. Along the way you will see real PyTorch code, learn how transformers are trained in practice, meet their famous cousins (BERT, GPT, Vision Transformers), and finish with a guided reading of the original paper, "Attention Is All You Need."
You do not need to be a mathematician. You need curiosity, basic Python, and the patience to work through small examples. Every chapter ends with a "For your research" box that connects the idea to publishable work you could actually do as an MS or early PhD student.
Learning objectives: - Explain why recurrent networks struggled with long sequences and parallel training - Describe attention using the query/key/value idea in plain words - Compute scaled dot-product self-attention by hand on a tiny example - Explain why multiple attention heads are better than one - Describe positional encoding and why transformers need it - Draw and explain the full encoder–decoder architecture from memory - List the key training choices that make transformers work (warmup, large batches, regularization) - Explain how Vision Transformers apply the same idea to images - Summarize efficient-attention ideas (Linformer, Performer, FlashAttention) at concept level - Use a pretrained transformer from Hugging Face for feature extraction and fine-tuning
Before transformers, there was a problem that haunted every researcher working with sequences — sentences, speech, time series, protein chains. The problem was simple to state and hard to solve: how does a machine remember what happened far back in a sequence while processing the next item?
A sequence is ordered data: word 1, word 2, word 3, … Each item's meaning depends on items that came before it. Consider the sentence:
"The trophy didn't fit in the suitcase because it was too big."
What does "it" refer to? The trophy, obviously — trophies can be too big. Now change one word:
"The trophy didn't fit in the suitcase because it was too small."
Now "it" refers to the suitcase. A single word at the end of the sentence flips the meaning of a pronoun that appeared earlier. To understand language, a model must carry information across long distances inside a sequence. This is called long-range dependency, and it is the central difficulty of sequence modeling.
The classic solution was the recurrent neural network (RNN). An RNN reads a sequence one item at a time, left to right, like a person reading a book. At each step it updates a hidden state — a vector of numbers that acts as its running memory of everything read so far.
Imagine reading a novel and keeping notes on a single index card. After each page, you must erase your old notes and rewrite a compressed summary that includes the new page. That index card is the hidden state. The RNN's rule for updating the card is learned from data.
The most successful RNN variants were the LSTM (long short-term memory, Hochreiter and Schmidhuber, 1997) and the GRU. They added gates — little learned switches that decide what to keep, what to forget, and what to add to the memory card. LSTMs were genuinely good: they powered the first strong machine translation and speech recognition systems of the 2010s.
But the index card has a fixed size. Suppose you translate a 50-word German sentence into English. The encoder RNN reads all 50 words and compresses everything into one final hidden state — a vector of, say, 1,000 numbers. Then the decoder must produce the entire English translation from that single vector.
Researchers noticed a pattern: translation quality was fine for short sentences but degraded badly for long ones. All the meaning of a long sentence simply could not fit through one fixed-size vector. This was called the information bottleneck. It was like asking someone to read a whole chapter and then summarize it from memory without looking back at the pages. The details are lost.
In 2014, Cho and colleagues formalized this encoder–decoder (sequence-to-sequence) setup for translation, and the bottleneck became visible in every experiment: longer input, worse output.
The second problem was about speed, and it was arguably worse. An RNN processes word 2 only after finishing word 1, word 3 only after word 2, and so on. This is inherently sequential.
Modern AI training runs on GPUs, which are fast precisely because they do thousands of operations in parallel. But an RNN starves the GPU: at each step, only a small computation is available, and the next step cannot start until the current one finishes. Training on long sequences was painfully slow. Researchers wanted to train on more data — everyone suspected bigger data would give better models — but the sequential bottleneck made that impractical.
Think of it as an assembly line where each worker must wait for the previous worker to finish the entire car before starting. You cannot hire more workers to speed it up, because the work itself is ordered.
The third problem was about learning. RNNs are trained with backpropagation through time: the error signal flows backward from the end of the sequence to the beginning, passing through every time step. Each step multiplies the signal by numbers that are usually smaller than one. After 50 steps, the signal reaching the early words is essentially zero — it has vanished.
The practical effect: the model could not learn connections between distant words. It learned that "not" negates the next word, but struggled to learn that a pronoun at position 40 refers to a noun at position 5. LSTMs and GRUs reduced this problem with their gates, but did not eliminate it.
In 2015, Bahdanau and colleagues published "Neural Machine Translation by Jointly Learning to Align and Translate," which introduced attention into the RNN encoder–decoder. The idea was beautiful and simple: instead of forcing the decoder to work from one final memory vector, let it look back at all the encoder's hidden states at every decoding step, and decide which ones matter right now.
When translating the French word that corresponds to "suitcase," the decoder could put most of its weight on the encoder state for "suitcase" and ignore the rest. Attention was a soft, learnable "looking back" mechanism. It helped with long sentences immediately, and translation quality jumped.
But notice: Bahdanau's attention was still bolted onto an RNN. The encoder still read one word at a time. The sequential bottleneck and the slow training remained.
This set the stage for the radical question that the transformer answered:
What if we keep the attention — the part that worked — and throw away the recurrence entirely?
What if, instead of reading one word at a time and maintaining a fragile memory card, the model looked at the entire sequence at once, and let every word directly gather information from every other word? No assembly line. No vanishing signal across time steps. Every connection is direct, and every computation can run in parallel on a GPU.
That is the transformer. The rest of this book builds it piece by piece.
"Gradients vanish" sounds abstract until you put numbers on it. During backpropagation through time, the error signal passes through one matrix multiplication per time step. Suppose each step shrinks the signal by a factor of 0.9 (typical for a well-behaved RNN). After 10 steps: 0.9^10 ≈ 0.35 — a third of the signal survives. After 50 steps: 0.9^50 ≈ 0.005 — half a percent. After 100 steps: 0.9^100 ≈ 0.00003 — effectively zero.
So a word at the start of a long paragraph contributes almost nothing to the learning signal by the time the error travels back to it. The model cannot learn to connect it to the end of the paragraph, no matter how much data you show it. LSTMs fought this with gates that let the signal skip steps unchanged (a factor close to 1.0 instead of 0.9), which is why they handled hundreds of steps while plain RNNs managed dozens. But gates were a patch, not a cure: the signal still traveled step by step, and very long dependencies still faded.
The transformer's answer is structural, not incremental: with self-attention, the path between any two tokens is a single step — O(1) path length, as the original paper's famous Table 1 puts it. No repeated multiplication, no decay. The gradient from the loss flows directly to every token in one hop.
A fair question: convolutional networks process whole inputs in parallel — why didn't they replace RNNs for sequences first? Researchers did try (notably with models like WaveNet for audio). The problem is the receptive field: a convolutional layer only sees a fixed local window (say, 3–5 tokens). To connect tokens 50 apart, you must stack many layers so the windows compose — depth grows with distance. Dilated convolutions stretched the windows, but the architecture still assumed locality: nearby tokens matter most. Language routinely violates that (our trophy/suitcase pronoun jumps the whole sentence). CNNs were parallel but myopic; RNNs were farsighted but sequential. Attention is both parallel and farsighted — that combination is what made it win.
Notice the pattern: each step removed a limitation while keeping the previous step's gains. Attention kept the look-back; the transformer kept the attention and removed the recurrence. Scientific progress here looks like sculpting — cutting away what constrains, keeping what works.
Take a 40-word product review and compress it into 10 numbers (a tiny hidden state), then hand those 10 numbers to a friend and ask them to reconstruct whether the review was positive, what product it was about, and what the main complaint was. They'll fail — 10 numbers can't hold it. Now allow 1,000 numbers: better, but the 40-word review has roughly 40 × (rich meaning each) worth of content, and your friend must answer questions you didn't know in advance. That is the fixed-vector bottleneck: the encoder must guess, ahead of time, which details the decoder will need. Attention abolished the guessing: the decoder looks back at the full encoded sequence and picks what it needs, when it needs it.
One of the cleverest pre-attention hacks came from Sutskever et al. (2014), who trained LSTM encoder–decoders on translation. They found that reversing the source sentence ("A B C" → "C B A") improved results significantly. Why? In "A B C → X Y Z" order, the first source word A is farthest from the first target word X in processing time — the signal must survive the whole encoding. Reversed, C (now first in) sits right next to X — short-term dependencies form easily, and the optimization gets a foothold before tackling long ones. It's a beautiful example of researchers working around the sequential bottleneck with data tricks instead of architectural change — and of how much pain that bottleneck caused.
The two great RNN variants differed in gating style: the LSTM (1997) keeps a separate cell state with input, forget, and output gates; the GRU (Cho et al., 2014) merges the cell and hidden state with reset and update gates — fewer parameters, similar performance on most tasks. Practitioners treated them as interchangeable workhorses. Both hit the same three walls (bottleneck, sequentiality, vanishing gradients); the GRU was cheaper per step but no more parallelizable. When the transformer arrived, it replaced both at once — which is why 2017–2019 saw entire codebases rewritten.
Before attention, researchers plotted translation quality (BLEU) against input sentence length and saw the same depressing curve everywhere: quality held up to ~20 words, then slid steadily downward. Longer sentence → more meaning crammed into the same fixed vector → worse output. Bahdanau's attention paper made its name partly by flattening that curve — with attention, long sentences degraded far less. When you evaluate any sequence model today, plot metric-vs-length first: it's the fastest diagnostic for whether your model truly handles long range or just memorizes short patterns.
The bottleneck was obvious, so people attacked it: bigger hidden states (diminishing returns — the vector grew but the compression problem didn't shrink), hierarchical encoders (sentence-level RNNs over word-level RNNs — helped document structure, not the core issue), bidirectional encoders (concatenating forward and backward passes — richer representations, same fixed-size output). Each helped at the margins. Attention was the first fix that changed the information flow rather than the capacity — the decoder no longer depended on one vector at all. Lesson for your research taste: capacity fixes scale costs; flow fixes change what's possible.
For your research: The history in this chapter is a template for how architectures are born: identify a concrete bottleneck (fixed memory, no parallelism, vanishing gradients), keep what works (attention), and remove what doesn't (recurrence). When you read any new architecture paper, ask: what bottleneck is it attacking? That single question will organize your literature review. Many publishable student projects are exactly this — taking a known bottleneck in a niche domain (Urdu text, medical time series, crop sensor data) and testing whether a transformer variant relieves it.
Key takeaways: - Sequences are hard because meaning depends on distant context (long-range dependencies). - RNNs read one item at a time into a fixed-size hidden state — a memory bottleneck. - RNNs cannot be parallelized across time steps, which wastes GPU power and slows training. - Gradients vanish over long sequences, so RNNs struggle to learn distant relationships. - Attention (2015) let decoders look back at all encoder states — the breakthrough the transformer kept, while discarding recurrence.
Attention is the one idea in this book you must understand deeply, because everything else is engineering around it. The good news: the intuition is something you already do every day.
Imagine you walk into a huge library and ask the librarian: "I need a beginner book on transformers." The librarian does three things:
Notice the separation: the key is what you match against (title, topic, level), and the value is what you actually receive (the book's content). Your query matched best with one card's key, so you received that card's value.
Now imagine a softer version. Instead of picking one book, the librarian brings you a blend: 70% of the beginner book, 20% of the visual guide, 10% of the math-heavy text — weighted by how well each key matched your query. That weighted blend is attention.
You are reading the sentence: "The animal didn't cross the street because it was too tired." You want to understand "it." Your brain automatically shines a spotlight back over the sentence: high weight on "animal," low weight on "street," almost nothing on "the." You then blend the meanings of the highlighted words to interpret "it."
That is self-attention: each word asks, "which other words should I pay attention to in order to understand myself in this context?" and then mixes their meanings accordingly.
In a transformer, every token (roughly, every word piece) plays three roles:
The process for one token: 1. Compare its query against every token's key (including its own). This gives a score for each pair — how relevant is that token to me? 2. Turn the scores into weights that sum to 1 (using softmax — more in Chapter 3). 3. Output a weighted blend of everyone's value vectors, using those weights.
The word "it" ends up as a mixture: mostly the meaning of "animal," a little of "street," a trace of everything else. Its representation is now context-aware. The same word "it" in a different sentence would blend different neighbors and get a different representation. That context-sensitivity is the superpower of transformers — and the reason fixed word vectors (like old word2vec embeddings) were left behind.
Beginners often ask: why not just compare words directly to words? Why the Q/K/V split?
Because a word plays different roles in different relationships. Consider "bank" in "the bank of the river." When "river" is figuring out its context, it queries for geographical things near me — and "bank" should match strongly as a key. But the value of "bank" (its meaning contribution) is different from its key (its matchability). Learned separately, the model can make keys about relevance and values about content. In practice, the model learns three different linear projections of each token's embedding, and training decides what each projection emphasizes.
A helpful mental image: the query is a question, the key is a label on a folder, the value is the folder's contents. You match questions to labels, but you read contents.
Computer scientists will recognize this pattern: it is a differentiable dictionary. In a normal dictionary, you look up one key and get one value. In attention, you look up with a query, get a similarity score against every key, and receive a weighted average of all values. Because every step is smooth (no hard choices), gradients flow through it, and the whole thing can be learned with backpropagation.
This "soft lookup" idea appears everywhere once you see it: database queries, search engines, human memory. The transformer just made it learnable and massively parallel.
Sentence: "The cat sat on the mat." We process the word "sat."
The representation of "sat" now contains its subject and location. A later layer can use this enriched representation to answer questions like "who sat?" without searching the sentence again.
Two common misconceptions to clear up early:
Attention is not the model "understanding" in a human sense. It is weighted averaging with learned weights. The magic is that weighted averaging, stacked in deep layers with learned projections, turns out to be an extremely expressive operation.
Attention weights are not always explanations. Researchers once hoped that looking at attention weights would reveal why a model decided something. Sometimes the weights are informative; often they are not — different weight patterns can produce the same output, and later layers transform everything anyway. Treat attention maps as clues, not proof. (This is itself an active research area — see Chapter 11.)
An alternative design would be hard attention: pick the single most relevant token and take only its value (like the librarian handing you exactly one book). This is intuitive but has a fatal flaw for learning: "pick the best" is a discrete, non-differentiable choice — you can't compute a gradient through an argmax, so backpropagation can't tune how the choice is made. Soft attention's weighted blend is differentiable end to end: every weight shifts a little with every training example, and the model gradually learns better matchings. The price is computation (you blend everything instead of picking one), but GPUs are built for exactly this kind of dense arithmetic. Differentiability beats discreteness — a recurring theme in deep learning.
Let's put numbers on the "Sara bought the car" example from earlier, for the word "she" under a hypothetical coreference head. Suppose the scaled scores are: Sara: 2.0, bought: 0.5, the: −1.0, car: 0.0, because: −0.5, she: 1.0, liked: 0.3, it: 0.8.
Softmax: e^2.0 ≈ 7.39, e^0.5 ≈ 1.65, e^−1.0 ≈ 0.37, e^0 ≈ 1.0, e^−0.5 ≈ 0.61, e^1.0 ≈ 2.72, e^0.3 ≈ 1.35, e^0.8 ≈ 2.23. Sum ≈ 17.32.
Weights: Sara ≈ 0.43, bought ≈ 0.10, the ≈ 0.02, car ≈ 0.06, because ≈ 0.04, she ≈ 0.16, liked ≈ 0.08, it ≈ 0.13. Check the sum: 0.43+0.10+0.02+0.06+0.04+0.16+0.08+0.13 = 1.02 (rounding). The head puts 43% of its mass on "Sara" — the correct referent — while keeping small weights everywhere else so gradients still flow. Notice "she" attends to itself with 0.16: self-attendance is common and healthy; a token's own value is often relevant to its meaning.
Only loosely. Human visual attention moves a spotlight: you look at one thing, and the rest blurs. Transformer attention never looks away from anything — it just turns volumes up and down. A better human analogy is a meeting where everyone speaks at once and you distribute your listening: 43% on Sara, 16% on yourself, a little on everyone else. The name "attention" is a metaphor, not a claim about brains. Don't let the metaphor do your reasoning for you — when in doubt, return to "weighted blend with learned weights."
Here's a viewpoint that connects transformers to another whole field: think of tokens as nodes in a fully connected graph, and attention as message passing — each node collects messages from all others, weighted by learned relevance. That's exactly the setup of graph neural networks (GNNs), specialized to the complete graph. This isn't just poetry: techniques transfer. Sparse attention (Chapter 9) is literally choosing a sparser graph; some researchers design the graph from domain structure (dependency trees for language, spatial adjacency for images) instead of using the complete graph. If you know GNNs, you already know half of attention theory — and if you don't, attention is a gentle on-ramp to them.
So far we've done self-attention: queries, keys, and values all come from the same sequence. But the machinery works across sequences too. In cross-attention, the queries come from sequence A (e.g., the decoder's English-so-far) while keys and values come from sequence B (e.g., the encoder's German representations). The English word being generated asks, "which German words are relevant to me right now?" and blends their values. Self-attention builds context within a sequence; cross-attention connects two sequences. The formula is identical — only the sources of Q vs. K/V differ. (You'll meet cross-attention properly in Chapter 6's decoder.)
Softmax has a hidden dial: temperature T. Compute softmax(scores / T): as T → 0, the weights collapse toward winner-take-all (hard attention); as T → ∞, they flatten toward uniform. T = 1 is standard attention. Temperature appears in two important places: knowledge distillation (Chapter 8's DeiT uses high-temperature teacher outputs to reveal "dark knowledge" about runner-up classes) and text generation (sampling with T < 1 makes output focused and deterministic; T > 1 makes it diverse and risky). It's the same mathematics as our attention weights — one more reminder that softmax is doing real work in the formula, not just normalizing.
For your research: The query/key/value framing is a design pattern you can reuse. Any time your problem involves "given X, find the relevant pieces of Y and blend them," you are looking at an attention-shaped problem: matching patient symptoms (query) against a knowledge base of cases (keys) to blend treatment notes (values); matching a satellite image patch (query) against historical patches (keys) to blend crop-yield records (values). Framing your problem this way in a paper's introduction instantly signals to reviewers that you understand the mechanism, not just the library call.
Key takeaways: - Attention = soft, weighted lookup: a query is matched against keys, and the matching values are blended. - Query asks, key advertises, value delivers. Matching happens between queries and keys; blending happens over values. - Self-attention applies this inside one sequence: every word gathers context from every other word. - The result is context-aware representation: the same word gets a different vector in different sentences. - Attention is weighted averaging with learned weights — simple, differentiable, and parallelizable.
This is the chapter many students fear and then, afterward, wonder what the fuss was about. The mathematics of self-attention is one matrix multiplication, one scaling, one softmax, and one more matrix multiplication. We will do all of it by hand on a tiny example, and you will check every number yourself.
The entire self-attention operation is this:
Attention(Q, K, V) = softmax(QK^T / √d_k) · V
Read it in four steps: 1. QK^T — compare every query against every key (dot products = similarity scores). 2. / √d_k — scale the scores down (d_k is the key dimension; scaling keeps numbers in a range where softmax behaves well). 3. softmax(...) — turn each row of scores into weights that sum to 1. 4. · V — blend the value vectors using those weights.
That's it. The rest of this chapter is understanding why each step exists, then computing it by hand.
We have queries Q (one row per token asking) and keys K (one row per token advertising). The matrix product QK^T computes every query·key dot product at once. Entry (i, j) of the result = "how much does token i's query match token j's key?"
The dot product is a natural similarity measure: it is large when two vectors point in the same direction, small or negative when they point in different directions. If the query of "it" points roughly the same way as the key of "animal," their dot product is large — a high relevance score.
Softmax has a quirk: if its inputs are large (say, 20 vs. 10), the output becomes almost one-hot — 1.0 for the largest, ~0 for the rest — and its gradients become tiny (the softmax saturation problem). Dot products grow with dimension: for random vectors of dimension d_k, the dot product has variance d_k. Dividing by √d_k brings the variance back to 1, keeping scores in the gentle range where softmax produces soft, learnable weights. It is a small detail with a large practical effect — the original paper's authors added it after finding that unscaled attention trained poorly at larger dimensions.
Softmax takes a row of scores [s₁, s₂, s₃] and returns [e^s¹/Σ, e^s²/Σ, e^s³/Σ] — positive numbers summing to 1. Big scores become big weights, but (unlike a hard maximum) every token keeps some weight, so gradients flow to everything. This softness is what makes attention differentiable.
Multiply the weight matrix by V. Row i of the output = Σⱼ weight(i,j) · Vⱼ — the weighted blend we described in Chapter 2. Note that the values V are blended, but the queries and keys are not: Q and K only decide the weights.
Let's compute everything for 3 tokens with tiny 2-dimensional vectors. d_k = 2, so √d_k ≈ 1.4142.
Queries, keys, values (rows = tokens 1, 2, 3):
Compute the output for token 1 only (one row keeps the arithmetic checkable; the same steps repeat for every row).
Step 1 — scores for q₁ = [1, 0]: - q₁·k₁ = 1×1 + 0×0 = 1 - q₁·k₂ = 1×0 + 0×1 = 0 - q₁·k₃ = 1×1 + 0×1 = 1
Step 2 — scale by √2 ≈ 1.4142: - 1 / 1.4142 ≈ 0.7071, 0 / 1.4142 = 0, 1 / 1.4142 ≈ 0.7071 - Scaled scores: [0.7071, 0, 0.7071]
Step 3 — softmax: - e^0.7071 ≈ 2.0281, e^0 = 1.0, e^0.7071 ≈ 2.0281 - Sum = 2.0281 + 1.0 + 2.0821… wait, let me add carefully: 2.0281 + 1.0000 + 2.0281 = 5.0562 - Weights: 2.0281/5.0562 ≈ 0.4011, 1.0000/5.0562 ≈ 0.1978, 2.0281/5.0562 ≈ 0.4011 - Check: 0.4011 + 0.1978 + 0.4011 = 1.0000 ✓
Step 4 — blend the values: - output₁ = 0.4011 × [4, 0] + 0.1978 × [0, 6] + 0.4011 × [2, 2] - = [1.6044, 0] + [0, 1.1868] + [0.8022, 0.8022] - = [2.41, 1.99] (rounded)
Token 1's new representation, [2.41, 1.99], is mostly a mix of v₁ and v₃ (each ~40% weight) with a smaller contribution from v₂ (~20%). Notice how the math exactly implements the library analogy: token 1's query matched keys 1 and 3 best, so it received mostly values 1 and 3.
Exercise for your calculator: compute the row for q₂ = [0, 1]. You should get scaled scores [0, 0.7071, 0.7071], the same weights shifted, and output₂ ≈ [0.80, 3.99]. (Solution in Exercise 1 at the end of the book.)
In a decoder that generates text left to right, token 5 must not peek at token 6 — that would be cheating during training (the model would just copy the answer). Masking enforces this: before softmax, we set the scores of forbidden positions to negative infinity. e^(−∞) = 0, so those tokens get exactly zero weight. The result is causal (left-to-right) attention — each token only blends information from itself and earlier tokens. Encoder self-attention typically uses no mask: every token sees everything.
If the sequence has n tokens, the score matrix QK^T is n × n. Computing it costs O(n²·d) time and O(n²) memory. For n = 512 this is fine; for n = 100,000 it is catastrophic. This quadratic cost is the transformer's main weakness and the motivation for every efficient variant in Chapter 9. Remember it — it shapes all transformer research.
Here is the formula as code. Read it line by line against the four steps above:
import torch
import torch.nn.functional as F
def self_attention(Q, K, V, mask=None):
d_k = Q.size(-1)
scores = Q @ K.transpose(-2, -1) / (d_k ** 0.5) # steps 1-2
if mask is not None:
scores = scores.masked_fill(mask == 0, float('-inf'))
weights = F.softmax(scores, dim=-1) # step 3
return weights @ V # step 4
# tiny demo with our hand-computed example (token 1)
Q = torch.tensor([[1., 0.]])
K = torch.tensor([[1., 0.], [0., 1.], [1., 1.]])
V = torch.tensor([[4., 0.], [0., 6.], [2., 2.]])
print(self_attention(Q, K, V))
# tensor([[2.4069, 1.9890]]) -- matches our hand computation [2.41, 1.99]
Eight lines. The whole revolution, in eight lines. Everything else in the transformer architecture is scaffolding that makes this operation deep, stable, and trainable.
Figure 1: Self-attention — every token connects to every other token, with connection strength set by query–key matching.
Our hand example used small scores, but real models produce scores in the tens or hundreds, and e^100 overflows floating-point arithmetic (it exceeds ~10^43). Libraries never compute softmax naively. The trick: subtract the row's maximum score before exponentiating. Softmax is shift-invariant — softmax([3, 1]) = softmax([103, 101]) — because multiplying every e^s by the same constant e^−max cancels in the division. So implementations compute softmax(s − max(s)): the largest input becomes e^0 = 1, and nothing overflows. PyTorch's F.softmax does this internally. You don't need to implement it, but you must know it exists — otherwise the first time you hand-roll softmax in NumPy for a debugging script, you'll get inf and nan and blame the model.
Let's apply a causal mask to our Chapter 3 example and recompute token 2's row. Recall q₂ = [0, 1] gave unmasked scaled scores [0, 0.7071, 0.7071] over tokens 1, 2, 3. With a causal mask, token 2 may see tokens 1–2 but not 3: set score(2,3) = −∞.
Compare with the unmasked result (≈ [1.59, 3.21] from Exercise 1's solution): masking removed token 3's contribution entirely and rebalanced the rest. During training, this guarantees the model predicts token 3's word using only tokens 1–2 — no peeking at the answer.
Softmax has three properties the job requires: outputs are positive (weights can't be negative), they sum to 1 (a proper blending), and it's smooth and differentiable everywhere. Alternatives exist — sparsemax produces sparse weights (exact zeros for losers), and some efficient variants use other normalizations — but softmax's smoothness makes optimization well-behaved, and its exponential sharpens differences just enough: a score gap of 2 becomes a ~7:1 weight ratio, letting the model be decisive without going one-hot. When you see a new paper replace softmax, ask which of the three properties it keeps and what it gains in return.
The original transformer applies dropout (rate 0.1) directly to the attention weights after softmax: randomly zero 10% of the weights and rescale the rest. Why drop weights rather than just activations? It prevents the model from leaning on a single token relationship — if the "it → animal" link is randomly cut during training, the model must also develop backup evidence paths. At test time dropout is off and the full blend is used. If your model's attention collapses onto one token per head (over-sharp, brittle behavior), increasing attention dropout is a targeted remedy.
Let's finish the whole example — all three rows — so you see the complete weight matrix. We have row 1 (token 1): weights [0.4011, 0.1978, 0.4011], output [2.41, 1.99]. Exercise 1 gives row 2 (token 2): weights [0.1978, 0.4011, 0.4011], output [1.59, 3.21]. Now row 3, q₃ = [1, 1]:
The complete attention weight matrix:
| attends to 1 | attends to 2 | attends to 3 | |
|---|---|---|---|
| token 1 | 0.40 | 0.20 | 0.40 |
| token 2 | 0.20 | 0.40 | 0.40 |
| token 3 | 0.25 | 0.25 | 0.50 |
Read row 3: token 3's query matched key 3 best (score 2 — the only self-match with both components aligned), so it keeps 50% of its own value and borrows 25% from each other token. Notice the matrix isn't symmetric — attention is directional: token 1 gives token 3 weight 0.40, but token 3 gives token 1 only 0.25. Each row answers "what do I need?" independently.
Our example was one head, one sequence. Production code adds two dimensions: batch (many sequences at once) and heads (Chapter 4). Tensors flow as (B, H, N, d_k): scores are (B, H, N, N), weights (B, H, N, N), output (B, H, N, d_k). The matrix multiplications broadcast over batch and heads automatically — that's why the reshape trick in Chapter 4's code works. When debugging shape errors, say the four dimensions aloud: "batch 8, heads 12, 64 tokens, 64 dims per head." Shape errors are the most common beginner bug; naming dimensions is the cure.
How does learning adjust the weights? Picture the loss telling output₁ "you should have been higher in dimension 2." That pressure flows backward along two paths: (1) to the values — v₁, v₂, v₃ get nudged in proportion to their weights (0.40, 0.20, 0.40): heavily-used values learn fastest; (2) to the queries and keys through the softmax — the q₁·k₁ and q₁·k₃ scores get pushed up (they were useful), q₁·k₂ pushed down. Over thousands of examples, queries learn to ask for what's useful and keys learn to advertise what's useful. The values carry the content; Q/K carry the routing. This division of labor is why the Q/K/V split (Chapter 2's question) was worth it.
For your research: Being able to hand-compute attention is a superpower when debugging. When your model's attention maps look wrong, you can construct a 3-token toy input with known Q/K/V, run it through your code, and compare against a calculator result. This "tiny verifiable example" habit is also excellent paper material: reviewers love a minimal example that isolates a claimed effect. Keep a notebook of 2–3 such toy cases for any attention variant you invent.
Key takeaways: - Self-attention = softmax(QK^T/√d_k)·V: score, scale, normalize, blend. - Dot products measure query–key similarity; √d_k scaling keeps softmax in its healthy range. - Softmax converts scores to weights summing to 1; the output is a weighted blend of values. - Masking (setting scores to −∞) enforces left-to-right generation in decoders. - Cost is O(n²): every token attends to every token — powerful, but expensive for long sequences.
Chapter 3 computed attention with a single set of queries, keys, and values. Real transformers use multi-head attention: they run the attention operation several times in parallel, each time with different learned projections, then combine the results. This chapter explains why.
A single attention head produces one set of weights per token — one "opinion" about which other tokens matter. But words relate to each other in many ways simultaneously. In "The animal didn't cross the street because it was too tired":
One weighted blend cannot easily capture all of these at once. It's like asking one person to simultaneously judge a singing contest on pitch, emotion, and stage presence with a single score — you'd rather have three judges, each specializing.
Multi-head attention creates h independent "judges" (heads). Each head has its own learned projection matrices W^Q, W^K, W^V that transform the input into that head's private query/key/value space. Each head then runs full self-attention in its smaller space. Finally, the heads' outputs are concatenated and passed through one more learned projection W^O.
In the original transformer: d_model = 512, h = 8 heads, each head works in d_k = d_v = 64 dimensions (512/8). The total computation is roughly the same as one 512-dimensional head — splitting into heads costs almost nothing extra, because the per-head dimensions shrink proportionally.
This is one of the most delightful findings in transformer research. When researchers visualized the attention weights of trained models, different heads had clearly specialized:
Nobody programmed these roles. They emerged from training because the task (predicting the next word, translating a sentence) is solved better when different heads track different relationships. This emergent specialization is strong evidence that multi-head attention is doing real linguistic work, not just adding parameters.
Sentence: "Sara bought the car because she liked it." Two heads process the word "she":
Concatenated, the representation of "she" simultaneously knows who she is and what she did. A single head would have to compromise between these two jobs.
import torch
import torch.nn as nn
import torch.nn.functional as F
class MultiHeadAttention(nn.Module):
def __init__(self, d_model=512, n_heads=8):
super().__init__()
assert d_model % n_heads == 0
self.d_k = d_model // n_heads
self.n_heads = n_heads
# one linear layer per role; heads are handled 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)
def forward(self, x, mask=None):
B, N, _ = x.shape # batch, sequence length
# project, then split the last dim into (heads, d_k)
Q = self.W_q(x).view(B, N, self.n_heads, self.d_k).transpose(1, 2)
K = self.W_k(x).view(B, N, self.n_heads, self.d_k).transpose(1, 2)
V = self.W_v(x).view(B, N, self.n_heads, self.d_k).transpose(1, 2)
scores = Q @ K.transpose(-2, -1) / (self.d_k ** 0.5)
if mask is not None:
scores = scores.masked_fill(mask == 0, float('-inf'))
weights = F.softmax(scores, dim=-1)
out = weights @ V # (B, heads, N, d_k)
out = out.transpose(1, 2).reshape(B, N, -1) # concat heads
return self.W_o(out)
The trick is the reshape: instead of h separate small linear layers, we use one big projection and split its output into heads. Mathematically identical, much friendlier to GPUs.
A common student mistake is assuming more heads = better model. Heads help up to a point; beyond it, you get redundant heads that learn the same pattern. Some research even prunes heads after training with little quality loss — evidence that many heads are backup singers, not soloists.
Two reasons. First, specialization: separate projection matrices let each head live in its own learned subspace, tracking different relationships. Second, and subtler: a single softmax must allocate its one unit of weight mass across all relationships, forcing compromise. Eight softmaxes give eight independent allocations — the model can attend to "Sara" for coreference and "bought" for verb structure at the same time, in different heads.
Figure 2: Multi-head attention — the input splits into parallel heads, each learning a different relationship pattern, then merges.
Let's count parameters for the original configuration (d_model = 512, h = 8, d_k = d_v = 64) to see why "8 heads cost the same as 1":
A single 512-dimensional head would need W^Q, W^K, W^V at 512×512 each = 786,432 + 262,144 output = ~1.05M — identical. Splitting changes the structure (8 independent softmaxes in 64-D subspaces vs. 1 softmax in 512-D) without changing the budget. This is a deep design lesson: representational power often comes from how parameters are organized, not how many there are.
For scale: BERT-base (d_model=768, h=12) uses ~7M attention parameters per layer across 12 layers; the FFNs add roughly twice that. When you read "110M parameters," about two-thirds sit in the feed-forward networks.
Head roles aren't just diverse — they're layered. Studies of trained transformers consistently find:
This hierarchy mirrors the classic NLP pipeline (syntax before semantics before task reasoning) — except nobody programmed the pipeline; depth plus attention discovered it. For your interpretability projects, always report which layer a head lives in; a "coreference head" in layer 2 would be surprising, in layer 7 expected.
Revisit "she" in "Sara bought the car because she liked it" with concrete numbers for two heads (weights over [Sara, bought, the, car, because, she, liked, it]):
Head A's output ≈ 0.72·v_Sara + …; Head B's ≈ 0.58·v_bought + …. Concatenated, "she" carries both facts in separate subspaces, and the final W^O projection learns how to combine them for downstream layers. Now imagine forcing one head to do both: its softmax would have to split mass between Sara and bought, diluting both signals — the compromise a single head can't avoid.
If heads specialize, are all of them needed? Michel et al. (2019) showed you can prune (remove) a large fraction of heads after training with minimal quality loss — many heads are redundant. But there's subtlety: which heads are prunable varies by task, and a head that's useless for translation might matter for question answering. Practical takeaways: (1) head count is a capacity knob with diminishing returns; (2) pruning is a legitimate compression technique for deployment; (3) when analyzing "what the model learned," focus on the heads that survive pruning — they're the load-bearing ones. A student project replicating head pruning on a new domain (with before/after quality curves) is a clean, honest paper.
Once you're comfortable, the whole operation compresses beautifully with Einstein summation notation:
import torch
# Q,K,V: (..., seq, d_k); the "..." covers batch and heads
scores = torch.einsum('...id,...jd->...ij', Q, K) / d_k**0.5
weights = torch.softmax(scores, dim=-1)
out = torch.einsum('...ij,...jd->...id', weights, V)
'...id,...jd->...ij': for each batch/head, output[i,j] = Σ_d Q[i,d]·K[j,d] — exactly QK^T. Many research codebases use einsum because the subscript labels are the documentation: you can read the shapes off the letters. Learn to read einsum — it pays off across all of deep learning, not just transformers.
Vaswani et al. didn't just assert multi-head was better — they ablated it. Reducing to a single head (with proportionally larger d_k to match parameters) degraded translation quality, confirming that the structure of independent heads matters beyond the parameter count. They also varied d_k and found 64 near-optimal for their setup. The meta-lesson: when you propose an architectural choice, the reviewers' first question is "did you try the simpler version?" — and the paper that survives is the one that already ran that experiment. Budget ablation compute into every project plan.
"Different heads learn different things" can be measured, not just eyeballed. Common metrics: average attention distance per head (how far, on average, a head looks — local vs. global), entropy of the weight distribution (focused vs. diffuse), and cosine similarity between heads' attention maps (low similarity = diverse roles). Studies using these found exactly the layered pattern described above — and also found some heads with near-identical maps (the redundancy behind head pruning). If you run an interpretability project, report one such metric alongside your visualizations: numbers make the pictures credible.
The original paper's d_k = 64 wasn't arbitrary — their ablations tested larger and smaller values. Too small (e.g., 16): each head's matching space is cramped; distinct relationships blur together. Too large (e.g., 256): dot products grow, the √d_k scaling works harder, and per-head softmaxes get peakier and harder to optimize — plus you can afford fewer heads per parameter budget. Sixty-four sits in the robust middle, and the field converged there: BERT uses 64, GPT-2 uses 64, most modern models use 64–128. When you design a small model, keep d_k in this range and scale the number of heads with d_model instead.
For your research: Attention heads are a ready-made interpretability project. Pick a trained model in your domain (even a small one you train yourself), visualize each head's attention on domain examples, and catalog what each head tracks. Papers have been published doing exactly this for new languages and modalities ("what do heads learn in Urdu BERT?" is a legitimate research question). If some heads are redundant, head-pruning for your domain's efficient deployment is another publishable angle.
Key takeaways: - Multi-head attention runs h parallel attention operations with separate learned projections, then concatenates. - Different heads specialize emergently: coreference, syntax, local context, delimiters. - Splitting costs little extra compute because per-head dimensions shrink (d_k = d_model / h). - Multiple heads let the model track several relationship types simultaneously without compromise. - Head count ~ d_model/64 is a good rule of thumb; more heads are not automatically better.
Here is a subtle problem. Self-attention treats its input as a set, not a sequence: the operation "every token attends to every other token" has no notion of left-to-right order. If you shuffle the words of a sentence, attention computes the same pairwise interactions — only the positions change. But "dog bites man" and "man bites dog" contain the same words in different order with opposite meanings. The transformer needs to know where each token sits.
RNNs got order for free: they read word 1, then word 2, then word 3 — position was baked into the process. Transformers threw away the sequential process, so they must inject order information explicitly. That injection is positional encoding.
A good positional encoding should: 1. Give every position a unique signature (position 7 must differ from position 8). 2. Work for sequences longer than those seen in training (ideally). 3. Let the model easily learn relative positions — "the word 3 positions to my left" matters more in language than "absolute position 42." 4. Not drown out the word's own meaning — it's a seasoning, not the main dish.
The original paper proposed a beautiful fixed (non-learned) encoding using sine and cosine waves of different frequencies:
PE(pos, 2i) = sin(pos / 10000^(2i/d_model)) PE(pos, 2i+1) = cos(pos / 10000^(2i/d_model))
For each position pos, we generate a d_model-dimensional vector. Dimension pair i uses a wave with wavelength growing geometrically from 2π to 10000·2π. Low dimensions oscillate fast (distinguishing neighboring positions); high dimensions oscillate slowly (giving a coarse sense of where in the sequence you are).
Why sines and cosines? Three elegant properties: - Uniqueness: the combination of frequencies makes each position's vector distinct. - Bounded: values stay in [−1, 1], so they never overwhelm the word embedding. - Relative positions as linear shifts: for any fixed offset k, PE(pos+k) is a linear function of PE(pos) — a rotation in each 2-D subspace. The model can therefore learn to attend to "k positions away" with a simple linear operation. This is the mathematical gem of the design.
Take d_model = 4 (two frequency pairs, i = 0 and i = 1): - i=0: divisor = 10000^0 = 1 → PE(pos,0) = sin(pos), PE(pos,1) = cos(pos) - i=1: divisor = 10000^(2/4) = 100 → PE(pos,2) = sin(pos/100), PE(pos,3) = cos(pos/100)
Position 0: [sin0, cos0, sin0, cos0] = [0, 1, 0, 1] Position 1: [sin1, cos1, sin0.01, cos0.01] ≈ [0.841, 0.540, 0.010, 1.000] Position 2: [sin2, cos2, sin0.02, cos0.02] ≈ [0.909, −0.416, 0.020, 1.000]
Each position gets a unique fingerprint. The fast pair (dims 0–1) already separates positions 0, 1, 2 clearly; the slow pair (dims 2–3) barely moves — it will distinguish position 5 from position 500.
The positional vector is added to the token's embedding (not concatenated): input = word_embedding + positional_encoding. Addition keeps the dimension at d_model (no growth), and the model learns to disentangle "what the word is" from "where it sits" — the embedding space is large enough to carry both signals.
A frequent beginner worry: "doesn't adding corrupt the word meaning?" In practice, no. The embedding dimensions are numerous (512+), the positional signal is structured and bounded, and training adjusts the word embeddings to coexist with it. Think of it as writing the page number lightly in the margin of each word — the text remains readable.
Instead of fixed sines, you can simply learn a vector per position (a lookup table: position 7 → learned vector p₇), trained with the rest of the model. GPT and BERT both did this. Learned encodings often perform slightly better on the training distribution, but they cannot handle positions beyond the maximum seen in training — position 600 has no learned vector if you only trained on length ≤ 512. Sinusoidal encodings extrapolate more gracefully (though not perfectly).
Practical rule: for fixed-maximum-length applications, learned encodings are fine and simple. For variable or very long sequences, fixed or relative schemes generalize better.
Later research (notably Shaw et al., 2018) argued that what really matters is relative distance — in "the cat sat," the model cares that "cat" is 1 position left of "sat," not that they sit at absolute positions 12 and 13. Relative schemes inject the distance (i − j) directly into the attention score. Modern large models use rotary embeddings (RoPE), which rotate query/key vectors by an angle proportional to position — a clever descendant of the sinusoidal idea that bakes relative distance into the dot product itself. The details are beyond this book, but the lineage is clear: position as rotation, from 2017 to today.
import torch
import math
def sinusoidal_positional_encoding(max_len, d_model):
pe = torch.zeros(max_len, d_model)
position = torch.arange(max_len).unsqueeze(1).float()
div = torch.exp(torch.arange(0, d_model, 2).float()
* -(math.log(10000.0) / d_model))
pe[:, 0::2] = torch.sin(position * div)
pe[:, 1::2] = torch.cos(position * div)
return pe # shape (max_len, d_model); add to embeddings
# usage
x = token_embeddings + sinusoidal_positional_encoding(seq_len, d_model)[:seq_len]
import torch.nn as nn
# learned alternative: one vector per position, trained with the model
pos_embed = nn.Embedding(max_len, d_model) # max_len = e.g. 512
positions = torch.arange(seq_len)
x = token_embeddings + pos_embed(positions) # same "add" pattern
Same addition pattern, different source of the vectors. BERT uses exactly this. The trade-off from the chapter bears repeating in code terms: nn.Embedding(512, ...) has no row 600 — ask for position 600 and you get an error (or garbage). Sinusoidal encoding computes position 600 on the fly. Choose based on whether your sequences have a hard maximum length.
Shaw et al. (2018) inject the distance between tokens directly into the score: score(i,j) = q_i·k_j + q_i·r_{i−j}, where r_{i−j} is a learned vector for each possible offset (clipped to a range like ±16). Example: in "the cat sat," for i = "sat" (position 2):
The model learns r_{−1} = "immediately left" as a reusable concept: whatever words occupy positions (i, i−1), the same offset vector applies. Absolute encodings must relearn "left neighbor" separately at every position; relative encodings learn it once. This parameter sharing is why relative schemes generalize better to longer sequences — and why modern models (via RoPE and relatives) all moved in this direction.
Rotary Position Embedding (RoPE) is the current standard in large language models, and it's a direct descendant of the sinusoidal idea. Instead of adding position vectors, RoPE rotates each query and key vector by an angle proportional to its position (different rotation speeds per dimension pair — the same geometric wavelengths as 2017). The dot product of a rotated q_i and rotated k_j then depends only on (i − j): relative position falls out of the math automatically, with no extra terms. Elegant, extrapolation-friendly, and cheap — which is why it won. You don't need its matrix form yet; remember the slogan: 2017 added waves; RoPE rotates by them.
Flat 1-D position numbers work for ViT, but images are 2-D: patch (row 3, col 7) is "near" patch (row 3, col 8) in a way 1-D numbering obscures. Options researchers use: separate learned row-embeddings plus column-embeddings (summed), or 2-D sinusoidal encodings (half the dimensions encode row, half encode column). For document understanding (forms, receipts), 2-D position is nearly mandatory — reading order alone doesn't capture "the total is at the bottom right." Another instance of the chapter's lesson: match the position scheme to the data's true geometry.
Train a tiny transformer (d_model=64, 2 layers) on a copying task with sequences of length ≤ 32, using (a) sinusoidal and (b) learned positional encodings. Test on length 64. You'll typically see (a) degrade gracefully while (b) collapses — its position-33+ vectors were never trained. This 30-minute experiment teaches more about positional encoding than any textbook paragraph, and it's the template for the ablation section of a methods paper: same model, same data, one factor varied, mechanism explained.
Vaswani et al. actually tried both and reported the result: learned positional embeddings performed nearly identically to sinusoidal on their translation task. They chose sinusoidal anyway — for its extrapolation behavior and because it might help the model attend by relative positions. This is exemplary science: run the comparison, report the tie, choose on principled secondary grounds, and say so. Too many papers bury tied ablations. When your experiment shows "no significant difference," that sentence belongs in the paper — it saves the next researcher a month.
What if you skip positional encoding entirely? Self-attention becomes permutation-equivariant: shuffle the input tokens, and the outputs shuffle identically — the model provably cannot distinguish "dog bites man" from "man bites dog." It's a bag-of-words model with extra steps. This is actually a useful diagnostic: if your model performs suspiciously well without positional encoding, your task may not need order at all (some classification tasks barely do) — and you should say so, because it reframes what the model is learning. Conversely, the Chapter 5 extrapolation experiment (train ≤32, test at 64) shows how much order-sensitivity the encoding provides.
ALiBi (Attention with Linear Biases, Press et al., 2022) takes a radically simple approach: no positional vectors at all. Instead, subtract a fixed penalty proportional to distance from each attention score: score(i,j) −= m·|i−j|, with a per-head slope m. Nearby tokens get full scores; distant tokens are penalized linearly. It extrapolates to longer sequences beautifully (the penalty just keeps growing) and costs nothing to compute. The lesson: positional information can live in the scores, not just the embeddings — another point in the design space your custom-encoding project (Chapter 5's research box) can explore.
The decoder uses the same positional encodings as the encoder — position 5 means "fifth token" on both sides. What differs is the causal mask (Chapter 3), which enforces that generation only uses the past. A subtle point: because the decoder's input during training is the target sequence shifted right (it predicts token t+1 from tokens ≤ t), position t's encoding must align with "I have seen t tokens." Off-by-one errors here are a classic implementation bug — if your from-scratch decoder generates gibberish while the encoder works, check the target shifting before anything else.
For d_model = 512, dimension pair i has wavelength 2π·10000^(2i/512):
| Pair i | Wavelength (tokens) | Role |
|---|---|---|
| 0 | ~6.3 | Distinguishes immediate neighbors |
| 64 | ~25 | Words within a phrase |
| 128 | ~628 | Sentences and paragraphs |
| 192 | ~15,800 | Whole documents |
| 255 | ~62,800 | Corpus-scale coarse position |
The geometric spacing means the encoding simultaneously represents position at every scale — like having rulers marked in millimeters, centimeters, and meters all at once. A head that needs "the previous word" reads the fast pairs; a head tracking "which paragraph" reads the slow ones. This multi-scale property is the real reason sinusoids beat a simple counter (position = one number): a single number forces the model to decode scale itself.
For your research: Positional encoding is an under-explored lever in applied work. If your data has natural structure beyond flat order — timestamps on sensor readings, 2-D layout of document pages, hierarchical sections — a custom positional scheme that reflects that structure is a legitimate, publishable contribution. "We replaced 1-D sinusoidal encoding with time-gap-aware encoding for irregular medical time series and gained X" is exactly the kind of focused paper students publish. Always ablate: compare against plain sinusoidal and learned baselines so reviewers see the gain comes from your idea.
Key takeaways: - Self-attention is order-blind; position must be injected explicitly. - Sinusoidal encoding gives each position a unique, bounded fingerprint of mixed-frequency waves. - PE(pos+k) is a linear function of PE(pos), letting the model learn relative offsets easily. - The encoding is added to (not concatenated with) token embeddings. - Learned positional embeddings work well within training length but don't extrapolate; relative schemes (Shaw 2018, RoPE) are the modern refinement.
We have built attention (Chapters 2–4) and position (Chapter 5). Now we assemble the complete transformer: stacks of identical layers, residual connections, normalization, and feed-forward networks — in three classic configurations.
The original transformer has two stacks: - Encoder (left): reads the entire input sequence at once and builds rich, context-aware representations of it. Used for understanding tasks. - Decoder (right): generates the output sequence one token at a time, looking back at its own past outputs and at the encoder's representations. Used for generation tasks.
Each stack is made of N = 6 identical layers (blocks) piled on top of each other. Representations get progressively more abstract as they rise: early layers track local patterns, later layers track long-range meaning.
An encoder layer does three things, in order:
In code-shaped pseudocode:
x = x + MultiHeadAttention(LayerNorm(x)) # (pre-norm variant, modern)
x = x + FFN(LayerNorm(x))
(The original paper applied norm after the residual add — "post-norm." Modern implementations usually normalize before each sub-layer — "pre-norm" — because it trains more stably in deep models. Know both; expect pre-norm in code you read.)
Each piece has a job: - Attention mixes information across positions (communication). - FFN transforms information within each position (computation). With 2048 hidden units, it holds most of the layer's parameters — researchers think of FFNs as the model's factual memory. - Residual connections (from ResNets, He et al., 2016) let gradients flow straight through the stack, enabling deep networks. Without them, 6+ layers would barely train. - Layer normalization (Ba et al., 2016) rescales each token's vector to stable statistics, preventing activations from drifting as they pass through layers.
A decoder layer has three sub-layers instead of two:
Each sub-layer is wrapped in its own residual + normalization. The decoder is deeper per layer because it juggles two information sources: its own past and the encoder's output.
Training feeds the whole target sentence at once (with masking to prevent cheating — this parallelism is a huge speedup over RNN training). But inference is still sequential: to translate, the decoder generates one token, appends it to its input, and repeats. The encoder runs once; the decoder runs once per output token, reusing cached keys/values from previous steps (the KV cache — an important inference optimization you'll meet in deployment work).
Nearly every transformer model is one of:
| Family | Uses | Examples | Best for |
|---|---|---|---|
| Encoder-only | Encoder stack alone | BERT (Devlin et al., 2019) | Understanding: classification, question answering, embeddings |
| Decoder-only | Decoder stack alone (no cross-attention) | GPT series (Radford et al., 2018; Brown et al., 2020) | Generation: chat, completion, code — and, surprisingly, almost everything else at scale |
| Encoder–decoder | Full original design | T5 (Raffel et al., 2020), original transformer | Sequence-to-sequence: translation, summarization |
A historical note for your literature reviews: the field started with encoder–decoder (2017), split into encoder-only BERT and decoder-only GPT (2018–2019), and then decoder-only models unexpectedly came to dominate — scaling a simple next-token predictor turned out to cover translation, summarization, and reasoning too. T5 showed encoder–decoder remains competitive when the task is explicitly framed as text-to-text.
Input (German): "Die Katze sitzt auf der Matte." Target (English): "The cat sits on the mat."
Figure 3: The full transformer — an encoder stack (left) feeding a decoder stack (right), with residual connections around every sub-layer.
Architecture is half the story; the pretraining objective (what the model is asked to predict) is the other half:
When you choose a pretrained model, match the objective to your task: understanding → BERT-style; generation → GPT-style; explicit input→output mapping → T5-style. Mismatching (e.g., forcing BERT to generate) fights the model's training.
The original paper: x = LayerNorm(x + Sublayer(x)) (post-norm — normalize after adding). Modern code: x = x + Sublayer(LayerNorm(x)) (pre-norm — normalize before). Why did the field switch? In post-norm, the residual stream's magnitude grows with depth (each layer adds unnormalized outputs), so gradients at early layers get scaled down relative to late layers — deep models (24+ layers) train unstably or need careful LR tuning. Pre-norm keeps the residual stream's scale controlled: each sub-layer reads a normalized input, and its contribution is added cleanly. Gradients flow evenly to all layers. Rule: post-norm can score slightly better on shallow models with tuned hyperparameters; pre-norm trains reliably at depth. For anything you build, use pre-norm.
Translating "Die Katze sitzt" → generating "The cat sits": the decoder runs once per output token. Naively, generating token 3 recomputes keys/values for tokens 1–2 from scratch — O(n) wasted work per step, O(n²) total. The KV cache stores each step's K and V: at step t, compute only the new token's q/k/v, append k/v to the cache, and attend over the cached keys/values. Per-step cost drops because the projections (the expensive d×d multiplications) run once per token instead of t times. Memory cost: 2 (K and V) × layers × heads × d_k × n floats — for long generations this cache, not the weights, dominates GPU memory. Every serving system (vLLM, TensorRT-LLM) is largely KV-cache engineering.
| Component | Parameters (base: 6 layers, d=512) |
|---|---|
| Token embeddings (37k vocab × 512) | ~19M |
| Encoder: 6 × (attention ~1.05M + FFN ~2.1M) | ~19M |
| Decoder: 6 × (2 attentions ~2.1M + FFN ~2.1M) | ~25M |
| Output softmax (tied with embeddings) | 0 (shared) |
| Total | ~65M |
Two lessons: the FFN (~2.1M/layer = 2×512×2048) is twice the attention — the "memory" dominates the "mixing." And embeddings are huge for large vocabularies — which is why subword vocabularies (~30–50k) are a sweet spot, not character-level (too long sequences) or word-level (too many rare words).
The original FFN used ReLU. BERT switched to GELU (Hendrycks and Gimpel, 2016) — a smooth, probabilistic variant (x·Φ(x)) that weights inputs by how likely they are to be "on" rather than hard-thresholding at zero. Empirically slightly better for transformers; now the default in nearly all implementations. When porting old code or reading ablations, treat ReLU→GELU as a small-but-real upgrade, not a cosmetic choice.
Recent interpretability research reframes the architecture: think of the residual connections as a single shared residual stream — a conveyor belt of dimension d_model running through all layers. Each attention head and FFN reads from the stream (via layer norm) and writes back into it (via its output projection). Layers don't hand off private representations; they all annotate one shared workspace. This view (developed in work on transformer circuits, e.g., Elhage et al., 2021) explains why you can sometimes "read" predictions directly from intermediate layers, and why ablating one head often barely matters — the stream carries redundant annotations. For interpretability projects, probe the residual stream at each layer boundary: it's where the model's evolving beliefs live.
In translation models, cross-attention weights often form clean word alignments: when generating the English word "cat," the head puts heavy weight on the German "Katze." Researchers plot these as heatmaps — and they look strikingly like the alignments from 1990s statistical machine translation (IBM models), learned here with no explicit supervision. Not all cross-attention heads align words (some track syntax, some spread broadly), but the alignment heads are a gift for analysis: they let you see what the model thinks corresponds to what. If your encoder–decoder project involves two sequences (speech-to-text, image captioning), visualizing cross-attention is step one of understanding it.
The original N=6 was a pragmatic choice for 2017 hardware, not a law. The pattern since: depth (layers) builds abstraction hierarchies, width (d_model) builds representational capacity per level. BERT-base: 12 layers × 768; BERT-large: 24 × 1024; GPT-3: 96 × 12288. Rough empirical wisdom: for a fixed parameter budget, balanced depth and width beat extremes — very deep + narrow starves each layer of capacity; very shallow + wide can't build hierarchies. When scaling down for your hardware, shrink both proportionally (e.g., 4 layers × 256) rather than keeping 12 paper-thin layers.
After the decoder stack, one linear layer maps d_model → vocabulary size (e.g., 512 → 37,000), and softmax turns it into next-token probabilities. With weight tying, this matrix is shared with the input embedding (transposed) — the model uses the same vectors to read tokens in and predict tokens out, saving millions of parameters and slightly improving quality. At inference, you don't need the full 37k-way softmax per step if you only want the top token — greedy decoding takes the argmax, beam search keeps the top-k hypotheses. Generation strategy (greedy vs. beam vs. sampling) is a whole topic of its own, but architecturally it's all "read the output head."
For your research: Architecture choice is a first-class research decision, not a default. For classification of fixed inputs, encoder-only (BERT-style) is the natural baseline. For generation, decoder-only. For tasks with distinct input and output sequences (translation, summarization, speech-to-text), encoder–decoder. In your paper's methodology, one sentence justifying the family ("we use an encoder-only architecture because the task is classification over complete inputs") signals competence to reviewers. And when resources are tight, remember: a well-tuned small model in the right family beats a poorly-trained large one in the wrong family.
Key takeaways: - Encoder layers: self-attention → FFN, each wrapped in residual + norm. Decoder layers add masked self-attention and cross-attention. - Attention mixes information across positions; the FFN (most parameters) transforms within positions. - Residuals enable depth; layer norm stabilizes training (modern code uses pre-norm). - Three families: encoder-only (understanding), decoder-only (generation), encoder–decoder (sequence-to-sequence). - Training is parallel (masked targets); inference generates one token at a time with a KV cache.
A transformer with random weights is useless; a transformer trained well is a revolution. Training is where many student projects silently fail — not because the architecture was wrong, but because the training recipe was. This chapter gives you the recipe: the handful of choices that separate "it doesn't learn" from "it works."
Two properties make transformers harder to train than, say, a small CNN:
The single most famous training detail from the original paper: don't start at full learning rate. Instead:
The original formula: lr = d_model^(−0.5) · min(step^(−0.5), step · warmup_steps^(−1.5)). You don't need to memorize it — every library implements get_linear_schedule_with_warmup or cosine variants. What you must remember: if your transformer loss explodes or flatlines in the first epoch, the warmup is the first suspect. This one trick rescues more student training runs than any other.
Transformers are almost always trained with Adam (Kingma and Ba, 2014) — an optimizer that adapts the step size per parameter using running averages of gradients and their squares. The original paper used β₁=0.9, β₂=0.98, ε=10⁻⁹. Modern practice prefers AdamW (Loshchilov and Hutter, 2019), which handles weight decay correctly (decay applied to the weights themselves, not folded into the gradient). Practical effect: slightly better generalization, and it's the default in Hugging Face trainers. Use AdamW unless you have a reason not to.
Transformer training uses very large batches — the original paper used ~25,000 tokens per batch; large language models use millions. Why? Attention gradients are noisy (each token's update depends on all others), and large batches average out that noise, giving the optimizer a reliable direction. If your hardware can't fit a big batch, use gradient accumulation: run several small batches, sum their gradients, and update once — mathematically close to one big batch, fitting in small GPU memory. This is the standard student trick for training on a single GPU.
Three regularizers appear in nearly every transformer recipe:
The original transformer trained on WMT 2014: 4.5 million English–German sentence pairs. That was large for 2017. The bitter lesson since: transformer quality scales remarkably predictably with data and compute — a finding that drove the entire large-language-model era. For your work, the practical reading is: before tuning the architecture, check whether you simply need more (or cleaner) data. Many student "architecture improvements" vanish when the baseline gets the same data budget.
Modern GPUs compute in 16-bit floating point (FP16/BF16) roughly twice as fast as 32-bit, using half the memory. Mixed-precision training keeps a 32-bit master copy of weights but computes forward/backward passes in 16-bit, with loss scaling to protect tiny gradients. In PyTorch it's a few lines (torch.cuda.amp), in Hugging Face Trainer it's a flag. There is almost no reason not to use it — it roughly doubles your effective batch size.
import torch
from torch.optim import AdamW
from transformers import get_linear_schedule_with_warmup
model = MyTransformer().cuda()
optimizer = AdamW(model.parameters(), lr=5e-4, weight_decay=0.01)
scheduler = get_linear_schedule_with_warmup(optimizer,
num_warmup_steps=4000, num_training_steps=100000)
scaler = torch.cuda.amp.GradScaler() # mixed precision
model.train()
for step, batch in enumerate(dataloader):
optimizer.zero_grad()
with torch.cuda.amp.autocast(): # 16-bit compute
loss = model(batch) # your forward + loss
scaler.scale(loss).backward()
scaler.step(optimizer) # gradient accumulation:
scaler.update() # call step() every k batches
scheduler.step() # warmup then decay
Picture the learning rate over 100,000 steps with 4,000 warmup steps and peak lr = 5e-4: a straight ramp from 0 up to 5e-4 (steps 0–4,000), then a long gentle curve down proportional to 1/√step, ending near 5e-4 × √(4000/100000) ≈ 1e-4. The ramp is short but critical — it's where attention patterns crystallize. The decay is long and boring — that's fine; boring is stable.
from transformers import get_linear_schedule_with_warmup
# linear warmup, then linear decay to zero (a common modern choice)
scheduler = get_linear_schedule_with_warmup(
optimizer, num_warmup_steps=4000, num_training_steps=100000)
Cosine decay (smooth cosine curve down to near-zero) is equally popular and sometimes slightly better. The choice between 1/√step, linear, and cosine decay matters far less than having a warmup at all. When reviewing papers, check warmup first, decay shape second.
"Batch size 32" means little for transformers if sentences vary in length — 32 tweets vs. 32 long paragraphs differ 10× in compute. Practitioners count tokens per batch instead: e.g., "25,000 tokens/batch" (the original paper). Implementation: pack sequences, or bucket by length and pad within buckets. Padding wastes compute on meaningless pad tokens, so packing (concatenating short sequences to fill the length budget, with attention masks keeping them separate) is standard in serious training. For your experiments, at least report tokens/batch — it makes your setup comparable across papers.
For language-model training, the headline metric is perplexity = e^(average negative log-likelihood). Intuition: a perplexity of 50 means the model is as confused as if choosing uniformly among 50 options at each step. It always decreases as training progresses (on training data), and its validation curve tells you when to stop: when validation perplexity rises while training perplexity falls, you're overfitting. For translation, BLEU is the final judge, but perplexity is the day-to-day compass — cheaper to compute and smoother to read. Log both.
Two students train the same 6-layer transformer on the same translation data. Student A uses constant lr = 1e-3, no warmup. Student B uses warmup 4,000 steps to peak 5e-4, then decay. After 2 hours: A's loss oscillates wildly and NaNs at step 9,000 — the early giant steps shattered the fragile initial attention patterns, and the model never recovered. B's loss descends smoothly; by hour 6 it beats A's best checkpoint by 4 BLEU. Same architecture, same data, same GPU — the schedule was the difference. This story replays in labs constantly. When a run dies young, suspect the first 4,000 steps before suspecting the model.
Even with warmup, occasional batches produce huge gradients (a weird sentence, a rare token). Gradient clipping caps the gradient norm: if ||g|| > 1.0, scale g down to norm 1.0. It doesn't change healthy updates (their norms are small) but prevents freak batches from undoing hours of training. One line in PyTorch (torch.nn.utils.clip_grad_norm_(model.parameters(), 1.0)), placed before optimizer.step(). Use it always — it's free insurance.
Adam keeps two running averages per parameter: m (momentum — the average of recent gradients, smoothing out noise) and v (the average of squared gradients — measuring how volatile each parameter's gradients are). The update is roughly m / √v: parameters with steady gradients take confident steps; parameters with wild gradients get scaled down automatically. That's why Adam tolerates one global learning rate across embeddings (sparse updates), attention (dense, noisy), and FFNs (huge matrices) — it adapts per parameter. The original paper's β₂=0.98 (vs. the usual 0.999) makes v forget faster, reacting quicker to the shifting gradient landscape of early transformer training. You don't need to tune these; you need to know why Adam is the default: heterogeneous layers need per-parameter adaptation.
In classic Adam, L2 regularization was implemented by adding λ·w to the gradient — but Adam then adapts that term per parameter, so the effective decay varies unpredictably. AdamW (Loshchilov and Hutter, 2019) decouples it: update with Adam, then separately shrink weights by (1 − λ·lr). Same idea, correct mechanics — and consistently better generalization for transformers. Practical note: with AdamW, the learning rate interacts with the decay strength, so retune λ (0.01 is the usual start) if you change lr drastically. Every Hugging Face Trainer uses AdamW by default; if you hand-roll training, import it, not Adam.
Transformer runs take hours to days — machines crash, preemptions happen. Save checkpoints regularly: model weights plus optimizer state (m, v), scheduler step, and RNG state. Resuming from weights alone restarts the optimizer cold — m and v reset to zero, effectively a new warmup the schedule doesn't know about, often causing a visible loss spike. Trainer handles this with save_steps and resume_from_checkpoint. Rule: checkpoint every N steps to two rotating files (so a crash mid-write doesn't corrupt your only copy), and always verify a checkpoint resumes bit-identically before trusting a week-long run to it.
Set Python, NumPy, and PyTorch seeds for reproducibility — but know their limits: cuDNN's fastest convolution/attention kernels are nondeterministic, so exact bit-reproducibility needs torch.use_deterministic_algorithms(True) at a speed cost. For papers, the honest standard isn't bit-identical reruns — it's multiple seeds with mean ± std. If your "improvement" is smaller than the seed variance, it's noise. Report both; reviewers increasingly demand it.
Training loss always falls; validation metrics eventually turn. Early stopping monitors a validation metric (perplexity, BLEU, F1) and halts after patience epochs without improvement, keeping the best checkpoint. It fights overfitting and saves compute — but set patience generously (transformers plateau before improving again) and monitor the task metric, not just loss. Combined with load_best_model_at_end, it's the standard recipe: train long, keep the peak.
For your research: Training details belong in your paper — reviewers increasingly check them. Report optimizer, peak learning rate, warmup steps, batch size (in tokens if sequences vary), dropout, and hardware/steps. A "reproducibility checklist" paragraph costs you five lines and buys credibility. Better: publish your training script alongside the paper. And when comparing against a baseline, give both the same training budget — nothing undermines a paper faster than a baseline that was obviously undertrained.
Key takeaways: - Warmup (gradual LR increase) then decay is the most important transformer training trick. - Use AdamW; large effective batches (via gradient accumulation if needed) stabilize noisy attention gradients. - Dropout (0.1), label smoothing (0.1), and weight decay are the standard regularization trio. - Quality scales with data — check data quantity/quality before blaming the architecture. - Mixed precision (~2× speedup, half memory) and gradient clipping are essentially free wins.
In 2020, the transformer escaped language. Dosovitskiy and colleagues asked: if attention works on word sequences, why not on image patches? The Vision Transformer (ViT) treated an image as a sequence — and matched or beat state-of-the-art CNNs. This chapter explains how, and why it matters far beyond images.
A transformer needs a sequence of vectors. An image is a grid of pixels — so ViT cuts it into fixed patches:
That's it. No convolutions anywhere. The model must learn from data that neighboring patches relate — the locality that CNNs get for free.
ViT's bet: with enough data, learned attention beats hand-designed locality. On JFT-300M (300 million images), ViT-Huge outperformed the best CNNs. But on small datasets like CIFAR-scale or a few thousand images, CNNs still win — their built-in assumptions (local patterns, translation invariance) are exactly right for small data, while ViT flounders without enough examples to learn those patterns from scratch.
The practical rule, confirmed across dozens of studies: ViT for large data or transfer learning; CNNs (or hybrids) for small data from scratch. And in practice, almost nobody trains ViT from scratch — you fine-tune a pretrained one (Chapter 10), which erases most of the data-hunger problem.
Visualizations of ViT attention heads show the same emergent specialization as language models: early layers attend locally (a patch looks at its neighbors — the model rediscovers convolution-like behavior), while deeper layers attend globally (a patch of a dog's ear attends to the dog's tail across the image). Some heads track object boundaries; the [CLS] token's attention often highlights the foreground object. The architecture discovered, from pixels and labels alone, the local-to-global processing that vision scientists designed by hand.
ViT's deeper message: anything you can cut into a sequence of tokens can go through a transformer. This launched a wave:
For a researcher, this is liberating: your domain's data probably tokenizes somehow, and then the entire transformer toolkit — pretraining, fine-tuning, efficient attention — becomes available.
A landmark follow-up, DeiT (Touvron et al., 2021), showed ViT could be trained on ImageNet-1k alone (1.3M images, no giant JFT dataset) by distilling from a CNN teacher: the student transformer learns to mimic the teacher's outputs via a special distillation token. Lesson for students: when data is scarce, distillation from a stronger model in a related modality is a principled, publishable strategy — not a hack.
import torch.nn as nn
class PatchEmbed(nn.Module): # the "tokenizer" for images
def __init__(self, img_size=224, patch=16, in_ch=3, d_model=768):
super().__init__()
self.proj = nn.Conv2d(in_ch, d_model,
kernel_size=patch, stride=patch)
self.n_patches = (img_size // patch) ** 2
def forward(self, x): # x: (B, 3, 224, 224)
x = self.proj(x) # (B, d_model, 14, 14)
return x.flatten(2).transpose(1, 2) # (B, 196, d_model)
# then: prepend CLS, add pos. embedding, run transformer encoder
(A convolution with kernel=stride=patch is exactly "cut into patches and linearly project each" — neat, and GPU-friendly.)
For a 224×224 image with 16×16 patches: 224/16 = 14 patches per side → 14×14 = 196 patch tokens, plus 1 [CLS] token = 197 tokens into the encoder. Each patch flattens to 16×16×3 = 768 numbers, projected to d_model (768 for ViT-Base — a coincidence of numbers, not a requirement). Attention cost per layer: 197² × 768 ≈ 30M operations — trivial for a GPU. Now try 1024×1024 with 16×16 patches: 4,096 tokens → 4096² × 768 ≈ 13B operations per layer — painful. This is why high-resolution vision transformers need the efficient attention of Chapter 9, or hierarchical designs that merge patches as they go deeper.
Try other patch sizes on 224×224: 32×32 patches → 49 tokens (coarse, fast, loses detail); 8×8 → 784 tokens (fine, slow). Patch size is a resolution-vs-cost knob, and the right setting depends on how small your objects of interest are — a tumor in a medical scan needs small patches; scene classification doesn't.
from transformers import ViTForImageClassification, ViTImageProcessor
proc = ViTImageProcessor.from_pretrained("google/vit-base-patch16-224")
model = ViTForImageClassification.from_pretrained(
"google/vit-base-patch16-224", num_labels=5) # your 5 classes
inputs = proc(images=[img1, img2], return_tensors="pt") # resize+normalize
outputs = model(**inputs) # logits: (2, 5)
# freeze the body, train only the head (feature-extraction mode):
for p in model.vit.parameters(): p.requires_grad = False
The processor handles the fiddly bits — resizing to 224, normalizing with ImageNet mean/std. Match the processor to the model (same rule as tokenizers in Chapter 10): a model pretrained with one normalization will silently underperform with another.
DeiT's trick: train the student ViT with two objectives — the normal classification loss, plus a distillation loss matching the CNN teacher's soft predictions, routed through a dedicated distillation token (a second special token alongside [CLS]; the two tokens' predictions are averaged at test time). The teacher's soft targets ("70% cat, 25% dog, 5% fox") carry far more information than hard labels ("cat") — they teach the student which confusions are reasonable. Result: DeiT trained competitively on ImageNet-1k alone (1.3M images) with no giant proprietary dataset. The general principle travels: whenever you lack data but can access a strong teacher (even in another modality), distillation compresses its knowledge into your student.
Two honest failure modes to know before you commit:
Pure ViT vs. pure CNN is a false choice. Many strong models use a convolutional stem (a few conv layers that turn pixels into feature maps) followed by transformer layers over the resulting tokens — convolutions handle local detail efficiently, attention handles global reasoning. For applied work, hybrids often give the best quality-per-compute on mid-size datasets. Don't let architectural purity override results.
Because ViT must learn locality instead of assuming it, training augmentation is load-bearing, not optional. The DeiT recipe: RandAugment (random photometric transforms), Mixup/CutMix (blending image pairs and their labels — forcing the model to use global evidence, since local patches may come from two images), stochastic depth (randomly dropping whole layers during training — the layer-wise cousin of dropout), and repeated augmentation. Ablations show ViT-Base trained without this recipe loses several accuracy points; a ResNet loses far less. Lesson: when your ViT underperforms, suspect the training recipe before the architecture — and when you publish ViT results, report the augmentation stack as carefully as the model.
Swin Transformer (Liu et al., 2021) tackles the high-resolution cost problem hierarchically: it computes attention only within local windows (e.g., 7×7 patches), then shifts the windows each layer so information crosses boundaries, while progressively merging patches (like CNN downsampling) to build a pyramid. Cost stays linear in image size, and the model handles detection/segmentation natively — tasks where plain ViT's single-scale tokens struggle. Conceptually it's "local attention + shifting + merging" — a hybrid of CNN wisdom and transformer machinery. For dense prediction tasks (find where, not just what), Swin-style hierarchical models are usually the stronger baseline than plain ViT.
ViT's learned positional embeddings are fixed at 196 positions — what if you fine-tune on 384×384 images (576 patches)? Interpolate the 14×14 position grid to 24×24 (bicubic), then fine-tune briefly. The model adapts in a few epochs and accuracy rises (higher resolution helps). This trick — pretrained at 224, fine-tuned at 384 — is standard practice and a good example of the field's pragmatism: the theory says positions are learned for 196 slots; practice says smooth interpolation plus fine-tuning just works. Always try the simple hack before designing a new positional scheme.
DETR (Carion et al., 2020) rebuilt object detection on transformers: a CNN backbone feeds image features to a transformer encoder–decoder, and the decoder takes learned object queries (like slots asking "is there an object for me?") and outputs boxes directly — no anchors, no non-maximum suppression, trained with bipartite matching loss. Performance matched Faster R-CNN with a dramatically simpler pipeline. DETR matters for this book because it shows the decoder (Chapter 6) working outside language: queries attending to image features via cross-attention is exactly translation's cross-attention, pointed at pixels. When your domain has "find all instances of X," DETR-style query decoding is a baseline worth knowing.
Video (TimeSformer, Bertasius et al., 2021): cut video into space-time patches (tubes spanning a few frames), embed each as a token, and run divided attention — spatial attention within frames plus temporal attention across frames. The divided design keeps cost manageable while letting the model track motion. Conceptually it's ViT plus a time axis — another instance of "tokenize everything."
Audio: convert sound to a spectrogram (time × frequency image) and patch it like ViT — the Audio Spectrogram Transformer does exactly this for sound classification. Or tokenize raw waveforms into frames. Speech recognition's transformer models (e.g., wav2vec 2.0's transformer layers over learned audio features) follow the same pattern: find the right tokenization, then attention does the rest. If your domain has a signal with time and frequency structure, spectrogram-patching is the first baseline to try.
ViT classifies from the [CLS] token, but an alternative is global average pooling — averaging all patch representations. Which is better? Studies find them close, with small edges depending on the task: [CLS] gives the model a dedicated "summary slot" with its own learned behavior (its attention pattern often highlights foreground objects), while average pooling is simpler and has no extra parameters. Some modern variants drop [CLS] entirely in favor of pooling plus a small attention-based aggregator. The takeaway for your projects: the aggregation choice is a legitimate ablation — try both, report both, and don't treat the original paper's choice as sacred.
For your research: ViT makes vision research accessible without CNN expertise. If your field has images (medical scans, crop photos, satellite tiles, industrial defects), the fastest path to a result is: take a pretrained ViT, fine-tune on your labeled images, and compare against the previous best method. That's a baseline paper. The interesting paper asks what the patches should be: for satellite imagery, should patches respect geographic scale? For medical scans, should the [CLS] token be replaced by region-aware pooling? Tokenization choices are domain knowledge meeting architecture — exactly where student contributions live.
Key takeaways: - ViT splits images into patches, embeds each as a token, and runs a standard transformer encoder. - A [CLS] token's final state serves as the image representation for classification. - ViT needs large data (or pretraining) because it must learn locality that CNNs assume; with enough data it matches or beats CNNs. - Attention heads rediscover local-to-global processing emergently. - The "tokenize anything" principle extends transformers to audio, video, time series, proteins, and multimodal data.
Standard self-attention costs O(n²): with n tokens, the attention matrix has n² entries. For n = 512, that's 262,144 — trivial. For n = 32,768 (a long document, a high-resolution image, a genome snippet), it's over a billion — impossible. This chapter surveys the main ideas for making attention cheaper. Concept level only: you should understand what each method sacrifices and keeps, not their proofs.
Two costs grow quadratically: - Compute: the QK^T multiplication does n²·d operations. - Memory: storing the n×n weight matrix (per head, per layer) for the backward pass.
For long sequences, memory usually breaks first. Every method below attacks one or both.
Observation: the n×n attention matrix is often effectively low-rank — its information can be captured with far fewer than n dimensions. Linformer (Wang et al., 2020) projects the keys and values from length n down to a smaller length k (e.g., k = 256) with learned linear projections, before attention. The attention matrix becomes n×k instead of n×n: cost O(n·k), linear in n.
Observation: softmax attention can be rewritten using a kernel trick: softmax(QK^T) = φ(Q)·φ(K)^T for a suitable (infinite-dimensional) feature map φ. Performer (Choromanski et al., 2021) approximates φ with random features (FAVOR+), computing attention in O(n) without ever forming the n×n matrix — it multiplies in a different order: (φ(Q)·(φ(K)^T·V)).
Observation: the n² matrix doesn't need to be stored — it can be recomputed in blocks. GPUs have small, blazing-fast on-chip memory (SRAM) and large, slower memory (HBM). Standard attention writes the whole n×n matrix to slow memory. FlashAttention (Dao et al., 2022) tiles the computation: it loads blocks of Q, K, V into fast memory, computes partial softmax with running statistics (a numerically careful online softmax), and never materializes the full matrix.
FlashAttention's lesson is profound: sometimes the bottleneck isn't the math, it's the memory hierarchy. An algorithms-plus-systems view beats a pure-math view.
Observation: most attention weights are near zero anyway — why compute them? Sparse methods restrict each token to attend to a subset: local windows (neighbors), strided patterns (every k-th token), or learned clusters. Longformer-style models combine local windows with a few global tokens that everyone attends to.
A newer family (state-space models like Mamba) returns to recurrence but with a parallelizable formulation: compress the past into a fixed-size state updated by a learned linear recurrence, computable in O(n) with a parallel scan. They trade attention's direct long-range links for linear scaling. Whether they displace transformers is open — but knowing they exist keeps your literature review honest.
| Situation | Consider |
|---|---|
| Training long-context models today | FlashAttention (exact, default choice) |
| Extremely long sequences, approximation OK | Performer / Linformer-style |
| Long documents with local structure | Sparse / windowed attention |
| Limited GPU memory at inference | KV-cache quantization, FlashAttention |
| Research novelty | Combine ideas: sparse + low-rank hybrids are still being explored |
Attention memory for one head, one layer, float32 (4 bytes per number), storing the n×n weight matrix:
| n (tokens) | n² entries | Memory |
|---|---|---|
| 512 | 262K | ~1 MB |
| 4,096 | 16.8M | ~67 MB |
| 16,384 | 268M | ~1.1 GB |
| 32,768 | 1.07B | ~4.3 GB |
| 131,072 | 17.2B | ~69 GB |
Multiply by heads × layers (e.g., 12 heads × 12 layers = 144×) and add gradients, and you see why 32k-context training needs either serious hardware or FlashAttention's tiling (which avoids storing the matrix at all). Keep this table in mind whenever someone proposes "just use longer context" — context length is a systems problem as much as a modeling one.
The most successful sparse recipe combines two patterns:
This combination preserves O(n) scaling while keeping a global communication channel — most of attention's practical power with a fraction of its cost. The design lesson: you don't need every pair; you need local detail plus global hubs.
If you evaluate an efficient method, report all four numbers or reviewers will (rightly) complain:
And always include two baselines: standard attention (the quality ceiling) and FlashAttention (the exact-but-fast challenger). A method that beats standard attention's memory but loses to FlashAttention's speed is a partial win — say so plainly.
Models like Mamba replace attention with a selective state-space recurrence: a hidden state h_t = A·h_{t−1} + B·x_t, where A, B (and the selectivity) are input-dependent and learned. Because the recurrence is linear, it can be computed with a parallel scan in O(n) — no n² matrix. The "selective" part lets the model decide what to remember or forget per token, recovering some of attention's content-awareness. Trade-off: the state is fixed-size (a compression bottleneck, echoing Chapter 1's RNN limits), while attention keeps the full history. Current status: competitive on long sequences, not yet dominant. For students, they're a healthy reminder that the architecture search isn't over — and a great "related work" section addition showing you know the frontier.
Linformer's core claim is empirical: the n×n attention matrix is effectively low-rank — its rows live near a k-dimensional subspace with k ≪ n (k=256 works for n=4096 in their experiments). So project keys and values down before attention: K' = E·K and V' = F·V, where E, F are learned k×n matrices. Attention becomes softmax(QK'^T/√d)·V' with an n×k score matrix — O(n·k) time and memory. The projections are learned per layer (sometimes shared across heads). Caveat for your benchmarks: the low-rank assumption holds well for trained models on natural data but can break for synthetic tasks with deliberately high-rank structure — which is exactly the kind of "where approximations break" finding (Chapter 9's research box) that makes a good paper.
Softmax attention can be written as a kernel: softmax(q·k) relates to φ(q)·φ(k) for a feature map φ. The Performer approximates φ with random features — random projections followed by a nonlinearity — so that E[φ(q)·φ(k)] equals the true softmax kernel (unbiased approximation). Then attention regroups: instead of forming (QK^T)V, compute Q'(K'^T V) — multiply K'^T·V first (d×d-ish), then by Q'. No n×n matrix ever exists: O(n) time and memory. More random features = lower variance = closer to exact. The price is randomness itself: two runs differ slightly, and quality needs enough features. It's the classic approximation trade — error bars included.
FlashAttention 2 (Dao, 2023) refined the tiling for newer GPUs (better parallelism, ~2× faster than v1); v3 targets Hopper-architecture specifics. You don't need to track versions — PyTorch's scaled_dot_product_attention (SDPA) and Hugging Face's attn_implementation="sdpa" automatically dispatch to the best available backend (FlashAttention, memory-efficient attention, or math fallback) for your GPU. Practical rule: write standard attention code, enable SDPA, and let the library choose. Only reach for explicit efficient-attention research code when your sequences exceed what SDPA handles — which, with FlashAttention under the hood, is now very long indeed.
Training must store activations for backprop — the n² matrix per head per layer is the memory wall (Chapter 9's table). Inference is kinder: with the KV cache, each generated token attends over cached keys in O(n) memory per step — but the compute per step still grows with n, so generating token 10,000 costs 10,000× the attention work of token 1. Long-context serving is therefore compute-bound at generation time and memory-bound (KV cache) at rest. This split explains the two optimization industries: FlashAttention-style kernels for training, KV-cache quantization and eviction policies for serving. Know which bottleneck your project faces before picking tools.
Approximation methods have knobs — here's how practitioners set them:
For your research: Efficiency papers are student-friendly because evaluation is crisp: same quality, less memory/time (or: longer context, same budget). You don't need a giant model to contribute — take a niche long-sequence problem in your domain (long Urdu documents, multi-day sensor streams, full-page document images), implement two efficient variants from open source, and benchmark quality-vs-cost honestly. Negative results ("method X's approximation breaks on our data because…") are publishable when the because teaches something. And always compare against FlashAttention as the strong exact baseline — reviewers will ask.
Key takeaways: - Standard attention is O(n²) in time and memory — the transformer's core weakness. - Linformer: compress keys/values to length k (low-rank) → O(n·k). - Performer: kernel trick with random features → O(n) approximation. - FlashAttention: block-wise exact computation respecting GPU memory hierarchy → same results, far less memory. - Sparse attention: hand-designed patterns (local windows + global tokens) skip near-zero weights.
Training a transformer from scratch on a big dataset is a luxury few students have. The practical path — and the one behind most published applied work — is transfer learning: take a model pretrained on massive data, and adapt it to your task. This chapter is your field manual.
Models like BERT are pretrained with masked language modeling: hide 15% of words in billions of sentences and learn to predict them. To do this well, the model must learn grammar, facts, reasoning patterns — general language competence baked into its weights. GPT-style models use next-token prediction at even larger scale. Fine-tuning then specializes this general competence to your task with far less data than training from scratch would need. It's the difference between hiring a literate adult and raising a child.
Hugging Face (the company and its open-source libraries, described in Wolf et al., 2020) is the GitHub of pretrained models: the transformers library gives you thousands of models behind a uniform API, plus tokenizers and training tools. The basic objects:
AutoModelForSequenceClassification — the pretrained body plus a task head.pipeline("sentiment-analysis") runs a full model in two lines.from transformers import pipeline, AutoTokenizer, AutoModelForSequenceClassification
import torch
# 1) zero-shot usage: sentiment in two lines
clf = pipeline("sentiment-analysis")
print(clf("The service at this restaurant was wonderful."))
# [{'label': 'POSITIVE', 'score': 0.9998}]
# 2) fine-tuning sketch: add your data, train the head + body
tok = AutoTokenizer.from_pretrained("bert-base-uncased")
model = AutoModelForSequenceClassification.from_pretrained(
"bert-base-uncased", num_labels=3)
inputs = tok(["The plot record looks correct.", "This entry is disputed."],
padding=True, truncation=True, return_tensors="pt")
with torch.no_grad():
logits = model(**inputs).logits # feature extraction: no training
print(logits.argmax(-1)) # predicted class per input
Feature extraction (frozen body): lock the pretrained weights; train only a small classifier on top of the model's outputs (often the [CLS] embedding or averaged token embeddings). Fast, cheap, hard to overfit — and often surprisingly strong. Use it when: your dataset is small (hundreds of examples), compute is limited, or you need a quick baseline.
Fine-tuning (train everything): unfreeze all weights and train end-to-end on your task with a small learning rate (typically 1e-5 to 5e-5 — much smaller than from-scratch training, so you don't destroy pretrained knowledge). Use it when: you have thousands+ of labeled examples and the task differs notably from pretraining.
The honest workflow: start with feature extraction as your baseline, then try fine-tuning, and report both. If fine-tuning barely beats frozen features, your dataset may be too small for full fine-tuning — say so; it's a finding, not a failure.
bert-base-uncased (understanding), roberta-base (stronger BERT variant), GPT-2 family (generation).bert-base-multilingual-cased or XLM-RoBERTa — multilingual models cover Urdu; check performance on your data, as low-resource languages vary.bert-base-uncased lowercases everything — fine for English sentiment, bad if case carries meaning.encode_plus/call; if you hand-build inputs, add them yourself.tok.tokenize("your sample text") before training.Pretrained models make it easy to get a number — and easy to fool yourself. Minimum discipline: fixed train/validation/test splits (never tune on test), report the frozen-feature baseline, and run 3+ random seeds (fine-tuning variance is real on small data). If you compare against prior work, match their splits exactly.
The Trainer handles the training loop, scheduling, checkpointing, and evaluation for you:
from transformers import (AutoTokenizer, AutoModelForSequenceClassification,
TrainingArguments, Trainer)
tok = AutoTokenizer.from_pretrained("bert-base-uncased")
model = AutoModelForSequenceClassification.from_pretrained(
"bert-base-uncased", num_labels=3)
args = TrainingArguments(
output_dir="ckpt", num_train_epochs=3,
per_device_train_batch_size=16, learning_rate=2e-5, # small LR!
warmup_steps=500, weight_decay=0.01,
eval_strategy="epoch", save_strategy="epoch",
load_best_model_at_end=True, fp16=True) # mixed precision
trainer = Trainer(model=model, args=args,
train_dataset=train_ds, eval_dataset=val_ds,
tokenizer=tok)
trainer.train()
Note the learning rate: 2e-5, roughly 10–25× smaller than from-scratch training. The pretrained weights are already good; you're nudging them, not rebuilding. Too large an LR causes catastrophic forgetting — the model unlearns its pretraining in the first epoch. If fine-tuning underperforms the frozen baseline, the LR is suspect number one.
Full fine-tuning updates all 110M+ parameters. Adapters insert tiny trainable bottleneck layers inside each transformer block and freeze everything else (~1–3% of parameters trainable). LoRA (Hu et al., 2021) goes further: it freezes the weights and learns low-rank update matrices (W + BA, where B and A are small) for the attention projections — often <1% of parameters, with quality matching full fine-tuning on many tasks.
Why students should care: LoRA fine-tuning of a 7B-parameter model fits on a single 24GB GPU; full fine-tuning doesn't. Your adapter weights are tiny files you can share, version, and swap per task while keeping one frozen base model. For domain adaptation papers, "LoRA-tuned X" is now a standard, respected experimental condition — not a compromise.
Between "use as-is" and "fine-tune on labels" lies a powerful middle step: continue the pretraining objective (masked LM) on your domain's unlabeled text before fine-tuning on labels. Example: take BERT, run masked-LM on 100k unlabeled Urdu legal documents (no labels needed), then fine-tune on your 2k labeled examples. Gains of 1–3 points are typical because the model learns domain vocabulary and style first. Cost: days on one GPU, no labeling. If your domain's language differs from the pretraining corpus (legal, medical, social-media dialects), this step often helps more than any architectural tweak.
A fine-tuned model in a notebook isn't a product. The deployment path:
model.save_pretrained(); consider ONNX for non-Python serving.Each step is its own skill set; for a student paper, reaching step 2 with latency/quality numbers is already a strong "deployment study" contribution.
Multilingual BERT and XLM-RoBERTa include Urdu, but with caveats students should verify rather than assume:
tok.tokenize on your data). High fertility (5+ pieces/word) eats your 512-token budget fast — consider a model with a bigger Urdu vocabulary or continued pretraining.Real datasets are rarely balanced — 95% "not spam," 5% "spam." A model that always predicts "not spam" scores 95% accuracy and is useless. Fixes: class weights in the loss (penalize minority-class errors more), oversampling the minority class, or focal loss (down-weights easy examples). And report F1 (harmonic mean of precision and recall), per-class scores, or balanced accuracy — never accuracy alone on skewed data. Reviewers check this reflexively; a paper reporting 97% accuracy on a 97%-majority dataset gets rejected on sight.
Pretrained tokenizers assume pretraining-like text. If your domain's vocabulary diverges sharply (medical terminology, code, a low-resource language), the tokenizer shreds words into many pieces (high fertility — Chapter 10's Urdu note), wasting context and blurring meaning. The fix: train a BPE/WordPiece tokenizer on your corpus (tokenizers library, ~20 lines), then either use it with a from-scratch model or — more advanced — adapt a pretrained model's embeddings to the new vocabulary. Rule of thumb: if average fertility exceeds ~2× the pretraining language's, tokenizer work will likely beat model-tweaking. It's unglamorous, high-leverage, and under-published.
A model that says "90% confident" should be right 90% of the time — that's calibration. Pretrained transformers are often overconfident (label smoothing, Chapter 7, helps). For high-stakes applications (medical, legal), report expected calibration error (ECE) alongside accuracy: a well-calibrated 85% model beats a cocky 87% one in deployment. Temperature scaling (Chapter 2's temperature knob, fit on validation data) is the one-line fix. Calibration plots are rare in student papers and impressive when present.
Before building on a pretrained model, check three things: (1) the model card — intended uses, known biases, evaluation details; (2) the license — many research models forbid commercial use; (3) the training data cutoff — a 2021 model knows nothing of 2024 events, which matters for news or medical applications. Citing the model card in your paper (not just the weights URL) is becoming standard practice. Boring? Yes. The kind of boring reviewers reward.
The most embarrassing failure in applied ML: test examples (or near-duplicates) hiding in the training set — inflated scores, paper retracted in spirit if not in fact. Transformer-era leakage sources: web-scale pretraining corpora that contain your test set (common with public benchmarks!), duplicated documents across splits, and temporal leakage (training on 2024 data to "predict" 2023 events). Defenses: deduplicate aggressively, split by time or source (not randomly) when appropriate, and run a canary test — insert a unique nonsense sentence into training and check whether the model reproduces it (if yes, memorization is in play). Report your dedup procedure in one sentence; it signals you thought about it.
With modern models you have three adaptation strategies, not two:
Decision guide: hundreds of labels + one GPU → fine-tune (or LoRA). No labels + API budget → prompt. Need a deployable classifier with latency guarantees → feature extraction or fine-tuned small model. Many published comparisons now evaluate all three — and the honest winner varies by task, which is itself a publishable finding.
For your research: A fine-tuning study on an under-served domain or language is one of the most reliable student publications: take a pretrained model, adapt it to (say) Urdu legal text, crop-disease Q&A, or local-language sentiment, evaluate rigorously, and release the model + dataset. The contribution is the resource and benchmark, not architectural novelty — and the community genuinely needs these. Frame it as such: "we present the first benchmarked transformer baseline for X," with error analysis showing where it fails. That error analysis is often the seed of your next paper.
Key takeaways: - Transfer learning: pretrain once on massive data (masked/next-token prediction), adapt cheaply to your task. - Hugging Face gives uniform access to models, tokenizers, pipelines (Wolf et al., 2020). - Feature extraction (frozen, cheap, strong baseline) vs. fine-tuning (unfrozen, small LR, needs more data) — run both, report both. - Match tokenizer to model; watch max length, casing, and language coverage. - Evaluate honestly: fixed splits, multiple seeds, frozen baseline included.
You understand the transformer now. The question every MS/PhD student asks next: what can I actually work on? This chapter maps research directions by the resources they need — so you can pick something that fits a student lab, not a tech giant's cluster.
1a. Domain adaptation studies. Take a pretrained transformer to an under-served domain or language: Urdu news classification, agricultural Q&A, local legal documents, medical notes in a regional language. Contribution: the first rigorous benchmark + released model/dataset. Cost: fine-tuning a base model fits on a single mid-range GPU.
1b. Attention-head interpretability in a new domain. Visualize and catalog what heads learn in (say) a code model, a protein model, or a multilingual model on your language. Do heads specialize the same way as in English? Publish the catalog + analysis. Cost: inference only — you don't even need to train.
1c. Efficiency benchmarking on real long-sequence tasks. Implement FlashAttention vs. Performer vs. sparse attention on your domain's long sequences (documents, sensor streams). Report quality-vs-memory-vs-speed honestly, including where approximations break. Negative results with explanations are publishable.
1d. Probing: what does the model know? Train tiny classifiers on frozen layer representations to test whether the model encodes (say) grammatical number, sentiment, or factual relations at each layer. Probing papers are cheap, principled, and frequently cited.
2a. Custom positional/tokenization schemes. If your data has structure — timestamps, 2-D layouts, hierarchies — design a positional encoding or tokenization that reflects it (Chapter 5's research box). Ablate against sinusoidal and learned baselines.
2b. Distillation for deployment. Distill a large transformer into a small one for a real constraint: a mobile app, an offline clinic, a low-power sensor gateway. Report the quality/size/latency trade-off curve. Practitioners cite these.
2c. Robustness studies. How does your fine-tuned model handle typos, dialect variation, or adversarial paraphrases in your language? Systematic robustness evaluation with a clear failure taxonomy is a solid workshop-to-conference path.
2d. Data-efficient fine-tuning. Compare full fine-tuning vs. adapters vs. prompt-tuning on small datasets in your domain. Adapters (small trainable modules inserted per layer, freezing the rest) are especially student-friendly: cheap to train, easy to share.
3a. New efficient-attention mechanisms. Only after mastering Chapter 9's landscape — your idea must beat FlashAttention-backed baselines, which is a high bar. Start by reproducing one existing method first.
3b. Multimodal fusion. Combine text with images, sensor data, or structured records in one transformer for a domain problem (e.g., crop disease: leaf photo + farmer's text description). Fusion architecture choices are still an open design space.
3c. Interpretability beyond attention maps. Attention weights alone don't explain decisions (Chapter 2's warning). Methods like integrated gradients, causal interventions on heads, or concept-based explanations applied to transformers in new domains are active, respected research.
For any direction above: find the 3–5 most-cited papers (Google Scholar, sorted by citations), read them with Chapter 12's method, then read their citing papers (the "cited by" link) to see what's already been tried — that's your novelty check. A direction is viable when you can name both the baseline you'd beat and the reason it might fail on your data.
Most student papers start as course projects. The upgrade path:
Steps 2–4 are what separate an exercise from research. Budget equal time for them as for the initial implementation.
Start with a workshop in your direction's community. The reviews will teach you what the field expects — cheaper tuition than a conference rejection cycle.
Related work isn't a list — it's an argument that your paper fills a gap. The formula per paragraph: "Line of work A does P well but assumes Q. Line of work B handles Q but not at scale R. We combine / we address R for domain D." Three to five paragraphs, each ending near your contribution. Common student mistake: summarizing papers one by one ("Smith et al. did X. Jones et al. did Y."). Instead, organize by approach family and show the gap your work fills. Reviewers skim related work looking for exactly one thing: do you know what came before, and is your delta real?
"Method X didn't work on our data" is publishable if you answer why. The template: (1) reproduce X faithfully on its home turf (show you implemented it right), (2) apply it to your domain with the same rigor, (3) diagnose the failure — e.g., "Performer's random-feature variance dominates on our short sequences because the kernel approximation needs n ≫ features," (4) propose the fix or the boundary condition. A negative result with a mechanism is a contribution; a negative result that's just "we tried and it failed" is a blog post. Chapters 9's benchmarking guide gives you the experimental backbone.
"I reproduced X and found Y" is underrated as a first publication. Pick a well-cited transformer paper whose code is public, reimplement the key result from the paper alone (not by copying the repo), and document every gap between paper and reality: unstated hyperparameters, preprocessing details, hardware-dependent choices. Then extend modestly — a new dataset, a harder setting. Replication notes are publishable at reproducibility workshops (a growing venue family), teach you the paper's innards better than any reading, and produce the verified-ground foundation Chapter 12 advocates. The field has a replication crisis precisely because too few students do this.
The best direction is one your supervisor can evaluate: if they work in NLP, a vision project leaves you unsupervised in the ways that matter (they can't spot a weak baseline). Bring them two candidate directions from this chapter's tiers with the five-point pick procedure (Chapter 11) filled in — question, baseline, dataset, metric, failure mode — and let them react to something concrete. Also align with your timeline: an MS thesis (~6 months) fits one Tier-1 project done well; a PhD's first year fits Tier 1 → Tier 2 progression. A project that needs data you can't get is not ambitious, it's stuck — kill it early.
Student norms worth knowing: the person who did the core work is typically first author; the supervisor is usually last. Discuss authorship before the work, not after — a five-minute conversation prevents most disputes. If you build on a labmate's codebase or dataset, credit them (co-authorship or prominent acknowledgment, agreed upfront). And when you release models or datasets, use your institutional affiliation consistently — your publication record is your career's compounding asset; keep it tidy from paper one.
One document per project: date, what you tried, the exact command/config, the result, one-line interpretation. Failed runs included — especially failed runs. Six months later, when a reviewer asks "did you try X?", your journal answers in seconds. When you write the paper's ablation section, the journal is the first draft. Tools don't matter (markdown file, lab notebook); the habit does. Every senior researcher you admire keeps one in some form.
Don't start training blind. The standard estimate for transformer training compute is C ≈ 6·N·D floating-point operations, where N is parameter count and D is training tokens (Kaplan et al., 2020): ~2ND for the forward pass, ~4ND for backward. Example: a 100M-parameter model on 1B tokens ≈ 6×10¹⁷ FLOPs. A GPU delivering 100 TFLOPs/s (realistic sustained, not peak) needs ~6×10¹⁷ / 10¹⁴ ≈ 6,000 seconds ≈ 1.7 hours — plus overhead, evaluation, and failed runs (budget 3×). This arithmetic tells you before coding whether your plan fits your hardware and deadline. Reviewers also use it to sanity-check your claims: "trained 1B parameters on 2 GPUs in a day" fails the 6ND test and invites skepticism.
Memory rule of thumb alongside it: weights (2 bytes in fp16) + gradients (2–4 bytes) + Adam states (8 bytes: m and v in fp32) ≈ 12–16 bytes per parameter for full fine-tuning — a 1B-parameter model needs ~12–16 GB just for optimizer states. That's why LoRA (which trains <1% of parameters) changes what's possible on student hardware: the optimizer-state bill collapses.
Before starting a direction, write one page with five headings: Question (one sentence), Baseline (the method you'll beat, with its reported number), Data (what you'll train/evaluate on, and how you'll get it), Metric (the single number that decides success), Failure mode (why it might not work). If any heading takes more than three sentences, the project is too vague. Share this page with your supervisor or a peer — ten minutes of their confusion now saves a month of yours later. Every project in this chapter's tiers fits on one page; if yours doesn't, split it into two projects.
For your research: This chapter is the research box. Copy the tier that matches your resources into your notes, pick one direction, and spend two weeks on the literature check before writing code. The most common student failure mode is not weak ideas — it's starting to code before knowing the baselines. A month of reading saves three months of training the wrong thing.
Key takeaways: - Student-viable research clusters around: domain adaptation, interpretability, efficiency benchmarking, probing (Tier 1); custom encodings, distillation, robustness, adapters (Tier 2). - Pick using unfair advantages (data, experts, hardware), the smallest real question, and a writable-in-advance abstract. - Always establish baselines first; "first rigorous benchmark for X" is itself a contribution. - Read cited-by chains to check novelty before coding.
You are ready for the source. Vaswani et al.'s "Attention Is All You Need" (NeurIPS 2017) is unusually readable for a landmark paper — eight pages, one big idea, clear ablations. This chapter walks you through it section by section, telling you what to notice and what to skip on first reading.
One sentence: a sequence model built only from attention — no recurrence, no convolution — trains faster and translates better than the best RNN systems of 2017. Everything in the paper serves that claim: the architecture (Section 3), why it should work (the complexity table), and the proof (translation results + ablations).
The abstract states the motivation in two lines: recurrence prevents parallelization; attention alone suffices. The introduction adds the key promise — "we achieve superior quality while being more parallelizable and requiring significantly less time to train." Notice the rhetorical move: they don't claim attention is a better idea than recurrence philosophically; they claim it wins on the metrics practitioners care about (quality, speed). When you write papers, imitate this: lead with the measurable win.
Half a page on RNN encoder–decoders and Bahdanau attention — exactly our Chapters 1–2. If those chapters made sense, you'll skim this. Note what they don't do: no lengthy related-work battle. They position against the dominant paradigm crisply and move on.
3.1 Encoder/Decoder stacks: N=6 layers each, d_model=512. You've built this in Chapter 6 — verify your mental model against theirs. Notice the residual + layer-norm wrapping ("we employ residual connections around each sub-layer, followed by layer normalization").
3.2 Attention: the scaled dot-product formula (your Chapter 3) and multi-head attention (your Chapter 4), with h=8, d_k=d_v=64. Read their justification for multiple heads carefully: "it allows the model to jointly attend to information from different representation subspaces." That's the whole argument — short, and enough.
3.3 Position-wise FFN: two linear layers with ReLU, d_ff=2048 — your Chapter 6's "computation within a position."
3.4 Embeddings and softmax: learned embeddings shared between encoder/decoder and the output layer (weight tying — a nice parameter saving to notice).
3.5 Positional encoding: the sinusoidal scheme from your Chapter 5, including the hypothesis that it helps the model "attend by relative positions." They even tested learned vs. sinusoidal and found them nearly identical — an honest ablation students should emulate (report the comparison even when it's a tie).
Table 1 compares self-attention vs. recurrent vs. convolutional layers on three axes: complexity per layer, sequential operations, and maximum path length (how many steps separate distant tokens — the vanishing-gradient story from Chapter 1). Self-attention: O(n²·d) per layer but O(1) sequential operations and O(1) path length. This table is the paper's theoretical argument complementing its empirical results. When you write architecture papers, build this table for your method — reviewers love it.
They give three explicit advantages: (1) lower per-layer complexity when n < d (usually true), (2) more parallelizable, (3) shorter paths between distant positions → easier learning of long-range dependencies. Each maps to a Chapter 1 problem. This is how you argue for an architecture: enumerate the old problems, show the new design dissolves each.
Note the specifics: Adam with the warmup schedule (warmup_steps=4000), dropout 0.1 everywhere including on attention weights and embeddings, label smoothing ε=0.1. Training data: WMT 2014 English–German (4.5M pairs) and English–French (36M pairs), byte-pair encoded (~37k shared vocabulary). Hardware honesty: base model trained 12 hours on 8 P100 GPUs; big model 3.5 days. They report what they used — imitate this transparency.
Table 2 (translation quality): the big model scores 28.4 BLEU on English–German and 41.0 on English–French (base: 27.3 / 38.1) — state of the art in 2017, at a fraction of the training cost of competitors. Note they report both base and big: the base shows the idea works cheaply; the big shows it scales.
Table 3 (parsing): they apply the same architecture to constituency parsing with minimal changes — evidence of generality, a strong move for any architecture paper (show it works on two tasks).
The ablations (varying heads, d_k, dropout, etc.) are brief but present. First-reading advice: check that ablations exist and what they vary; study the numbers only if you're replicating.
Three sentences restated: attention-only works, it's fast, future work is images/audio/video — which the field then spent five years doing (your Chapter 8). Landmark papers often end by pointing at the next mountain; notice how this one did, correctly.
The fastest way to trust the paper: build a tiny transformer (2 layers, 4 heads, d_model=128) and train it on the copy task — input a random digit sequence like [3, 7, 1, 9], output the same sequence. A transformer learns this in minutes on a CPU. Then run three ablations:
Each ablation reproduces a paper claim at toy scale. Write up the three plots with two paragraphs each — that's a blog post that will teach others, and the habit that will make your future paper ablations rigorous.
The paper doesn't handle very long sequences (quadratic cost — your Chapter 9), doesn't explain why attention trains well (warmup was empirical), and its generality claim rests on two tasks (translation, parsing). Recognizing these boundaries isn't criticism — it's how you find the next paper. Every limitation listed here became someone's research program: long sequences → efficient attention; training dynamics → optimization theory; generality → BERT, GPT, ViT.
Read them in this order with the four-pass method from this chapter, and you'll have the complete 2017–2020 story — the foundation everything since builds on.
Let's practice skeptical reading on Table 2 (translation results). The big model: 28.4 BLEU English–German, 41.0 English–French. Ask four questions:
Run these four questions on every results table you read. They take five minutes and catch most over-claiming.
A landmark paper is a node in a graph. Two moves exploit this:
If you read this paper with peers (recommended — arguing about papers is how understanding hardens), use these:
The title is a provocation — a riff on the Beatles' "All You Need Is Love," and a dare to the field: recurrence was never necessary. Great paper titles do argumentative work: this one tells you the thesis before you read a word. When you title your own papers, imitate the confidence but earn it: the title's claim must be the thing your experiments actually prove. "Attention Is All You Need" survived because the BLEU tables backed the swagger. A bold title with weak evidence is the fastest route to harsh reviews.
For your research: Reproduce a small result from this paper before inventing anything — e.g., train a tiny transformer on a toy translation or copying task and verify that removing positional encoding breaks it, or that one head underperforms eight. A reproduced ablation, documented in your notes or a blog post, is the foundation every strong researcher builds on: you now trust the paper because you've touched its claims. Then your novel idea starts from verified ground, not from hope.
Key takeaways: - The paper's claim: attention-only models translate better and train faster than RNN systems. - Read architecture papers in passes: claim first (abstract/figures/tables), method with pen and paper, training details only when replicating. - Study Table 1 (complexity/path length) as the template for arguing architectural advantages. - Note the honest reporting: base and big results, learned-vs-sinusoidal tie, exact hardware and data. - Reproduce one small ablation yourself before building on any paper.
| Step | Operation | Shape (single head) | Purpose |
|---|---|---|---|
| 1. Project | Q = XW^Q, K = XW^K, V = XW^V | (n, d_k) each | Give each token query/key/value roles |
| 2. Score | S = QK^T | (n, n) | Pairwise relevance: how much i attends to j |
| 3. Scale | S / √d_k | (n, n) | Keep softmax in its healthy range |
| 4. Mask (optional) | S[forbidden] = −∞ | (n, n) | Enforce causality in decoders |
| 5. Normalize | W = softmax(S) rows | (n, n) | Weights summing to 1 per row |
| 6. Blend | Output = W·V | (n, d_k) | Weighted mix of values |
Multi-head: run h copies with d_k = d_model/h, concatenate → project with W^O (d_model × d_model). Full formula: Attention(Q,K,V) = softmax(QK^T/√d_k)·V. Cost: O(n²·d) time, O(n²) memory.
| RNN (LSTM) | Transformer encoder | Transformer decoder | Encoder–decoder | |
|---|---|---|---|---|
| Order of processing | Sequential, one token at a time | All tokens in parallel | Parallel in training, sequential in generation | Encoder parallel; decoder as left |
| Long-range links | Weak (vanishing gradient) | Direct, O(1) path | Direct, O(1) path (causal) | Direct via cross-attention |
| Positional info | Built into recurrence | Injected (positional encoding) | Injected | Injected |
| Cost per layer | O(n·d²) | O(n²·d) | O(n²·d) | O(n²·d) |
| Classic use | Old translation, speech | BERT: classification, QA | GPT: generation, chat | T5: translation, summarization |
| Key weakness | No parallelism, forgets | Quadratic in length | Quadratic; slow inference | Most parameters, complex |
| Term | Plain meaning |
|---|---|
| Token | The model's unit of input — a word piece, image patch, or similar chunk |
| Embedding | A vector of numbers representing a token's meaning |
| d_model | The width of the model (e.g., 512) — size of every token vector |
| Head | One independent attention operation; models use 8–96 in parallel |
| d_k | Key/query dimension per head (d_model ÷ heads, often 64) |
| Logits | Raw unnormalized scores before softmax (e.g., next-token scores) |
| Causal mask | Blocker that stops a token from seeing future tokens |
| Cross-attention | Decoder queries attending to encoder keys/values |
| KV cache | Saved keys/values from past steps to avoid recomputation during generation |
| Warmup | Gradually raising the learning rate at training start |
| Fine-tuning | Training a pretrained model further on your task |
| Distillation | Training a small model to mimic a large one |
| BLEU | Classic translation quality score (higher = better) |
| Perplexity | How "surprised" a language model is by text (lower = better) |
| Hyperparameter | Original paper | BERT-base | Your small experiments |
|---|---|---|---|
| Layers | 6 + 6 | 12 | 2–4 |
| d_model | 512 | 768 | 128–256 |
| Heads | 8 | 12 | 4 |
| d_k per head | 64 | 64 | 32–64 |
| d_ff | 2048 | 3072 | 4 × d_model |
| Dropout | 0.1 | 0.1 | 0.1–0.3 |
| Warmup steps | 4,000 | 10,000 | 5–10% of total steps |
| Peak LR (from scratch) | ~7e-4 (formula) | 1e-4 | 1e-4–5e-4 |
| Peak LR (fine-tune) | — | — | 1e-5–5e-5 |
| Batch (tokens) | ~25,000 | 131,072 | as large as fits + accumulation |
| Optimizer | Adam (β₂=0.98) | AdamW | AdamW |
Start from the right column; move toward the middle only when results justify the compute.
[1] A. Vaswani et al., "Attention is all you need," in Proc. Adv. Neural Inf. Process. Syst., vol. 30, Long Beach, CA, USA, 2017, pp. 5998–6008.
[2] J. Devlin, M.-W. Chang, K. Lee, and K. Toutanova, "BERT: Pre-training of deep bidirectional transformers for language understanding," in Proc. Conf. North Amer. Chapter Assoc. Comput. Linguistics (NAACL-HLT), Minneapolis, MN, USA, 2019, pp. 4171–4186.
[3] A. Dosovitskiy et al., "An image is worth 16x16 words: Transformers for image recognition at scale," in Proc. Int. Conf. Learn. Representations (ICLR), Vienna, Austria, 2021.
[4] D. Bahdanau, K. Cho, and Y. Bengio, "Neural machine translation by jointly learning to align and translate," in Proc. Int. Conf. Learn. Representations (ICLR), San Diego, CA, USA, 2015.
[5] J. L. Ba, J. R. Kiros, and G. E. Hinton, "Layer normalization," arXiv preprint arXiv:1607.06450, 2016.
[6] K. He, X. Zhang, S. Ren, and J. Sun, "Deep residual learning for image recognition," in Proc. IEEE Conf. Comput. Vis. Pattern Recognit. (CVPR), Las Vegas, NV, USA, 2016, pp. 770–778.
[7] S. Hochreiter and J. Schmidhuber, "Long short-term memory," Neural Computation, vol. 9, no. 8, pp. 1735–1780, 1997.
[8] T. B. Brown et al., "Language models are few-shot learners," in Proc. Adv. Neural Inf. Process. Syst., vol. 33, 2020, pp. 1877–1901.
[9] T. Dao, D. Y. Fu, S. Ermon, A. Rudra, and C. Ré, "FlashAttention: Fast and memory-efficient exact attention with IO-awareness," in Proc. Adv. Neural Inf. Process. Syst., vol. 35, New Orleans, LA, USA, 2022, pp. 16344–16359.
[10] K. Choromanski et al., "Rethinking attention with Performers," in Proc. Int. Conf. Learn. Representations (ICLR), Vienna, Austria, 2021.
[11] S. Wang, B. Z. Li, M. Khabsa, H. Fang, and H. Ma, "Linformer: Self-attention with linear complexity," arXiv preprint arXiv:2006.04768, 2020.
[12] T. Wolf et al., "Transformers: State-of-the-art natural language processing," in Proc. Conf. Empirical Methods Natural Lang. Process. (EMNLP): System Demonstrations, 2020, pp. 38–45.
Exercise 1 — Hand-compute attention (worked in Chapter 3). Using Q rows q₁=[1,0], q₂=[0,1], q₃=[1,1]; K rows k₁=[1,0], k₂=[0,1], k₃=[1,1]; V rows v₁=[4,0], v₂=[0,6], v₃=[2,2]: compute the self-attention output row for token 2 (q₂) by hand — scores, scaled scores, softmax weights, blended output. Then verify with the PyTorch snippet from Chapter 3. Solution: scores = [0,1,1]; scaled = [0, 0.7071, 0.7071]; softmax weights ≈ [0.1978, 0.4011, 0.4011]; output = 0.1978·[4,0] + 0.4011·[0,6] + 0.4011·[2,2] ≈ [1.59, 3.21].
Exercise 2 — Mask it. Take the same Q, K, V. Apply a causal mask so token 2 can only attend to tokens 1 and 2 (set score(2,3) = −∞ before softmax). Recompute token 2's output by hand. How does it change, and why does this matter for generation? Solution: scores [0, 0.7071, −∞] → weights ≈ [0.3302, 0.6698, 0] → output ≈ [1.32, 4.02]. The future token contributes nothing — during training the model can't cheat by copying the answer.
Exercise 3 — One head vs. two. Write down a 4-token sentence where a single attention head cannot simultaneously resolve a pronoun's referent and a verb's subject, but two heads could. Explain which query–key match each head would learn. (Research angle: this is the intuition behind head-specialization studies.)
Exercise 4 — Positional encoding by hand. With d_model = 4, compute the sinusoidal positional vectors for positions 0, 1, 2, 3 (formulas in Chapter 5). Verify all four are distinct. Then explain in two sentences why adding (not concatenating) is the standard choice.
Exercise 5 — Build a micro-transformer. In PyTorch, implement a 2-layer, 2-head transformer encoder (d_model=64) for a toy task: binary sentiment on 200 short sentences you write yourself. Train it, then visualize one head's attention on three examples. Write 300 words on what the heads appear to track. (Research angle: this is a miniature interpretability paper.)
Exercise 6 — Warmup ablation. Train the Exercise 5 model twice: once with linear warmup (500 steps), once with a constant learning rate. Plot both loss curves. Which trains more stably? Write down the exact failure mode you observe without warmup (divergence? plateau?). This is the Chapter 7 debugging skill in action.
Exercise 7 — Frozen vs. fine-tuned. Pick a small text-classification dataset (or 500 labeled examples from your domain). Compare (a) frozen BERT-base features + logistic regression vs. (b) full fine-tuning. Report accuracy, training time, and variance across 3 seeds. When does (b) justify its cost? Frame your answer as a one-paragraph "experiment" section for a paper.
Exercise 8 — Research design. Choose one Tier-1 direction from Chapter 11 for your field. Write: (i) the research question in one sentence, (ii) the baseline you'd compare against, (iii) the dataset, (iv) the metric, (v) one reason the idea might fail. If you can't fill in (v), the question is too vague — refine it.
Exercise 9 — Paper dissection. Read "Attention Is All You Need" Sections 3–4 with the guided method from Chapter 12. Then answer: (a) Why do the authors divide by √d_k — what goes wrong without it? (b) What does Table 1's "maximum path length" column argue, in your own words? (c) Name one ablation the authors ran and what it showed.
Exercise 10 — Efficiency trade-off analysis. For a sequence length n = 16,384 and d = 512, estimate the attention matrix memory (float32) for standard attention, and the rough reduction Linformer would give with k = 256. Then write 200 words on which efficient method (Chapter 9) you'd pick for classifying 16k-token legal documents, and why FlashAttention alone might already suffice. (Research angle: honest cost modeling is the backbone of efficiency papers.)
End of Book 14 — "Transformers Explained for Beginners." Next in the series: Book 15.