Multi-head attention, step by step¶
Begin with colour/material/detail, subject/location, and object/event examples. These illustrate possible head roles, not measured attention. Then return to Part II’s river-bank example and compute two complete head messages. Run the implementation lab, including the batch axes and a comparison with PyTorch.
The toy uses hand-chosen parameters. The final links lead to genuinely trained TinyStories models, not this worksheet. As in Part II, E stores embedding rows, M is the mask, H = AV stores message rows, and E′ = E + ΔE. Superscripts label heads; subscripts label tokens.
Part III · Live models · Download code and notebooks
Run all cells from this directory. No data download or training run is needed. The single optimizer step demonstrates learning; it does not create the models used in the benchmark.
import json
import math
from pathlib import Path
import torch
from torch import nn
from torch.nn import functional as F
from IPython.display import display, HTML
from multihead_from_scratch import (TinyMultiHeadLM, ScratchMultiHead,
load_worksheet_weights, copy_to_pytorch)
torch.manual_seed(7)
worksheet = json.loads(Path('multihead-worksheet.json').read_text())
word_to_id = {word: i for i, word in enumerate(worksheet['vocab'])}
river_ids = [word_to_id[w.lower()] for w in worksheet['sentences']['river']]
cheque_ids = [word_to_id[w.lower()] for w in worksheet['sentences']['cheque']]
print('River tokens:', worksheet['sentences']['river'])
print('River IDs:', river_ids)
River tokens: ['The', 'fisherman', 'sat', 'beside', 'the', 'river', 'bank', 'and', 'watched', 'the'] River IDs: [0, 1, 2, 3, 0, 4, 5, 6, 7, 0]
The visual story¶
The same figures appear in the lecture. Short code excerpts here are explained visually; the executable lab below builds their inputs and runs each operation.
Back to the river bank¶
Which parts of this prefix would help you choose a continuation?
One token can need several kinds of detail¶
At “coat”, colour is useful—but so are material and shape. One head can read several words. The question is whether we want the same mixture for every feature.
These next examples illustrate possible readings, not attention measured in our trained language model. A query is a numerical vector, not an English question. Each source also has a numerical key and value. The arrows highlight plausible useful sources; they do not mean all other sources are masked or have exactly zero weight. In deeper layers, source rows can already contain context.
Three heads can keep three readings separate¶
One head could favour red, another wool, another hooded. Each receives the same input rows but can return a different weighted message. The output projection combines those messages.
This is a possible division of work, not three manually assigned jobs. All three heads see all seven rows, including coat. Head-specific projections can expose different matching and value features. A single head may also encode several attributes; the extra capability is independent source weighting, not a rule that each head can understand only one concept.
The nearest noun is not always the subject¶
The plural subject and the location are both useful clues. Separate heads can preserve them without having to favour the same source.
The receiver is usually, the last known token—not the as-yet-unknown word bark. “The dogs near the gate usually bark at visitors” is one possible continuation. Dogs supports bark, not barks; the nearer singular noun gate does not control subject agreement. Both useful source words occur before the receiver. These semantic roles are illustrative, not measured trained heads.
A reference and an event are different clues¶
At “It”, one reading could retrieve the object; another could retrieve what happened to it. Combining glass and dropping makes “broke” a plausible continuation, not a certainty.
We treat punctuation as a token in this illustration: It is receiver 6. Both useful sources are in its known prefix. The two heads run in parallel, so the event head does not wait for the reference head’s result. Contextual input rows from earlier layers can help make these matches.
Change the event; keep the object¶
The glass is unchanged. “Washed” supports a different continuation, such as “gleamed”. Separate messages carry both the object clue and the event clue.
The arrows are a schematic of useful information, not a promise that a trained attention map remains unchanged after a word substitution. All head projections learn jointly from prediction loss. Interpretable roles sometimes emerge, but some heads overlap or can be pruned; see Voita et al. (2019). More heads are a capacity choice to test, not an automatic improvement.
One head gives “the” one weighted message¶
One query chooses how much to read from each token. The head mixes their value vectors into one message for the final “the”.
The prefix is “The fisherman sat beside the river bank and watched the ___”. The final known token, the at slot 10, is the receiver. The blank has no query yet. One head can read several words and carry several features in its message. It uses one shared set of source weights for all coordinates of that message. “Setting clues?” describes a possible role of the numerical query, not a literal question or a manually assigned task.
Two heads give “the” two separate messages¶
Same receiver, two queries. Each head has its own keys, values and weights. One can favour river, the other fisherman. We combine their messages afterward.
Each head learns its own W_Q, W_K and W_V, applied to the same input rows. At this receiver that produces one query per head. Both heads can read all ten known tokens and run in parallel. Training can discover useful roles, but does not assign setting and person labels or guarantee that the heads specialize this way.
Optional: why separate weights can help
In a separate two-source toy, let river have value [10, 1] and fisherman [2, 8]. A single head with weights [0.8, 0.2] returns [8.4, 2.4]. Both coordinates use that same mixture. Two heads can instead select one value coordinate each: a setting head with weights [0.8, 0.2] gives 8.4, and a person head with weights [0.2, 0.8] gives 6.6. With these fixed values, one river weight cannot be both 0.8 and 0.2. This illustrates independent source weighting, not a proof that every one-head network fails. These invented numbers are separate from the full-sentence worksheet that follows.
# Optional arithmetic for the reading note, not the full-sentence model.
# This separate two-source illustration explains independent mixtures.
values = torch.tensor([[10., 1.], [2., 8.]]) # river, fisherman
setting_weights = torch.tensor([0.8, 0.2])
person_weights = torch.tensor([0.2, 0.8])
one_head = setting_weights @ values
two_heads = torch.stack([setting_weights @ values[:, 0],
person_weights @ values[:, 1]])
print('One shared mixture:', one_head.tolist())
print('Two separate mixtures:', two_heads.tolist())
torch.testing.assert_close(one_head, torch.tensor([8.4, 2.4]))
torch.testing.assert_close(two_heads, torch.tensor([8.4, 6.6]))
# For fixed values, one river weight cannot meet both targets.
a_for_setting = (8.4 - 2.) / (10. - 2.)
a_for_person = (6.6 - 8.) / (1. - 8.)
assert math.isclose(a_for_setting, 0.8)
assert math.isclose(a_for_person, 0.2)
assert not math.isclose(a_for_setting, a_for_person)
print('Required river weights:', a_for_setting, a_for_person)
One shared mixture: [8.399999618530273, 2.4000000953674316] Two separate mixtures: [8.399999618530273, 6.599999904632568]
Required river weights: 0.8 0.20000000000000004
Setting clues in the full sentence¶
Back to all ten tokens. Our hand-chosen head 1 gives river a large weight. Its values carry setting information to the receiver.
These are hand-chosen two-head parameters applied to Part II’s exact input rows. “Setting” and “person” name the intended behaviour of this example, not jobs assigned to trained heads. Arrows show information moving from source to receiver; their widths encode computed attention weights.
Another head can read who is there¶
The sentence and receiver stay fixed. A different head gives fisherman a large weight.
Keep both readings¶
Two heads can retain different source mixtures at the same token. They run in parallel, not one after the other.
How does each head choose what to read?¶
Visual inspiration: 3Blue1Brown’s attention lesson and Jay Alammar’s Illustrated Transformer. We keep our own river-bank example, numbers and Part II row-vector convention. A head-specific superscript labels a head; a subscript still labels a token.
Recall the Maya example: query, key and value¶
In Part II’s illustrative later-layer example, the query asks for a person. Maya’s key can match that request. Her value supplies useful details. A head uses separate projections for these roles.
This is the verbal example from Part II. It illustrates possible later-layer representations, not measured model outputs. Earlier layers may have gathered the preceding facts into the second Maya row. Every token has a query, key and value, even though this example follows only She’s query and Maya’s key/value. The actual vectors contain numbers, not written questions or records.
The two heads and the output projection¶
Each head keeps all 10 tokens. (A) weights source tokens; (H=AV) contains their messages. Concatenation joins coordinates, so the result is (10\times4), not (20\times2).
This diagram shows one sentence, without a batch axis. E has shape [10, 4]. Each head has separate W_Q, W_K and W_V of shape [4, 2], producing Q, K and V each [10, 2]. QKᵀ multiplies [10, 2] by [2, 10] to give [10, 10] scores. Scaling, causal masking and row-wise softmax preserve that shape. A has one row per receiver and one column per source. H = AV multiplies [10, 10] by [10, 2] to produce [10, 2]: one message per receiver. The final receiver’s individual q, k, v and message rows are each [1, 2].
H¹ and H² concatenate to [10, 4]. W_O [4, 4] gives ΔE [10, 4], then E′ = E + ΔE is also [10, 4]. Its last row [1, 4] feeds the prediction MLP. Here the joined width already equals the embedding width. W_O still learns to mix head outputs into embedding coordinates. In general its shape is [n_heads × d_v, d_model]. It is not an inverse of the input projections.
The same roles in the river example¶
Head 1 can ask about the setting; head 2 can ask about the person. These are interpretations of our chosen numbers. Training learns the projections without assigning these jobs.
Head 1: computing the setting query¶
In this toy, the final “the” has only a nonzero glue coordinate. The two columns of (W_Q) turn that input into water and finance matching features.
E contains position-aware embedding rows, not token IDs. The water, finance, person and glue axes are invented for teaching; glue is our toy feature for function words such as “the”. W_Q reads all four input coordinates. Each query coordinate is a dot product with one column of W_Q. These sparse matrices are hand-chosen; learned projections need not be sparse or interpretable. The head superscripts label heads, not powers. Query/key width is two per head here; Part II’s single-head toy used three.
Head 2: computing the person query¶
The input embedding stays [0, 0, 0, 2.3]. Head 2 uses its own (W_Q). Its query requests the person feature and gives zero weight to the glue matching feature.
Source rows supply keys and values¶
Each head has its own (W_K) and (W_V), both (4\times2) here. For easy arithmetic, they select the same coordinates. Keys determine matching; values carry the numbers we mix.
Head 1 uses W_K = W_V = [[1,0],[0,1],[0,0],[0,0]], selecting water and finance. Head 2 uses W_K = W_V = [[0,0],[0,0],[1,0],[0,1]], selecting person and glue. Each matrix is [4,2] and acts on every source row, including sources not shown here. W_K and W_V are distinct parameters with equal numerical entries in this worksheet. They need not be equal in a trained model. In Part II’s Maya example they served different roles too.
# Every projected row is computed from an input embedding.
case = worksheet['headsLesson']['cases']['river']
E = torch.tensor(case['E'])
for h, projection in enumerate(worksheet['headsLesson']['projections']):
Q, K, V = [E @ torch.tensor(projection[kind], dtype=torch.float32)
for kind in ['Q', 'K', 'V']]
source = 5 if h == 0 else 1 # river or fisherman, zero-based index
print(f'Head {h+1}: final query', Q[-1].tolist())
print('Source:', case['tokens'][source],
'key:', K[source].tolist(), 'value:', V[source].tolist())
for kind, actual in [('Q', Q), ('K', K), ('V', V)]:
torch.testing.assert_close(actual, torch.tensor(case['heads'][h][kind]))
Head 1: final query [2.299999952316284, 2.299999952316284] Source: river key: [3.0999999046325684, -0.10000000149011612] value: [3.0999999046325684, -0.10000000149011612] Head 2: final query [2.299999952316284, 0.0] Source: fisherman key: [2.0999999046325684, 0.0] value: [2.0999999046325684, 0.0]
Head 1: the query and key matrices¶
Every token has a query and a key. We follow the last query row, for the final “the”, and compare it with all ten key rows.
# Head 1: these are its own projection matrices.
head_index = 0
projection = worksheet['headsLesson']['projections'][head_index]
Q, K, V = [E @ torch.tensor(projection[kind], dtype=torch.float32)
for kind in ['Q', 'K', 'V']]
print('Q:', Q.shape)
print(Q)
print('K:', K.shape)
print(K)
q = Q[-1:] # final "the", kept as a [1, 2] row matrix
assert q.shape == (1, 2) and K.T.shape == (2, 10)
Q: torch.Size([10, 2])
tensor([[2.3000, 2.3000],
[0.0000, 0.0000],
[1.9000, 1.9000],
[2.0000, 2.0000],
[2.3000, 2.3000],
[0.1000, 0.1000],
[0.8000, 0.8000],
[2.2000, 2.2000],
[1.7000, 1.7000],
[2.3000, 2.3000]])
K: torch.Size([10, 2])
tensor([[ 0.1000, 0.0000],
[ 2.0000, 0.1000],
[ 0.3000, 0.0000],
[ 0.6000, -0.1000],
[ 0.1000, 0.1000],
[ 3.1000, -0.1000],
[ 0.7000, 0.7000],
[-0.1000, 0.1000],
[ 0.7000, 0.0000],
[ 0.0000, 0.0000]])
Head 1: one dot product per key¶
Transposing (K) puts each source key in a column. Multiplying the query row by (K^\top) gives one raw dot product per source.
raw = q @ K.T # [1, 2] @ [2, 10] -> [1, 10]
for j, word in enumerate(case['tokens']):
products = q[0] * K[j]
print(j + 1, word, 'coordinate products:', products.tolist(),
'sum:', raw[0, j].item())
torch.testing.assert_close(products.sum(), raw[0, j])
1 The coordinate products: [0.23000000417232513, 0.0] sum: 0.23000000417232513 2 fisherman coordinate products: [4.599999904632568, 0.23000000417232513] sum: 4.829999923706055 3 sat coordinate products: [0.6899999976158142, 0.0] sum: 0.6899999976158142 4 beside coordinate products: [1.3799999952316284, -0.23000000417232513] sum: 1.149999976158142 5 the coordinate products: [0.23000000417232513, 0.23000000417232513] sum: 0.46000000834465027 6 river coordinate products: [7.12999963760376, -0.23000000417232513] sum: 6.899999618530273 7 bank coordinate products: [1.6099998950958252, 1.6099998950958252] sum: 3.2199997901916504 8 and coordinate products: [-0.23000000417232513, 0.23000000417232513] sum: 0.0 9 watched coordinate products: [1.6099998950958252, 0.0] sum: 1.6099998950958252 10 the coordinate products: [0.0, 0.0] sum: 0.0
Head 1: turning scores into weights¶
Divide each dot product by √2, then apply softmax across these ten sources. The final query can see all ten tokens. The unrounded weights sum to one.
Every source contributes to the denominator, including the receiver itself. All displayed numbers are rounded; calculations use full precision. At receiver row 10, the causal mask permits source positions 1 through 10. Earlier query rows cannot read later sources. The notebook also computes the full masked attention matrix. A numerically stable softmax subtracts the maximum score before exponentiating; this leaves the normalized weights unchanged.
scores = raw / math.sqrt(q.shape[-1])
exponentials = scores.exp()
total = exponentials.sum(dim=-1, keepdim=True)
weights = exponentials / total
# Direct exponentials are safe for these small worksheet scores.
# torch.softmax is numerically stable for general inputs.
torch.testing.assert_close(weights, scores.softmax(dim=-1))
torch.testing.assert_close(weights.sum(dim=-1), torch.ones(1))
torch.testing.assert_close(weights[0], torch.tensor(case['heads'][head_index]['A'][-1]))
for j, word in enumerate(case['tokens']):
print(word, 'score:', round(scores[0, j].item(), 3),
'exp:', round(exponentials[0, j].item(), 3),
'weight:', round(weights[0, j].item(), 3))
print('Shared denominator:', total.item())
The score: 0.163 exp: 1.177 weight: 0.006 fisherman score: 3.415 exp: 30.427 weight: 0.166 sat score: 0.488 exp: 1.629 weight: 0.009 beside score: 0.813 exp: 2.255 weight: 0.012 the score: 0.325 exp: 1.384 weight: 0.008 river score: 4.879 exp: 131.504 weight: 0.718 bank score: 2.277 exp: 9.746 weight: 0.053 and score: 0.0 exp: 1.0 weight: 0.005 watched score: 1.138 exp: 3.122 weight: 0.017 the score: 0.0 exp: 1.0 weight: 0.005 Shared denominator: 183.24386596679688
Head 1: the weights for each source¶
(\alpha_{10,j}) tells us how much the final “the” reads from source (j). Keep these weights as we bring in the values.
These are the ten entries of row 10 of this head’s attention matrix A. We display them vertically so each weight can sit beside its source word and value row. The receiver stays fixed at token 10. Only the source index j changes.
# The final receiver uses row 10 of its head's attention matrix.
# Python index 9 is the tenth token. Each j below names a source.
assert weights.shape == (1, 10)
for j, word in enumerate(case['tokens']):
print(f'source {j + 1}: {word}, alpha = {weights[0, j].item():.3f}')
source 1: The, alpha = 0.006 source 2: fisherman, alpha = 0.166 source 3: sat, alpha = 0.009 source 4: beside, alpha = 0.012 source 5: the, alpha = 0.008 source 6: river, alpha = 0.718 source 7: bank, alpha = 0.053 source 8: and, alpha = 0.005 source 9: watched, alpha = 0.017 source 10: the, alpha = 0.005
Head 1: the value beside each weight¶
Each source supplies a two-number value row from this head’s V matrix. The source words and weights stay in place, so we can see which weight belongs to which value.
V = E W_V has one row per source. These are projected value vectors, not token IDs or attention weights. Our hand-chosen worksheet gives W_K and W_V equal entries, but they remain separate parameters with different roles. The keys helped compute the weights. V supplies the content to mix.
assert V.shape == (10, 2)
for j, word in enumerate(case['tokens']):
print(word, 'weight:', round(weights[0, j].item(), 3),
'value row:', V[j].tolist())
The weight: 0.006 value row: [0.10000000149011612, 0.0] fisherman weight: 0.166 value row: [2.0, 0.10000000149011612] sat weight: 0.009 value row: [0.30000001192092896, 0.0] beside weight: 0.012 value row: [0.6000000238418579, -0.10000000149011612] the weight: 0.008 value row: [0.10000000149011612, 0.10000000149011612] river weight: 0.718 value row: [3.0999999046325684, -0.10000000149011612] bank weight: 0.053 value row: [0.699999988079071, 0.699999988079071] and weight: 0.005 value row: [-0.10000000149011612, 0.10000000149011612] watched weight: 0.017 value row: [0.699999988079071, 0.0] the weight: 0.005 value row: [0.0, 0.0]
Head 1: each source’s contribution¶
Multiply both value coordinates by the same source weight. The new column shows exactly what each source contributes. Displayed numbers are rounded.
The calculation below the rows expands the highlighted source’s scalar–vector multiplication coordinate by coordinate. A small weight scales down the entire value vector. The sign of each coordinate comes from the value. All products use the unrounded weights.
# One scalar multiplies both coordinates of its source's value.
source = 5 if head_index == 0 else 1 # river or fisherman
alpha = weights[0, source]
value = V[source]
print(case['tokens'][source], ':', alpha.item(), '*', value.tolist())
print('Contribution:', (alpha * value).tolist())
# Line up the ten weights with the ten value rows. Broadcasting
# applies each [10, 1] weight to both columns of V [10, 2].
contributions = weights.T * V
assert contributions.shape == (10, 2)
for word, row in zip(case['tokens'], contributions):
print(word, row.tolist())
river : 0.7176441550254822 * [3.0999999046325684, -0.10000000149011612] Contribution: [2.2246968746185303, -0.07176441699266434] The [0.0006420988356694579, 0.0] fisherman [0.33209142088890076, 0.016604570671916008] sat [0.002666770713403821, 0.0] beside [0.007383772637695074, -0.0012306286953389645] the [0.000755497720092535, 0.000755497720092535] river [2.2246968746185303, -0.07176441699266434] bank [0.037231165915727615, 0.037231165915727615] and [-0.0005457208608277142, 0.0005457208608277142] watched [0.011925802566111088, 0.0] the [0.0, 0.0]
Head 1: adding the contributions¶
Add the first coordinates to get the first message coordinate. Add the second coordinates to get the second. All ten sources contribute. The sum uses unrounded products.
This is one row of H = AV: m₁₀ = Σⱼ α₁₀,ⱼ vⱼ. The original weight row has shape [1,10] and V has shape [10,2], giving a [1,2] message. We arranged the weights vertically only to make the source correspondence visible. The displayed products are rounded, but the sums use full precision.
# Sum down the ten sources, keeping the two coordinates separate.
message = contributions.sum(dim=0, keepdim=True)
print('First coordinate:', contributions[:, 0].sum().item())
print('Second coordinate:', contributions[:, 1].sum().item())
assert message.shape == (1, 2)
# The matrix product performs exactly the same multiply-and-sum.
torch.testing.assert_close(message, weights @ V)
torch.testing.assert_close(message[0], torch.tensor(case['heads'][head_index]['messages'][-1]))
print('Message:', message.tolist())
First coordinate: 2.616847515106201 Second coordinate: -0.01785808987915516 Message: [[2.616847515106201, -0.01785808987915516]]
Head 2: the query and key matrices¶
Every token has a query and a key. We follow the last query row, for the final “the”, and compare it with all ten key rows.
# Head 2: these are its own projection matrices.
head_index = 1
projection = worksheet['headsLesson']['projections'][head_index]
Q, K, V = [E @ torch.tensor(projection[kind], dtype=torch.float32)
for kind in ['Q', 'K', 'V']]
print('Q:', Q.shape)
print(Q)
print('K:', K.shape)
print(K)
q = Q[-1:] # final "the", kept as a [1, 2] row matrix
assert q.shape == (1, 2) and K.T.shape == (2, 10)
Q: torch.Size([10, 2])
tensor([[2.3000, 0.0000],
[0.0000, 0.0000],
[1.9000, 0.0000],
[2.0000, 0.0000],
[2.3000, 0.0000],
[0.1000, 0.0000],
[0.8000, 0.0000],
[2.2000, 0.0000],
[1.7000, 0.0000],
[2.3000, 0.0000]])
K: torch.Size([10, 2])
tensor([[0.1000, 2.3000],
[2.1000, 0.0000],
[0.6000, 1.9000],
[0.1000, 2.0000],
[0.0000, 2.3000],
[0.0000, 0.1000],
[0.1000, 0.8000],
[0.1000, 2.2000],
[0.7000, 1.7000],
[0.0000, 2.3000]])
Head 2: one dot product per key¶
Transposing (K) puts each source key in a column. Multiplying the query row by (K^\top) gives one raw dot product per source.
raw = q @ K.T # [1, 2] @ [2, 10] -> [1, 10]
for j, word in enumerate(case['tokens']):
products = q[0] * K[j]
print(j + 1, word, 'coordinate products:', products.tolist(),
'sum:', raw[0, j].item())
torch.testing.assert_close(products.sum(), raw[0, j])
1 The coordinate products: [0.23000000417232513, 0.0] sum: 0.23000000417232513 2 fisherman coordinate products: [4.8299994468688965, 0.0] sum: 4.8299994468688965 3 sat coordinate products: [1.3799999952316284, 0.0] sum: 1.3799999952316284 4 beside coordinate products: [0.23000000417232513, 0.0] sum: 0.23000000417232513 5 the coordinate products: [0.0, 0.0] sum: 0.0 6 river coordinate products: [0.0, 0.0] sum: 0.0 7 bank coordinate products: [0.23000000417232513, 0.0] sum: 0.23000000417232513 8 and coordinate products: [0.23000000417232513, 0.0] sum: 0.23000000417232513 9 watched coordinate products: [1.6099998950958252, 0.0] sum: 1.6099998950958252 10 the coordinate products: [0.0, 0.0] sum: 0.0
Head 2: turning scores into weights¶
Divide each dot product by √2, then apply softmax across these ten sources. The final query can see all ten tokens. The unrounded weights sum to one.
Every source contributes to the denominator, including the receiver itself. All displayed numbers are rounded; calculations use full precision. At receiver row 10, the causal mask permits source positions 1 through 10. Earlier query rows cannot read later sources. The notebook also computes the full masked attention matrix. A numerically stable softmax subtracts the maximum score before exponentiating; this leaves the normalized weights unchanged.
scores = raw / math.sqrt(q.shape[-1])
exponentials = scores.exp()
total = exponentials.sum(dim=-1, keepdim=True)
weights = exponentials / total
# Direct exponentials are safe for these small worksheet scores.
# torch.softmax is numerically stable for general inputs.
torch.testing.assert_close(weights, scores.softmax(dim=-1))
torch.testing.assert_close(weights.sum(dim=-1), torch.ones(1))
torch.testing.assert_close(weights[0], torch.tensor(case['heads'][head_index]['A'][-1]))
for j, word in enumerate(case['tokens']):
print(word, 'score:', round(scores[0, j].item(), 3),
'exp:', round(exponentials[0, j].item(), 3),
'weight:', round(weights[0, j].item(), 3))
print('Shared denominator:', total.item())
The score: 0.163 exp: 1.177 weight: 0.027 fisherman score: 3.415 exp: 30.427 weight: 0.693 sat score: 0.976 exp: 2.653 weight: 0.06 beside score: 0.163 exp: 1.177 weight: 0.027 the score: 0.0 exp: 1.0 weight: 0.023 river score: 0.0 exp: 1.0 weight: 0.023 bank score: 0.163 exp: 1.177 weight: 0.027 and score: 0.163 exp: 1.177 weight: 0.027 watched score: 1.138 exp: 3.122 weight: 0.071 the score: 0.0 exp: 1.0 weight: 0.023 Shared denominator: 43.908485412597656
Head 2: the weights for each source¶
(\alpha_{10,j}) tells us how much the final “the” reads from source (j). Keep these weights as we bring in the values.
These are the ten entries of row 10 of this head’s attention matrix A. We display them vertically so each weight can sit beside its source word and value row. The receiver stays fixed at token 10. Only the source index j changes.
# The final receiver uses row 10 of its head's attention matrix.
# Python index 9 is the tenth token. Each j below names a source.
assert weights.shape == (1, 10)
for j, word in enumerate(case['tokens']):
print(f'source {j + 1}: {word}, alpha = {weights[0, j].item():.3f}')
source 1: The, alpha = 0.027 source 2: fisherman, alpha = 0.693 source 3: sat, alpha = 0.060 source 4: beside, alpha = 0.027 source 5: the, alpha = 0.023 source 6: river, alpha = 0.023 source 7: bank, alpha = 0.027 source 8: and, alpha = 0.027 source 9: watched, alpha = 0.071 source 10: the, alpha = 0.023
Head 2: the value beside each weight¶
Each source supplies a two-number value row from this head’s V matrix. The source words and weights stay in place, so we can see which weight belongs to which value.
V = E W_V has one row per source. These are projected value vectors, not token IDs or attention weights. Our hand-chosen worksheet gives W_K and W_V equal entries, but they remain separate parameters with different roles. The keys helped compute the weights. V supplies the content to mix.
assert V.shape == (10, 2)
for j, word in enumerate(case['tokens']):
print(word, 'weight:', round(weights[0, j].item(), 3),
'value row:', V[j].tolist())
The weight: 0.027 value row: [0.10000000149011612, 2.299999952316284] fisherman weight: 0.693 value row: [2.0999999046325684, 0.0] sat weight: 0.06 value row: [0.6000000238418579, 1.899999976158142] beside weight: 0.027 value row: [0.10000000149011612, 2.0] the weight: 0.023 value row: [0.0, 2.299999952316284] river weight: 0.023 value row: [0.0, 0.10000000149011612] bank weight: 0.027 value row: [0.10000000149011612, 0.800000011920929] and weight: 0.027 value row: [0.10000000149011612, 2.200000047683716] watched weight: 0.071 value row: [0.699999988079071, 1.7000000476837158] the weight: 0.023 value row: [0.0, 2.299999952316284]
Head 2: each source’s contribution¶
Multiply both value coordinates by the same source weight. The new column shows exactly what each source contributes. Displayed numbers are rounded.
The calculation below the rows expands the highlighted source’s scalar–vector multiplication coordinate by coordinate. A small weight scales down the entire value vector. The sign of each coordinate comes from the value. All products use the unrounded weights.
# One scalar multiplies both coordinates of its source's value.
source = 5 if head_index == 0 else 1 # river or fisherman
alpha = weights[0, source]
value = V[source]
print(case['tokens'][source], ':', alpha.item(), '*', value.tolist())
print('Contribution:', (alpha * value).tolist())
# Line up the ten weights with the ten value rows. Broadcasting
# applies each [10, 1] weight to both columns of V [10, 2].
contributions = weights.T * V
assert contributions.shape == (10, 2)
for word, row in zip(case['tokens'], contributions):
print(word, row.tolist())
fisherman : 0.6929605603218079 * [2.0999999046325684, 0.0] Contribution: [1.4552171230316162, 0.0] The [0.002679679309949279, 0.06163262575864792] fisherman [1.4552171230316162, 0.0] sat [0.03625689446926117, 0.11481349170207977] beside [0.002679679309949279, 0.05359358713030815] the [0.0, 0.05238167196512222] river [0.0, 0.0022774641402065754] bank [0.002679679309949279, 0.02143743447959423] and [0.002679679309949279, 0.058952946215867996] watched [0.049770113080739975, 0.12087027728557587] the [0.0, 0.05238167196512222]
Head 2: adding the contributions¶
Add the first coordinates to get the first message coordinate. Add the second coordinates to get the second. All ten sources contribute. The sum uses unrounded products.
This is one row of H = AV: m₁₀ = Σⱼ α₁₀,ⱼ vⱼ. The original weight row has shape [1,10] and V has shape [10,2], giving a [1,2] message. We arranged the weights vertically only to make the source correspondence visible. The displayed products are rounded, but the sums use full precision.
# Sum down the ten sources, keeping the two coordinates separate.
message = contributions.sum(dim=0, keepdim=True)
print('First coordinate:', contributions[:, 0].sum().item())
print('Second coordinate:', contributions[:, 1].sum().item())
assert message.shape == (1, 2)
# The matrix product performs exactly the same multiply-and-sum.
torch.testing.assert_close(message, weights @ V)
torch.testing.assert_close(message[0], torch.tensor(case['heads'][head_index]['messages'][-1]))
print('Message:', message.tolist())
First coordinate: 1.551962971687317 Second coordinate: 0.5383411645889282 Message: [[1.551962971687317, 0.5383411645889282]]
The two complete head calculations¶
Each head has now computed its own scores, weights and message. We keep both messages, then concatenate and project them into the embedding update.
As in Part II, αᵢⱼ is one attention weight and A stores the full attention matrix. The walkthrough follows its final row, A[10, :], using one-based token positions. Q/K choose weights and V supplies the weighted content. All ten source contributions enter each message. The calculations are independent across heads and can run in parallel.
What if we made one head wider?¶
A wider query/key can compare more features. Their products still add into one score per source, followed by one softmax.
For this controlled comparison, concatenate the two existing Q, K and V projections into width-four projections. A single wide head scales its dot products by √4. The two narrower heads each scale by √2 and normalize separately. The example changes no projection entries and does not compare separately trained models.
Two heads keep separate source preferences¶
A wider value vector carries more features with one shared weight row. Two heads can favour river for setting features and fisherman for person features.
Every displayed attention row sums to one over all ten sources; “others” combines the remaining eight. One wide head uses its single row for all four value coordinates. The two heads use different rows for their two-coordinate values, preserving both mixtures before W_O combines them. This is a useful architectural choice, not proof that one wide head cannot learn useful relationships or that more heads always improve accuracy. See Attention Is All You Need, §3.2.2.
# Same projected coordinates, one wide softmax or two narrow softmaxes.
case = worksheet['headsLesson']['cases']['river']
Q_wide, K_wide, V_wide = [
torch.cat([torch.tensor(head[kind]) for head in case['heads']], dim=-1)
for kind in ['Q', 'K', 'V']
]
scores_wide = Q_wide[-1] @ K_wide.T / math.sqrt(4)
weights_wide = scores_wide.softmax(-1) # all ten sources allowed at row 10
message_wide = weights_wide @ V_wide
torch.testing.assert_close(weights_wide, torch.tensor(case['wide']['A'][-1]))
torch.testing.assert_close(message_wide, torch.tensor(case['wide']['messages'][-1]))
for label, row in [('one wide head', weights_wide)] + [
(f'head {h+1}', torch.tensor(head['A'][-1]))
for h, head in enumerate(case['heads'])
]:
print(label, 'river:', round(row[5].item(), 3),
'fisherman:', round(row[1].item(), 3))
one wide head river: 0.178 fisherman: 0.708 head 1 river: 0.718 fisherman: 0.166 head 2 river: 0.023 fisherman: 0.693
Put the messages side by side¶
Concatenation keeps the two messages separate. It does not add them coordinate by coordinate.
Map the messages back to embedding space¶
(W_O) learns how the head messages contribute to the embedding update. Four joined coordinates become four update coordinates.
Equivalently, split W_O into a two-row block per head. Then Δe = m¹ W_O¹ + m² W_O²: each head contributes a four-coordinate update. Concatenation followed by one matrix multiplication computes exactly that sum. We retain the standard concat notation from the original Transformer paper.
Add context to the original embedding¶
Exactly as in Part II: (e_i) is the input embedding row, (\Delta e_i) is the context update, and (e_i^{\prime}=e_i+\Delta e_i).
Return to the next-token prediction¶
The last updated row still feeds the prediction MLP. More heads change how it reads context—not what the target means.
Try the other bank¶
Switch contexts in the optional worked reference to see both reading patterns change. The executable lab below computes the river and cheque examples from the same parameters.
Change to the cheque sentence. These controls recalculate Q, K, V, both attention rows, the messages and the vocabulary prediction from the hand-chosen parameters. This is a worked arithmetic explorer, not a trained language model.
Project every row with the same head matrix¶
Ten input rows × four coordinates. Multiplying by a 4 × 2 projection gives ten query rows × two coordinates.
One head makes one attention grid¶
Keep Part II’s notation: (M) is the causal mask, (A=\operatorname{softmax}(QK^\top/\sqrt{d_k}+M)), and (H=AV) stores the message rows.
Two heads make two attention grids¶
Both heads receive the same E. Each uses its own projections, attention grid and values. Neither reads the other head’s output.
Join coordinates, not token rows¶
(\Delta E=\operatorname{Concat}(H^{(1)},H^{(2)})W_O), then (E^{\prime}=E+\Delta E). The superscript labels the head; each matrix still has ten token rows.
One head is still the Part II calculation¶
No new attention rule is needed. We reuse the same calculation with different learned matrices.
def head(E, W_Q, W_K, W_V):
Q, K, V = E @ W_Q, E @ W_K, E @ W_V
scores = Q @ K.T / math.sqrt(Q.shape[-1])
future = torch.ones(len(E), len(E), dtype=torch.bool).triu(1)
A = scores.masked_fill(future, -torch.inf).softmax(-1)
return A @ V
Call it with two sets of parameters¶
This unbatched example keeps one row per token. The notebook checks these outputs against the printed numbers.
H1 = head(E, W_Q1, W_K1, W_V1)
H2 = head(E, W_Q2, W_K2, W_V2)
joined = torch.cat([H1, H2], dim=-1) # [10, 4]
delta_E = joined @ W_O
E_prime = E + delta_E
The same loss trains both heads¶
Next-token cross-entropy trains all the projections together. We do not label training examples “setting head” or “person head”.
At training time, the observed next-token ID is the target for the vocabulary logits. At inference time, weights remain fixed: tokenize a prefix, run the model, choose a next token and append it. These steps are unchanged from Part II; the full executable loop remains in Notebook 7.
A bias adds a learned offset¶
The offset [0.2, −0.1] is illustrative. With bias=True, training learns it along with the projection weights. We use bias=False so the code matches the worksheet.
A projection bias is shared across token positions (and batches). Position embeddings depend on the position, so these are different operations. A bias can shift the projection even for a zero input. Omitting it keeps this example simpler; it is not a general claim that bias-free attention is better.
Which projections use the bias flag?¶
PyTorch defaults to bias=True. In this lesson, bias=False removes the Q, K, V and output offsets. It leaves the mask, position embeddings and separate prediction layers unchanged.
With total width four, bias=True adds four query offsets, four key offsets, four value offsets and four output offsets: 16 additional trainable scalars. Within each head the Q/K/V bias slice has width two. PyTorch stores the three input offsets together in in_proj_bias and the output offset in out_proj.bias. The distinct add_bias_kv option is not the bias flag discussed here. Official API documentation. Both our scratch implementation and trained attention variants omit attention projection biases; their prediction MLP layers retain biases.
# A learned offset is broadcast to each token row.
Q_no_bias = torch.tensor(case['heads'][0]['Q'])
b_Q = torch.tensor([0.2, -0.1]) # illustrative, not a fitted parameter
Q_with_bias = Q_no_bias + b_Q
print('Receiver 10:', Q_no_bias[-1].tolist(), '->', Q_with_bias[-1].tolist())
torch.testing.assert_close(Q_with_bias[-1], torch.tensor([2.5, 2.2]))
without_bias = nn.MultiheadAttention(4, 2, bias=False)
with_bias = nn.MultiheadAttention(4, 2, bias=True)
count = lambda layer: sum(p.numel() for p in layer.parameters())
print('Projection parameters:', count(without_bias), 'vs', count(with_bias))
print('Packed Q/K/V bias:', tuple(with_bias.in_proj_bias.shape))
print('Output bias:', tuple(with_bias.out_proj.bias.shape))
assert count(with_bias) - count(without_bias) == 16
assert without_bias.in_proj_bias is None
assert without_bias.out_proj.bias is None
Receiver 10: [2.299999952316284, 2.299999952316284] -> [2.5, 2.200000047683716] Projection parameters: 64 vs 80 Packed Q/K/V bias: (12,) Output bias: (4,)
PyTorch packages the head calculation¶
Here E has shape [10, 4]: one unbatched sequence. PyTorch handles the projections and mixing. A contains two 10 × 10 weight grids.
mha = nn.MultiheadAttention(embed_dim=4, num_heads=2, bias=False)
future = torch.ones(10, 10, dtype=torch.bool).triu(1)
delta_E, A = mha(E, E, E, attn_mask=future,
average_attn_weights=False)
E_prime = E + delta_E
The library layer does not add position embeddings, the residual or the vocabulary classifier. attn_mask=future blocks future sources. average_attn_weights=False returns each head’s weight grid separately; it changes the returned diagnostic weights, not how the head messages combine. The notebook copies identical weights from our scratch implementation to PyTorch and checks both ΔE and the per-head attention weights. See the PyTorch documentation.
Back to the complete multi-head model¶
We have computed both heads, joined their messages, applied the output projection and added the residual. The final updated row still goes to the familiar prediction MLP.
This repeats the same two-head diagram after the detailed arithmetic, rather than introducing a new architecture. The worksheet uses ten rows of width four. Revisit Head 1’s calculations, Head 2’s calculations, or the output projection. Notebook 7 executes the complete forward pass, including the prediction MLP and vocabulary probabilities. This teaching model has one attention update and a prediction MLP; it is not a diagram of every component in a deep production Transformer.
The live MLP reads a flattened window¶
Now look at the genuinely trained browser models. The MLP reads all 64 token rows in fixed order. Its hidden layer receives 4,096 inputs.
This diagram follows FixedWindowMLP in wordlm.py and the saved comparison protocol. Each token embedding has 64 coordinates; concatenating 64 rows gives 4,096. A 4,096→256 hidden layer with ReLU feeds a 256→4,000 vocabulary layer. The fixed concatenation order already distinguishes slots. This baseline has no separate learned position table. Its larger first MLP matrix helps explain its parameter count.
One head adds a learned context update¶
Add learned position rows before Q, K and V. One head mixes the sources, the output projection transforms its message, and the residual preserves E. Only the last updated row feeds the prediction MLP.
The trained single-head model uses d_model = d_k = d_v = 64. The full attention matrix has shape 64×64, with future and padding masks. The optimized browser inference path computes only the final query row, which gives the same next-token logits as taking the last row of this full diagram. The 64→256 prediction MLP replaces the baseline’s much wider 4,096→256 input layer.
Four heads keep four source-weight patterns¶
The input and output widths stay 64. Each head now uses 16 coordinates and its own softmax. Concatenation restores width 64 before the output projection, residual and the same prediction MLP.
This follows MultiHeadAttentionLM in multihead.py. Each head has Q, K and V of shape 64×16 and its own 64×64 attention matrix. Its H = AV has shape 64×16. Concatenating four H matrices gives 64×64, followed by W_O [64×64]. Both attention variants add learned absolute position rows to token rows. The comparison does not isolate the effect of position encoding, because it does not train a position-free attention ablation. Head count, total width, parameter count, measured losses and device timings must be distinguished.
Keep the total width fixed¶
The trained comparison uses width 64. Four smaller heads give four attention patterns with the same attention parameter count as one wide head.
Implementations often concatenate the per-head W_Q matrices into one large W_Q (and similarly for K and V). They project E first and reshape the projected coordinates into heads. They do not assign disjoint slices of raw E to different heads. Our two-head arithmetic toy uses width 4; this experiment uses width 64.
Compare held-out next-token prediction¶
For each run, perplexity = exp(cross-entropy using natural logs). Lower means higher probability for observed tokens. The table averages three seeds on the same test stories.
Same 6,000-document TinyStories subset, tokenizer, 64-token context, 6,000-update budget and validation selection. Seeds 11, 29 and 47. More heads helped this experiment; this is not a guarantee for every dataset, head count or generated continuation. Full experiment and variation across seeds.
Compare size and training time too¶
Mean training times on Apple M2 Max/MPS, including validation. The two attention models have equal parameter counts; four heads took slightly longer.
Give all three models the same prompt¶
Open the live models ↗ · Work through the code · Inspect the trained experiment
The demo runs actual checkpoints with WebGPU or WebAssembly. It includes training, held-out and outside-domain prompts and measures generation locally. Normalization, full Transformer blocks and complexity remain in the optional Part 2B reference. Formula source: Attention Is All You Need, §3.2.2. Visual teaching references: 3Blue1Brown and Jay Alammar.
Image patches can supply the embedding rows¶
Replace word tokens with image patches. Flatten each patch and apply the same learned projection. Attention still receives a sequence of embedding rows.
This is a tiny shape illustration, not a trained vision model: an 8×8 grayscale image split into four 4×4 patches. Each patch has 16 pixel values; a shared 16×4 projection produces four-dimensional rows. For RGB images the flattened width is P²×3. The general sequence length is (H/P)×(W/P). The projection is learned from the image task.
One updated row can classify the whole image¶
The [CLS] row gathers image information. Its final representation feeds a classifier trained with the image label. Next: Vision I, from pixels to an image class ↗
For this two-class shape example, prepend one learned CLS token to the four patch embeddings, then add position embeddings to the five rows. A Transformer encoder updates all five rows. The final CLS row feeds a two-logit classifier; softmax produces class probabilities and cross-entropy uses the observed image label. Image classification does not need a causal future-token mask. The full encoder also includes per-row MLPs, residuals and normalization, which are grouped here rather than silently omitted. This is a roadmap, not a benchmark or a completed classifier. See An Image is Worth 16×16 Words. Cross-attention remains a separate optional continuation.
The executable lab¶
The lecture has now shown the whole idea. This optional lab slows down the code:
we create two examples, project the embeddings, expose the head axis, calculate
the messages and check the results. Here D means the same model width as
d_model in the lecture. The head axis has size 2; it is not the message matrix H.
The expanded figures below accompany the code rather than add new lecture slides.
Which earlier characters might help?¶
Part I predicted the next character in a name such as aabid. The question stays the same; we are changing how the model reads the known context.
Back to the river bank¶
One query can benefit from several clues at once. What setting are we in? Who is there?
Keep the two examples separate¶
The worksheet explains the arithmetic. Held-out results later tell us whether the trained models improved.
From one weight row to two¶
What changes when the same input passes through two sets of projections?
One head produces one message¶
One head can already read several sources. Its value coordinates share the same attention weights.
Two heads keep two messages¶
Each head has its own Q, K and V projections and its own softmax. Both receive the full E.
Project first; split the projected coordinates¶
We split Q, K and V after projection. We do not give head 1 the first half of the raw embedding and head 2 the second half.
Separate heads can use different input features¶
These are two 4×2 matrices placed side by side. Every head can learn from every input coordinate.
Start with the same position-aware rows¶
E already includes token embedding + position embedding. Token IDs are not embedding coordinates.
The first query asks about the setting¶
Our chosen projection copies the glue coordinate into both query coordinates: [2.3, 2.3].
The second query uses a different projection¶
The same input row now produces [2.3, 0.0]. The query changes because the projection matrix changes.
The source keys are different too¶
Head 1 exposes water/finance features. Head 2 exposes person/glue features. These labels belong to this designed worksheet.
Head 1: scores become a weight row¶
Each weight is exp(score) divided by the sum of exp(scores) in this row. All ten sources are allowed for the final query.
Head 2 has its own softmax¶
Normalize over sources within head 2. Do not normalize across heads or average their scores.
Head 1 sends water and finance information¶
The river contribution is its weight × its value row. Add the contributions from every source to get the two-coordinate message.
Head 2 sends a different kind of message¶
The second head mixes its own values with its own weights. Matching features choose where to read; values carry the information.
Bring the two messages back together¶
How do two small message rows become one update to the original row?
Concatenate the messages, not the tokens¶
Two 2-coordinate rows become one 4-coordinate row. Concatenation preserves both feature groups; it does not average them.
W_O can mix information across heads¶
Water update = 2.617 + 0.25 × 1.552 = 3.005. W_O maps the joined message back to model width.
Add the update to the original row¶
The residual is the same addition as in Part II. There is still one updated row per input token.
The prediction MLP stays in place¶
Choose a token only after vocabulary softmax. Attention weights choose source positions; vocabulary probabilities choose possible next tokens.
Change the context; inspect both heads¶
The matching slide lets you switch contexts and inspect either head. Here, compute both sentences and print their final-query messages. These are hand-chosen worksheet parameters, not trained results.
for name in ['river', 'cheque']:
ids = torch.tensor([[word_to_id[w.lower()] for w in worksheet['sentences'][name]]])
demo = load_worksheet_weights(TinyMultiHeadLM(len(word_to_id)), worksheet)
E_demo = demo.token_embedding(ids) + demo.position_embedding(torch.arange(10))
with torch.no_grad():
_, weights = demo.attention(E_demo)
values = demo.attention.split_heads(demo.attention.W_V(E_demo))
messages = weights @ values
print(name, 'final messages by head:', messages[0, :, -1].tolist())
river final messages by head: [[2.6168477535247803, -0.017858127132058144], [1.5519635677337646, 0.5383409261703491]] cheque final messages by head: [[0.05300545319914818, 2.471384286880493], [2.3669915199279785, 0.4609062373638153]]
Create the tables, projections and prediction MLP¶
The download defines every layer. Load the hand-chosen worksheet weights to reproduce the printed numbers; ordinary training starts from random parameters.
The complete implementation¶
This is the same source file imported above. Read the small snippets that follow alongside the diagram, then return here to see how they fit together. load_worksheet_weights copies the printed parameters so our outputs match the figures. It is not part of an ordinary training loop.
"""Part III: explicit multi-head self-attention and a small next-token model."""
import math
import torch
from torch import nn
from torch.nn import functional as F
class ScratchMultiHead(nn.Module):
def __init__(self, width=4, heads=2):
super().__init__()
if heads < 1 or width % heads:
raise ValueError('width must be divisible by a positive head count')
self.width, self.heads = width, heads
self.head_width = width // heads
self.W_Q = nn.Linear(width, width, bias=False)
self.W_K = nn.Linear(width, width, bias=False)
self.W_V = nn.Linear(width, width, bias=False)
self.W_O = nn.Linear(width, width, bias=False)
def split_heads(self, rows):
B, T, _ = rows.shape
return rows.reshape(B, T, self.heads, self.head_width).transpose(1, 2)
def forward(self, E, padding_mask=None):
B, T, _ = E.shape
Q = self.split_heads(self.W_Q(E))
K = self.split_heads(self.W_K(E))
V = self.split_heads(self.W_V(E))
scores = Q @ K.transpose(-2, -1) / math.sqrt(self.head_width)
blocked = torch.ones(T, T, dtype=torch.bool, device=E.device).triu(1)
if padding_mask is not None:
# Ignore PAD sources for real queries. PAD-query outputs are unused;
# leave their causal prefix available to avoid an all-masked softmax.
blocked = blocked | (padding_mask[:, None, :] & ~padding_mask[:, :, None])
blocked = blocked[:, None]
weights = scores.masked_fill(blocked, float('-inf')).softmax(dim=-1)
messages = weights @ V
joined = messages.transpose(1, 2).contiguous().reshape(B, T, self.width)
delta = self.W_O(joined)
return delta, weights
class TinyMultiHeadLM(nn.Module):
def __init__(self, vocab_size, context=10, width=4, heads=2, hidden=8):
super().__init__()
self.context = context
self.token_embedding = nn.Embedding(vocab_size, width)
self.position_embedding = nn.Embedding(context, width)
self.attention = ScratchMultiHead(width, heads)
self.hidden = nn.Linear(width, hidden)
self.readout = nn.Linear(hidden, vocab_size)
def forward(self, ids):
T = ids.shape[1]
if T > self.context:
raise ValueError('Crop the prompt to the configured context window')
E = self.token_embedding(ids) + self.position_embedding(torch.arange(T, device=ids.device))
delta, _ = self.attention(E)
updated = E + delta
return self.readout(F.relu(self.hidden(updated[:, -1])))
def copy_to_pytorch(scratch):
"""Use exactly the same parameters, not a newly randomized comparison."""
layer = nn.MultiheadAttention(scratch.width, scratch.heads, bias=False,
dropout=0.0, batch_first=True)
layer = layer.to(device=scratch.W_Q.weight.device, dtype=scratch.W_Q.weight.dtype)
with torch.no_grad():
layer.in_proj_weight.copy_(torch.cat([scratch.W_Q.weight,
scratch.W_K.weight,
scratch.W_V.weight], dim=0))
layer.out_proj.weight.copy_(scratch.W_O.weight)
return layer
def load_worksheet_weights(model, worksheet):
"""Load the printed, hand-chosen example; this is not a training algorithm."""
heads = worksheet['headsLesson']['projections']
with torch.no_grad():
model.token_embedding.weight.copy_(torch.tensor([worksheet['tok_emb'][w] for w in worksheet['vocab']]))
model.position_embedding.weight.copy_(torch.tensor(worksheet['pos_emb'][:model.context]))
for letter in ['Q', 'K', 'V']:
packed = [a + b for a, b in zip(heads[0][letter], heads[1][letter])]
getattr(model.attention, 'W_' + letter).weight.copy_(torch.tensor(packed).T)
model.attention.W_O.weight.copy_(torch.tensor(worksheet['headsLesson']['W_O']).T)
model.hidden.weight.copy_(torch.tensor(worksheet['W_hidden']).T)
model.hidden.bias.copy_(torch.tensor(worksheet['b_hidden']))
model.readout.weight.copy_(torch.tensor(worksheet['W_vocab']).T)
model.readout.bias.copy_(torch.tensor(worksheet['b_vocab']))
return model
model = TinyMultiHeadLM(vocab_size=20, context=10,
width=4, heads=2, hidden=8)
load_worksheet_weights(model, worksheet)
TinyMultiHeadLM(
(token_embedding): Embedding(20, 4)
(position_embedding): Embedding(10, 4)
(attention): ScratchMultiHead(
(W_Q): Linear(in_features=4, out_features=4, bias=False)
(W_K): Linear(in_features=4, out_features=4, bias=False)
(W_V): Linear(in_features=4, out_features=4, bias=False)
(W_O): Linear(in_features=4, out_features=4, bias=False)
)
(hidden): Linear(in_features=4, out_features=8, bias=True)
(readout): Linear(in_features=8, out_features=20, bias=True)
)
Two input windows form one batch¶
These are integer IDs, not embeddings. Each example has one observed target.
X = torch.tensor([river_ids, cheque_ids])
y = torch.tensor([word_to_id["water"], word_to_id["teller"]])
Find the embedding lookup on the map¶
We have token IDs. Next, look up their learned rows and add the position rows. Both heads will read the result.
Look up token rows and add position rows¶
The embedding tables are learned parameters. E contains floating-point vectors with shape [2,10,4].
positions = torch.arange(X.shape[1])
E = model.token_embedding(X) + model.position_embedding(positions)
Project the full input three times¶
Different learned matrices give matching queries, matching keys and transmitted values.
Q = model.attention.W_Q(E)
K = model.attention.W_K(E)
V = model.attention.W_V(E)
Make the head axis explicit¶
reshape groups projected coordinates; transpose puts heads before tokens. Apply the same operation to K and V.
Q = Q.reshape(2, 10, 2, 2).transpose(1, 2)
K = K.reshape(2, 10, 2, 2).transpose(1, 2)
V = V.reshape(2, 10, 2, 2).transpose(1, 2)
Compute every query–key pair within each head¶
The first two axes keep examples and heads separate. Matrix multiplication contracts only the matching-coordinate axis.
scores = Q @ K.transpose(-2, -1)
scores = scores / math.sqrt(2)
Block future sources in every head¶
True marks a blocked entry. The same causal mask broadcasts across both examples and both heads.
future = torch.ones(10, 10, dtype=torch.bool).triu(1)
scores = scores.masked_fill(future, float("-inf"))
Each row becomes a distribution over sources¶
The final axis indexes source tokens. Every allowed row sums to one; masked entries have weight zero.
A = scores.softmax(dim=-1)
assert A.shape == (2, 2, 10, 10)
expected = torch.tensor([[h['A'] for h in worksheet['headsLesson']['cases'][name]['heads']]
for name in ['river', 'cheque']])
torch.testing.assert_close(A, expected)
torch.testing.assert_close(A.sum(-1), torch.ones(2, 2, 10))
assert not A.triu(1).any()
print('All 400 head weights match the worksheet; no future source receives weight.')
All 400 head weights match the worksheet; no future source receives weight.
Multiply the weights by the values¶
The source-token axis is summed out. Each head keeps its own two-coordinate message.
messages = A @ V
Put each token’s head messages side by side¶
Transpose before reshaping. A direct reshape of the original layout would mix token rows with head rows.
joined = messages.transpose(1, 2).contiguous()
joined = joined.reshape(2, 10, 4)
Return to the shared prediction path¶
The two head messages are joined. W_O mixes them, and the residual adds that update to the original E.
Project, add, and keep the final row¶
W_O combines the heads. The residual keeps E. Only the last updated row feeds this next-token loss.
delta = model.attention.W_O(joined)
updated = E + delta
logits = model.readout(F.relu(model.hidden(updated[:, -1])))
expected = torch.tensor([worksheet['headsLesson']['cases'][name]['logits'][-1]
for name in ['river', 'cheque']])
torch.testing.assert_close(logits, expected)
torch.testing.assert_close(model(X), logits)
print('All 40 vocabulary logits match the independent worksheet.')
All 40 vocabulary logits match the independent worksheet.
The rest of the learning loop is familiar¶
More heads change the context update, not the definition of the next-token target.
Score the two observed targets¶
Cross-entropy reads logits and observed token IDs. Do not sample a generated token to make the training target.
logits = model(X)
loss = F.cross_entropy(logits, y)
One loss trains all the projections¶
Heads are not assigned jobs or separate labels. They receive gradients from the same prediction loss.
optimizer = torch.optim.AdamW(model.parameters(), lr=0.001)
optimizer.zero_grad()
loss.backward()
optimizer.step()
assert all(p.grad is not None and p.grad.isfinite().all() for p in model.parameters())
print('All learned tables, head projections and prediction layers received finite gradients.')
All learned tables, head projections and prediction layers received finite gradients.
Start generation from known token IDs¶
We already encoded the river sentence. A real application tokenizes and looks up a new prompt with the same vocabulary.
history = torch.tensor([river_ids])
model.eval()
TinyMultiHeadLM(
(token_embedding): Embedding(20, 4)
(position_embedding): Embedding(10, 4)
(attention): ScratchMultiHead(
(W_Q): Linear(in_features=4, out_features=4, bias=False)
(W_K): Linear(in_features=4, out_features=4, bias=False)
(W_V): Linear(in_features=4, out_features=4, bias=False)
(W_O): Linear(in_features=4, out_features=4, bias=False)
)
(hidden): Linear(in_features=4, out_features=8, bias=True)
(readout): Linear(in_features=8, out_features=20, bias=True)
)
Generate one token, then repeat¶
At inference, weights stay fixed. The current prefix becomes longer; we keep only the configured context window.
with torch.no_grad():
logits = model(history[:, -model.context:])
next_id = logits.argmax(dim=-1, keepdim=True)
history = torch.cat([history, next_id], dim=1)
Trace one request through the full model¶
Start at the known IDs. Name the shape at each arrow. Where does the observed target enter during training? It enters only at the loss.
Replace the head calculation, not the whole model¶
The PyTorch layer replaces both heads, their concatenation and W_O. We still add the residual and use our prediction MLP.
Create the multi-head attention layer¶
batch_first=True means [B,T,D]. bias=False and dropout=0 match our scratch implementation.
mha = nn.MultiheadAttention(
embed_dim=4, num_heads=2, bias=False,
dropout=0.0, batch_first=True)
Pass E as query, key and value input¶
PyTorch performs the learned projections internally. average_attn_weights=False preserves the head axis in the returned weight tensor.
delta, A = mha(E, E, E, attn_mask=future,
average_attn_weights=False)
updated = E + delta
Compare the same weights, not two random layers¶
The notebook copies the scratch Q/K/V and W_O parameters into PyTorch, then checks both outputs numerically.
mha = copy_to_pytorch(model.attention)
delta_scratch, A_scratch = model.attention(E)
delta_api, A_api = mha(E, E, E, attn_mask=future,
average_attn_weights=False)
torch.testing.assert_close(delta_api, delta_scratch)
torch.testing.assert_close(A_api, A_scratch)
print('Scratch and PyTorch updates agree:', tuple(delta_api.shape))
print('Individual head weights agree:', tuple(A_api.shape))
Scratch and PyTorch updates agree: (2, 10, 4) Individual head weights agree: (2, 2, 10, 10)
What the attention block does not include¶
nn.MultiheadAttention does not add position embeddings, the residual, a vocabulary head or a training loop.
Do more heads help this experiment?¶
Move from the small worksheet to the actual TinyStories checkpoints.
More heads need not mean more parameters¶
Q, K, V and W_O remain 64×64. The number of attention patterns grows, while each pattern uses fewer matching coordinates.
Four heads improved held-out prediction here¶
Three-seed means. Same stories, tokenizer, context and update budget; validation-selected checkpoints. Perplexity = exp(cross-entropy). Lower is better. This does not guarantee better text for every prompt.
Compare the cost as well as the score¶
Mean training time on Apple M2 Max/MPS, including validation checks. Both attention models have equal parameter counts; the MLP is larger.
Try the same prompt with all three models¶
Open the live browser demo ↗ Compare continuations, vocabulary coverage and generation time. This runs real checkpoints with WebGPU or WASM.
Next: read a different sequence¶
So far, Q, K and V came from the same sequence. Part IV changes that: cross-attention reads another sequence.
Continue with the trained experiment¶
Notebook 6 inspects the four-head TinyStories checkpoint, per-head weights and all three-seed results. The browser demo loads the actual exported checkpoints and measures generation on your device.
Visual teaching references: 3Blue1Brown and Jay Alammar. Formula and API references: Attention Is All You Need, §3.2.2 and PyTorch MultiheadAttention. The diagrams and worked numbers here are original adaptations of our Part II example.