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.

Visual story · Executable lab

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.

In [1]:
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¶

Optional worked reference

The river-bank prefix from Part IIThe1fisherman2sat3beside4the5river6bank7and8watched9the10___The known prefix ends here.The updated final “the” will predict the next word.

Which parts of this prefix would help you choose a continuation?

One token can need several kinds of detail¶

Optional worked reference

Illustrative coat head readingsPossible head roles — an illustration, not a trained model’s attention1Maya2wore3a4hooded5red6wool7coatreceiverHead / query: what to look forKey: a matching sourceValue: what it can send1 · Which colour?redcolour: redStill to retrieve: material and 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¶

Optional worked reference

Illustrative coat head readingsPossible head roles — an illustration, not a trained model’s attention1Maya2wore3a4hooded5red6wool7coatreceiverHead / query: what to look forKey: a matching sourceValue: what it can send1 · Which colour?redcolour: red2 · Which material?woolmaterial: wool3 · Which detail?hoodeddetail: hoodKeep colour, material and detail as separate messages.

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¶

Optional worked reference

Illustrative grammar head readingsPossible head roles — an illustration, not a trained model’s attention1The2dogs3near4the5gate6usuallyreceiverHead / query: what to look forKey: a matching sourceValue: what it can send1 · Who is the subject?dogsplural subject2 · Where are they?gatelocation cluePossible next words: “bark at visitors”“dogs” controls agreement; the nearer noun “gate” is singular.

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¶

Optional worked reference

Illustrative event head readingsPossible head roles — an illustration, not a trained model’s attention1Ravi2dropped3the4glass5.6ItreceiverHead / query: what to look forKey: a matching sourceValue: what it can send1 · What does “It” refer to?glassobject: glass2 · What happened to it?droppedevent: droppingPossible continuation: “broke”The object alone does not tell us what happened.

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¶

Optional worked reference

Illustrative event head readings after changing the actionPossible head roles — an illustration, not a trained model’s attention1Ravi2washed3the4glass5.6ItreceiverHead / query: what to look forKey: a matching sourceValue: what it can send1 · What does “It” refer to?glassobject: glass2 · What happened to it?washedevent: washingPossible continuation: “gleamed”Same object. A different event gives a different clue.

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¶

Optional worked reference

The final the uses one query to retrieve one weighted messageThe1fisherman2sat3beside4the5river6bank7and8watched9the10___Receiver: the final “the”. Its updated row predicts the blank.Query at this receiverWeights from matching keysMix source valuesHead 1 · q(1)Setting clues?Favour riverAlso read the other tokensMessage 1setting informationIllustrative head roles. The queries and messages are numerical vectors.

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¶

Optional worked reference

The same final the receives two independently weighted messagesThe1fisherman2sat3beside4the5river6bank7and8watched9the10___Receiver: the final “the”. Its updated row predicts the blank.Query at this receiverWeights from matching keysMix source valuesHead 1 · q(1)Setting clues?Favour riverAlso read the other tokensMessage 1setting informationHead 2 · q(2)Person clues?Favour fishermanAlso read the other tokensMessage 2person informationIllustrative head roles. The queries and messages are numerical vectors.

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.

In [2]:
# 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¶

Optional worked reference

Head 1 reads the known prefixHead 1: setting cluesThe1fisherman2sat3beside4the5river6bank7and8watched9the10___0.71810 · theThicker arrow = more attention weightreceiver

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¶

Optional worked reference

Head 2 reads the known prefixHead 2: person cluesThe1fisherman2sat3beside4the5river6bank7and8watched9the10___0.69310 · theThicker arrow = more attention weightreceiver

The sentence and receiver stay fixed. A different head gives fisherman a large weight.

Keep both readings¶

Optional worked reference

Two different reading patternsHead 1: setting cluesThe1fisherman2sat3beside4the5river6bank7and8watched9the100.71810 · theThicker arrow = more attention weightreceiverHead 2: person cluesThe1fisherman2sat3beside4the5river6bank7and8watched9the100.69310 · theThicker arrow = more attention weightreceiver

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?¶

Optional worked reference

Queries, keys and values in each headQueries, keys and values in each headEach head learns its own WQ, WK and WV.

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¶

Optional worked reference

Recall Part II: the query asks, the key matches, the value supplies contentMaya cycled home in the rain. Cold and tired, Maya reachedfor a hooded red wool coat. She …q: what is needed?She × WQWhich earlier person?k: what can match?Maya × WKPerson candidatev: what is sent?Maya × WVCold, tired; cycled in rain“Maya” and “She” here mean their current embedding rows.

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¶

Optional worked reference

The same next-token path, with every matrix shape in two parallel headsOne sentence: 10 tokens, width 4. Two heads, width 2 each.Input E[10×4]Head 1Q [10×2]K [10×2]V [10×2]QKᵀ: [10×10]scale, causal mask, softmaxA: [10×10]H(1) = A(1)V(1)[10×10] × [10×2]= [10×2]Head 2Q [10×2]K [10×2]V [10×2]QKᵀ: [10×10]scale, causal mask, softmaxA: [10×10]H(2) = A(2)V(2)[10×10] × [10×2]= [10×2]Concatenatetwo [10×2] → [10×4]Project with WO[10×4] × [4×4]keep original E [10×4]ΔE [10×4]+E′ = E + ΔE[10×4], same as ELast row [1×4]next-token prediction MLP

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¶

Optional worked reference

The familiar query, key and value roles, now repeated in two headsReceiver: the final “the” in the river-bank prefixHeadQuery asks aboutA useful sourceValue carries1the settingriversetting clues2the personfishermanperson cluesEach head computes q, k and v for every token.We follow one receiver and two useful sources.

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¶

Optional worked reference

Head 1: multiply the same four-coordinate embedding by its own query projectionReceiver: final “the”q₁₀(1) = e₁₀ WQ(1)e₁₀: four input coordinates0.0water0.0finance0.0person2.3glue×WQ(1) [4 × 2]water?finance?00000011q₁₀(1) [1 × 2]2.32.3First query coordinate: 0×0 + 0×0 + 0×0 + 2.3×1 = 2.3Second query coordinate: 0×0 + 0×0 + 0×0 + 2.3×1 = 2.3

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¶

Optional worked reference

Head 2: multiply the same four-coordinate embedding by its own query projectionReceiver: final “the”q₁₀(2) = e₁₀ WQ(2)e₁₀: four input coordinates0.0water0.0finance0.0person2.3glue×WQ(2) [4 × 2]person?glue?00000010q₁₀(2) [1 × 2]2.30.0First query coordinate: 0×0 + 0×0 + 0×0 + 2.3×1 = 2.3Second query coordinate: 0×0 + 0×0 + 0×0 + 2.3×0 = 0.0

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¶

Optional worked reference

Each source embedding supplies a key for matching and a value for the weighted messageSource input eⱼ: [water, finance, person, glue]Head 1: river[3.1, −0.1, 0.0, 0.1]WK(1) selects water, financeWV(1) selects water, financek = [3.1, −0.1]v = [3.1, −0.1]Head 2: fisherman[2.0, 0.1, 2.1, 0.0]WK(2) selects person, glueWV(2) selects person, gluek = [2.1, 0.0]v = [2.1, 0.0]kⱼ = eⱼ WK sets the match. vⱼ = eⱼ WV supplies the message.

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.

In [3]:
# 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¶

Optional worked reference

Head 1: all ten query rows and key rows, with the receiver and example source highlightedToken and positionQ(1) = E WQ(1)[10 × 4] × [4 × 2] = [10 × 2]water?finance?2.32.30.00.01.91.92.02.02.32.30.10.10.80.82.22.21.71.72.32.3K(1) = E WK(1)[10 × 4] × [4 × 2] = [10 × 2]waterfinance0.10.02.00.10.30.00.6−0.10.10.13.1−0.10.70.7−0.10.10.70.00.00.0 1 The 2 fisherman 3 sat 4 beside 5 the 6 river 7 bank 8 and 9 watched10 theFollow query row 10 (“the”) and key row 6 (“river”).

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.

In [4]:
# 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¶

Optional worked reference

Head 1: one query times the transposed key matrix yields ten dot productsq₁₀(1) =2.32.3Receiver 10: final “the”q [1 × 2] × Kᵀ [2 × 10] = raw scores r [1 × 10]Kᵀwaterfinanceraw r1The0.10.00.232fisherman2.00.14.833sat0.30.00.694beside0.6−0.11.155the0.10.10.466river3.1−0.16.907bank0.70.73.228and−0.10.10.009watched0.70.01.6110the0.00.00.00Column 6: query · river key2.3 × 3.1 + 2.3 × (−0.1) = 7.13 + (−0.23) = 6.90

Transposing (K) puts each source key in a column. Multiplying the query row by (K^\top) gives one raw dot product per source.

In [5]:
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¶

Optional worked reference

Head 1: scale all ten dot products and divide each exponential by the sum of all ten exponentialsTen matching scores, normalized within this headSourceRaw dot rs = r / √2exp(s)Weight α 1 The0.230.1631.1770.006 2 fisherman4.833.41530.4270.166 3 sat0.690.4881.6290.009 4 beside1.150.8132.2550.012 5 the0.460.3251.3840.008 6 river6.904.879131.5040.718 7 bank3.222.2779.7460.053 8 and0.000.0001.0000.005 9 watched1.611.1383.1220.01710 the0.000.0001.0000.005Shared total: Σ exp(s) = 183.244river weight: 131.504 / 183.244 ≈ 0.718

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.

In [6]:
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¶

Optional worked reference

Head 1: attention weights for each sourceReceiver: the final “the”, at position 10Source j 1 The 2 fisherman 3 sat 4 beside 5 the 6 river 7 bank 8 and 9 watched10 theWeight α10,j0.0060.1660.0090.0120.0080.7180.0530.0050.0170.005α10,6 ≈ 0.718Receiver 10 reads source 6: river71.8% of this head’s total weightThese are the softmax weights we just calculated.

(\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.

In [7]:
# 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¶

Optional worked reference

Head 1: pair each source weight with its value rowReceiver: the final “the”, at position 10Source j 1 The 2 fisherman 3 sat 4 beside 5 the 6 river 7 bank 8 and 9 watched10 theWeight α10,j0.0060.1660.0090.0120.0080.7180.0530.0050.0170.005Value row vⱼ[0.1, 0.0][2.0, 0.1][0.3, 0.0][0.6, −0.1][0.1, 0.1][3.1, −0.1][0.7, 0.7][−0.1, 0.1][0.7, 0.0][0.0, 0.0]V = E WV10 source rows × 2river: its weight pairs with its own two-number value row.

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.

In [8]:
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¶

Optional worked reference

Head 1: multiply both value coordinates by the same source weightReceiver: the final “the”, at position 10Source j 1 The 2 fisherman 3 sat 4 beside 5 the 6 river 7 bank 8 and 9 watched10 theWeight α10,j0.0060.1660.0090.0120.0080.7180.0530.0050.0170.005Value row vⱼ[0.1, 0.0][2.0, 0.1][0.3, 0.0][0.6, −0.1][0.1, 0.1][3.1, −0.1][0.7, 0.7][−0.1, 0.1][0.7, 0.0][0.0, 0.0]Contribution α10,j vⱼ[0.001, 0.000][0.332, 0.017][0.003, 0.000][0.007, −0.001][0.001, 0.001][2.225, −0.072][0.037, 0.037][−0.001, 0.001][0.012, 0.000][0.000, 0.000]×=river: [α10,6 × (3.1), α10,6 × (−0.1)] ≈ [2.225, −0.072]

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.

In [9]:
# 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¶

Optional worked reference

Head 1: sum all ten weighted value rows to obtain the messageReceiver: the final “the”, at position 10Source j 1 The 2 fisherman 3 sat 4 beside 5 the 6 river 7 bank 8 and 9 watched10 theWeight α10,j0.0060.1660.0090.0120.0080.7180.0530.0050.0170.005Value row vⱼ[0.1, 0.0][2.0, 0.1][0.3, 0.0][0.6, −0.1][0.1, 0.1][3.1, −0.1][0.7, 0.7][−0.1, 0.1][0.7, 0.0][0.0, 0.0]Contribution α10,j vⱼ[0.001, 0.000][0.332, 0.017][0.003, 0.000][0.007, −0.001][0.001, 0.001][2.225, −0.072][0.037, 0.037][−0.001, 0.001][0.012, 0.000][0.000, 0.000]×=Sum each contribution coordinate:m₁₀(1) =2.617−0.018weight row [1 × 10] × V [10 × 2] = message [1 × 2]

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.

In [10]:
# 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¶

Optional worked reference

Head 2: all ten query rows and key rows, with the receiver and example source highlightedToken and positionQ(2) = E WQ(2)[10 × 4] × [4 × 2] = [10 × 2]person?glue?2.30.00.00.01.90.02.00.02.30.00.10.00.80.02.20.01.70.02.30.0K(2) = E WK(2)[10 × 4] × [4 × 2] = [10 × 2]personglue0.12.32.10.00.61.90.12.00.02.30.00.10.10.80.12.20.71.70.02.3 1 The 2 fisherman 3 sat 4 beside 5 the 6 river 7 bank 8 and 9 watched10 theFollow query row 10 (“the”) and key row 2 (“fisherman”).

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.

In [11]:
# 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¶

Optional worked reference

Head 2: one query times the transposed key matrix yields ten dot productsq₁₀(2) =2.30.0Receiver 10: final “the”q [1 × 2] × Kᵀ [2 × 10] = raw scores r [1 × 10]Kᵀpersonglueraw r1The0.12.30.232fisherman2.10.04.833sat0.61.91.384beside0.12.00.235the0.02.30.006river0.00.10.007bank0.10.80.238and0.12.20.239watched0.71.71.6110the0.02.30.00Column 2: query · fisherman key2.3 × 2.1 + 0.0 × 0.0 = 4.83 + (0.00) = 4.83

Transposing (K) puts each source key in a column. Multiplying the query row by (K^\top) gives one raw dot product per source.

In [12]:
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¶

Optional worked reference

Head 2: scale all ten dot products and divide each exponential by the sum of all ten exponentialsTen matching scores, normalized within this headSourceRaw dot rs = r / √2exp(s)Weight α 1 The0.230.1631.1770.027 2 fisherman4.833.41530.4270.693 3 sat1.380.9762.6530.060 4 beside0.230.1631.1770.027 5 the0.000.0001.0000.023 6 river0.000.0001.0000.023 7 bank0.230.1631.1770.027 8 and0.230.1631.1770.027 9 watched1.611.1383.1220.07110 the0.000.0001.0000.023Shared total: Σ exp(s) = 43.908fisherman weight: 30.427 / 43.908 ≈ 0.693

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.

In [13]:
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¶

Optional worked reference

Head 2: attention weights for each sourceReceiver: the final “the”, at position 10Source j 1 The 2 fisherman 3 sat 4 beside 5 the 6 river 7 bank 8 and 9 watched10 theWeight α10,j0.0270.6930.0600.0270.0230.0230.0270.0270.0710.023α10,2 ≈ 0.693Receiver 10 reads source 2: fisherman69.3% of this head’s total weightThese are the softmax weights we just calculated.

(\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.

In [14]:
# 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¶

Optional worked reference

Head 2: pair each source weight with its value rowReceiver: the final “the”, at position 10Source j 1 The 2 fisherman 3 sat 4 beside 5 the 6 river 7 bank 8 and 9 watched10 theWeight α10,j0.0270.6930.0600.0270.0230.0230.0270.0270.0710.023Value row vⱼ[0.1, 2.3][2.1, 0.0][0.6, 1.9][0.1, 2.0][0.0, 2.3][0.0, 0.1][0.1, 0.8][0.1, 2.2][0.7, 1.7][0.0, 2.3]V = E WV10 source rows × 2fisherman: its weight pairs with its own two-number value row.

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.

In [15]:
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¶

Optional worked reference

Head 2: multiply both value coordinates by the same source weightReceiver: the final “the”, at position 10Source j 1 The 2 fisherman 3 sat 4 beside 5 the 6 river 7 bank 8 and 9 watched10 theWeight α10,j0.0270.6930.0600.0270.0230.0230.0270.0270.0710.023Value row vⱼ[0.1, 2.3][2.1, 0.0][0.6, 1.9][0.1, 2.0][0.0, 2.3][0.0, 0.1][0.1, 0.8][0.1, 2.2][0.7, 1.7][0.0, 2.3]Contribution α10,j vⱼ[0.003, 0.062][1.455, 0.000][0.036, 0.115][0.003, 0.054][0.000, 0.052][0.000, 0.002][0.003, 0.021][0.003, 0.059][0.050, 0.121][0.000, 0.052]×=fisherman: [α10,2 × (2.1), α10,2 × (0.0)] ≈ [1.455, 0.000]

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.

In [16]:
# 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¶

Optional worked reference

Head 2: sum all ten weighted value rows to obtain the messageReceiver: the final “the”, at position 10Source j 1 The 2 fisherman 3 sat 4 beside 5 the 6 river 7 bank 8 and 9 watched10 theWeight α10,j0.0270.6930.0600.0270.0230.0230.0270.0270.0710.023Value row vⱼ[0.1, 2.3][2.1, 0.0][0.6, 1.9][0.1, 2.0][0.0, 2.3][0.0, 0.1][0.1, 0.8][0.1, 2.2][0.7, 1.7][0.0, 2.3]Contribution α10,j vⱼ[0.003, 0.062][1.455, 0.000][0.036, 0.115][0.003, 0.054][0.000, 0.052][0.000, 0.002][0.003, 0.021][0.003, 0.059][0.050, 0.121][0.000, 0.052]×=Sum each contribution coordinate:m₁₀(2) =1.5520.538weight row [1 × 10] × V [10 × 2] = message [1 × 2]

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.

In [17]:
# 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¶

Optional worked reference

Two complete head calculations produce two messages before concatenation and output projectionHead 1q₁₀(1) K(1)ᵀ / √2ten scoressoftmaxOwn weight rowriver: 0.718× V(1)2.617−0.018m₁₀(1) [1 × 2]Head 2q₁₀(2) K(2)ᵀ / √2ten scoressoftmaxOwn weight rowfisherman: 0.693× V(2)1.5520.538m₁₀(2) [1 × 2]Both heads read the same input E. Their score lists and value sums stay separate.

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?¶

Optional worked reference

A four-coordinate dot product adds both matching contributions before one softmaxJoin the same query and key coordinates into one wider head.q₁₀ =2.32.32.30.0setting coordinatesperson coordinatesriver6.90+0.00= 6.90÷ √43.45fisherman4.83+4.83= 9.66÷ √44.83setting dotperson dotone score

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¶

Optional worked reference

The same projection coordinates yield one attention row or two separately normalized rowsReceiver 10 readsriverfishermanothersvalue widthOne wide head0.1780.7080.1134Head 1: setting0.7180.1660.1162Head 2: person0.0230.6930.2842One shared mixture of four coordinates, or two separate mixtures of two.

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.

In [18]:
# 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

What reaches the receiving token?¶

Optional worked reference

Combining the head messagesCombining the head messagesConcatenate, project into embedding space, then add the update.

Put the messages side by side¶

Optional worked reference

Concatenate two messages; do not average themm₁₀(1) · setting messagem₁₀(2) · person message2.617−0.0181.5520.5382.617−0.0181.5520.5382 coordinates + 2 coordinates → 4 coordinates

Concatenation keeps the two messages separate. It does not add them coordinate by coordinate.

Map the messages back to embedding space¶

Optional worked reference

Map the joined message back to the four input coordinatesjoined message2.617−0.0181.5520.538×WO · 4 × 4100001000.250100001head 1 rowshead 2 rowsΔe₁₀ =3.005−0.0181.5520.538First coordinate: 2.617 + 0.25 × 1.552 = 3.005

(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¶

Optional worked reference

Keep the original embedding row and add the context updateoriginal e₁₀0.0000.0000.0002.300waterfinancepersonglue+ update Δe₁₀3.005−0.0181.5520.538= updated e′₁₀3.005−0.0181.5522.838

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¶

Optional worked reference

The same next-token path, with every matrix shape in two parallel headsOne sentence: 10 tokens, width 4. Two heads, width 2 each.Input E[10×4]Head 1Q [10×2]K [10×2]V [10×2]QKᵀ: [10×10]scale, causal mask, softmaxA: [10×10]H(1) = A(1)V(1)[10×10] × [10×2]= [10×2]Head 2Q [10×2]K [10×2]V [10×2]QKᵀ: [10×10]scale, causal mask, softmaxA: [10×10]H(2) = A(2)V(2)[10×10] × [10×2]= [10×2]Concatenatetwo [10×2] → [10×4]Project with WO[10×4] × [4×4]keep original E [10×4]ΔE [10×4]+E′ = E + ΔE[10×4], same as ELast row [1×4]next-token prediction MLP

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¶

Optional worked reference

Two different reading patternsHead 1: setting cluesThe1fisherman2sat3beside4the5river6bank7and8watched9the100.71810 · theThicker arrow = more attention weightreceiverHead 2: person cluesThe1fisherman2sat3beside4the5river6bank7and8watched9the100.69310 · theThicker arrow = more attention weightreceiver

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.

Can every token do this at once?¶

Optional worked reference

The same calculation for every tokenThe same calculation for every tokenTrace the receiver row through the full matrices.

Project every row with the same head matrix¶

Optional worked reference

Project ten four-coordinate input rows into ten two-coordinate rowsE10 × 4×WQ(1)4 × 2Q(1)10 × 2Every token uses the same projection matrix within this head.

Ten input rows × four coordinates. Multiplying by a 4 × 2 projection gives ten query rows × two coordinates.

One head makes one attention grid¶

Optional worked reference

One head: query-key scores, a causal weight grid, then weighted valuesQ(1)10 × 2×K(1)ᵀ2 × 10÷ √2masksoftmaxA(1)10 × 10××××××××××××××××××××××××××××××××××××××××××××××V(1)10 × 2H(1)10 × 2Each highlighted row follows receiver 10. Columns of A are source positions.

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¶

Optional worked reference

Two independent attention grids computed from the same input snapshotsame E10 × 4Head 1own Q, K, Vown projectionsA(1)10 × 10×××××××××××××××××××××××××××××××××××××××××××××× VH(1)10 × 2Head 2own Q, K, Vown projectionsA(2)10 × 10×××××××××××××××××××××××××××××××××××××××××××××× VH(2)10 × 2

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¶

Optional worked reference

Concatenate columns, project, and preserve the token axisH(1)10 × 2H(2)10 × 2;Concat(H(1), H(2))10 × 4×WO4 × 4ΔE10 × 410 message rows → 10 update rows. The token count stays fixed.

(\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¶

Optional worked reference

One head in codeE → Q, K, V → scores → weights → message

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¶

Optional worked reference

Combine two head outputs in codeH(1), H(2) → concatenate columns → WO → ΔE

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¶

Optional worked reference

The same next-token path, with every matrix shape in two parallel headsOne sentence: 10 tokens, width 4. Two heads, width 2 each.Input E[10×4]Head 1Q [10×2]K [10×2]V [10×2]QKᵀ: [10×10]scale, causal mask, softmaxA: [10×10]H(1) = A(1)V(1)[10×10] × [10×2]= [10×2]Head 2Q [10×2]K [10×2]V [10×2]QKᵀ: [10×10]scale, causal mask, softmaxA: [10×10]H(2) = A(2)V(2)[10×10] × [10×2]= [10×2]Concatenatetwo [10×2] → [10×4]Project with WO[10×4] × [4×4]keep original E [10×4]ΔE [10×4]+E′ = E + ΔE[10×4], same as ELast row [1×4]next-token prediction MLPLossobserved target y

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¶

Optional worked reference

A bias adds a learned offset after the projection; the displayed offset is illustrativeOne head: qᵢ = eᵢ WQ + bQbias=Falsee₁₀ WQ2.32.3bias=True2.32.3+0.2−0.1=2.52.2bQ: learned offset [2]The same bQ is added to every token row in this head.

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?¶

Optional worked reference

The bias flag controls all query, key, value and output projection biases, not the separate prediction MLPInside nn.MultiheadAttentionQ = E WQ + bQoffset after input projectionK = E WK + bKoffset after input projectionV = E WV + bVoffset after input projectionΔE = Concat(H(1), H(2)) WO + bOoffset after output projectionOur worksheet and trained attention layers: bias=False.The separate prediction MLP still has its own biases.

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.

In [19]:
# 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¶

Optional worked reference

PyTorch returns the projected update, before our residualE → nn.MultiheadAttention → ΔE

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¶

Optional worked reference

The same next-token path, with every matrix shape in two parallel headsOne sentence: 10 tokens, width 4. Two heads, width 2 each.Input E[10×4]Head 1Q [10×2]K [10×2]V [10×2]QKᵀ: [10×10]scale, causal mask, softmaxA: [10×10]H(1) = A(1)V(1)[10×10] × [10×2]= [10×2]Head 2Q [10×2]K [10×2]V [10×2]QKᵀ: [10×10]scale, causal mask, softmaxA: [10×10]H(2) = A(2)V(2)[10×10] × [10×2]= [10×2]Concatenatetwo [10×2] → [10×4]Project with WO[10×4] × [4×4]keep original E [10×4]ΔE [10×4]+E′ = E + ΔE[10×4], same as ELast row [1×4]next-token prediction MLP

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¶

Optional worked reference

The trained MLP concatenates all token rows before its prediction layersSame 64-token input window · vocabulary size C = 4,000Token lookup64 rows × 64Concatenate1 × 4,096MLP hidden layer4,096 → 256 + ReLULogits256 → 4,000Different input slots occupy different parts of the flattened vector.No attention layer or separate position table in this baseline.Vocabulary softmax → next-token probabilities.

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¶

Optional worked reference

The trained single-head browser model, including learned positions, output projection, residual and prediction MLPSame 64-token input window · vocabulary size C = 4,000Token lookup64 rows × 64Position lookup64 rows × 64+Input E64 × 64One attention headQ, K, V: each 64 × 64A = softmax(masked scores)H = AV: 64 × 64Output WO64 × 64keep E+E′ = E + ΔElast row: 1 × 64Prediction MLP64 → 256 + ReLUVocabulary logits256 → 4,000Vocabulary softmax → next-token probabilities.

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¶

Optional worked reference

The trained four-head browser model, including learned positions, output projection, residual and prediction MLPSame 64-token input window · vocabulary size C = 4,000Token lookup64 rows × 64Position lookup64 rows × 64+Input E64 × 64Head 1: H is 64 × 16Head 2: H is 64 × 16Head 3: H is 64 × 16Head 4: H is 64 × 16Concat H: 64 × 64Output WO64 × 64keep E+E′ = E + ΔElast row: 1 × 64Prediction MLP64 → 256 + ReLUVocabulary logits256 → 4,000Vocabulary softmax → next-token probabilities.

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¶

Optional worked reference

More heads divide a fixed total width into smaller per-head projections1 head × 64 coordinates644 heads × 16 coordinates16161616WQ, WK, WV and WO stay 64 × 64.

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¶

Optional worked reference

Model; Cross-entropy ↓; Perplexity ↓ModelCross-entropy ↓Perplexity ↓Embedding → MLP3.94351.59One attention head3.44531.34Four attention heads3.36028.78

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¶

Optional worked reference

Model; Parameters; Training / seedModelParametersTraining / seedEmbedding → MLP2,332,83258.2 sOne attention head1,321,12065.6 sFour attention heads1,321,12069.7 s

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¶

Optional worked reference

Compare real language models in the browserTraining prefix → held-out story → a different kind of promptCompare the continuation and generation time on your device.

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¶

Optional worked reference

Image patches become embedding rows through one shared projectionA tiny 8 × 8 grayscale imageFour 4 × 4 patchesFlatten each patch12344 rows × 16 pixel valuesWpatch16 × 4Patch embeddingsEpatch: 4 × 4The shared projection makes one embedding row per patch.

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¶

Optional worked reference

A CLS row gathers image information before a two-class predictionFour patch rows + one learned [CLS] row + position rows[CLS]patch 1patch 2patch 3patch 4E: 5 × 4Transformer encoderMHA: 2 heads × 2 coordinatesthen a per-row MLPresiduals + normalization insideUpdated rows: 5 × 4All image patches are visible: no future-token mask.Final [CLS]: 1 × 4Classifier: 4 → 2Two class logitsSoftmax → vertical-stripe / horizontal-stripe probabilities

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?¶

a b i → Embeddings → Read context → Next charactera b i3 known IDsEmbeddings3 learned rowsRead contextMLP or attentionNext characterfor example, d

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¶

The river sentence from Part IIThe1fisherman2sat3beside4the5river6bank7and8watched9the10Predict the next token after the final “the”.Which setting?Who is in the scene?

One query can benefit from several clues at once. What setting are we in? Who is there?

Keep the two examples separate¶

Worksheet; Trained experimentWorksheetTrained experiment2 heads × 2 coordinates4 heads × 16 coordinatesHand-chosen projectionsLearned from TinyStoriesPart II token + position rowsSame 64-token benchmark window

The worksheet explains the arithmetic. Held-out results later tell us whether the trained models improved.

From one weight row to two¶

Known inputs → Two heads → One predictionKnown inputsTwo headsOne prediction

What changes when the same input passes through two sets of projections?

One head produces one message¶

Full E → Q, K, V → Weights A → Message AVFull Eall coordinatesQ, K, Vlearned projectionsWeights Aone row per queryMessage AVone weighted sum

One head can already read several sources. Its value coordinates share the same attention weights.

Two heads keep two messages¶

Two heads inside the same next-token modelKnown tokensIDs [B,T]Lookup + positionE [B,T,D]Head 1Q¹ · K¹ · V¹Head 2Q² · K² · V²Concatenatemessages [B,T,D]W_O + residualE′ = E + ΔELast row → MLPnext-token logitsKeep E for the residual

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¶

E [B,T,4] → W_Q [4,4] → Q [B,T,4] → 2 heads × 2E [B,T,4]full input rowsW_Q [4,4]learned mixingQ [B,T,4]projected coordinates2 heads × 2one view per head

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¶

Input coordinate; W_Q: head 1; W_Q: head 2Input coordinateW_Q: head 1W_Q: head 2water[0, 0][0, 0]finance[0, 0][0, 0]person[0, 0][0, 0]glue[1, 1][1, 0]

These are two 4×2 matrices placed side by side. Every head can learn from every input coordinate.

Two heads, one receiving token¶

Known inputs → Two heads → One predictionKnown inputsTwo headsOne prediction

Keep the final “the” at position 10 as the query throughout.

Start with the same position-aware rows¶

Position / token; water; finance; person; gluePosition / tokenwaterfinancepersonglue2 fisherman2.00.12.10.06 river3.1−0.10.00.110 the0.00.00.02.3

E already includes token embedding + position embedding. Token IDs are not embedding coordinates.

The first query asks about the setting¶

e₁₀ → W_Q¹ → q₁₀¹e₁₀[0.000, 0.000, 0.000, 2.300]W_Q¹4 × 2q₁₀¹[2.300, 2.300]

Our chosen projection copies the glue coordinate into both query coordinates: [2.3, 2.3].

The second query uses a different projection¶

e₁₀ → W_Q² → q₁₀²e₁₀[0.000, 0.000, 0.000, 2.300]W_Q²4 × 2q₁₀²[2.300, 0.000]

The same input row now produces [2.3, 0.0]. The query changes because the projection matrix changes.

The source keys are different too¶

Source; Key in head 1; Key in head 2SourceKey in head 1Key in head 2fisherman[2.000, 0.100][2.100, 0.000]river[3.100, −0.100][0.000, 0.100]bank[0.700, 0.700][0.100, 0.800]the[0.000, 0.000][0.000, 2.300]

Head 1 exposes water/finance features. Head 2 exposes person/glue features. These labels belong to this designed worksheet.

Head 1 scores the river key¶

Calculation; ValueCalculationValueq₁₀¹[2.3, 2.3]k₆¹ for river[3.1, −0.1]Dot product2.3 × 3.1 + 2.3 × (−0.1) = 6.9Divide by √24.879

The scaling uses the head width: two coordinates, not four.

Head 2 scores the fisherman key¶

Calculation; ValueCalculationValueq₁₀²[2.3, 0.0]k₂² for fisherman[2.1, 0.0]Dot product2.3 × 2.1 + 0.0 × 0.0 = 4.83Divide by √23.415

We compare each head’s query only with keys from that same head.

Head 1: scores become a weight row¶

Attention weights for head 1Head 1 · final query at position 10The0.160.006fisherman3.420.166sat0.490.009beside0.810.012the0.330.008river4.880.718bank2.280.053and0.000.005watched1.140.017the0.000.005scoreweightOne softmax over the ten sources. Row sum = 1.

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¶

Attention weights for head 2Head 2 · final query at position 10The0.160.027fisherman3.420.693sat0.980.060beside0.160.027the0.000.023river0.000.023bank0.160.027and0.160.027watched1.140.071the0.000.023scoreweightOne softmax over the ten sources. Row sum = 1.

Normalize over sources within head 2. Do not normalize across heads or average their scores.

Head 1 sends water and finance information¶

Source; Weight; Value row; Weighted contributionSourceWeightValue rowWeighted contributionriver0.718[3.100, −0.100][2.225, −0.072]All other sources[0.392, 0.054]Total message[2.617, −0.018]

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¶

Source; Weight; Value row; Weighted contributionSourceWeightValue rowWeighted contributionfisherman0.693[2.100, 0.000][1.455, 0.000]All other sources[0.097, 0.538]Total message[1.552, 0.538]

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¶

Known inputs → Two heads → One predictionKnown inputsTwo headsOne prediction

How do two small message rows become one update to the original row?

Concatenate the messages, not the tokens¶

Head; Message coordinates; MessageHeadMessage coordinatesMessage1water / finance[2.617, −0.018]2person / glue[1.552, 0.538]Joinedhead 1, then head 2[2.617, −0.018, 1.552, 0.538]

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¶

Joined coordinate; Output weightsJoined coordinateOutput weights1[1, 0, 0, 0]2[0, 1, 0, 0]3[0.25, 0, 1, 0]4[0, 0, 0, 1]

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¶

Row; Four coordinatesRowFour coordinatesOriginal e₁₀[0.000, 0.000, 0.000, 2.300]Attention update Δe₁₀[3.005, −0.018, 1.552, 0.538]Updated e′₁₀[3.005, −0.018, 1.552, 2.838]

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¶

e′₁₀ → Hidden + ReLU → 20 logits → Softmaxe′₁₀4 coordinatesHidden + ReLU8 activations20 logitsone per vocabulary itemSoftmaxnext-token probabilities

Choose a token only after vocabulary softmax. Attention weights choose source positions; vocabulary probabilities choose possible next tokens.

Change the context; inspect both heads¶

Attention weights for head 1Head 1 · final query at position 10The0.160.006fisherman3.420.166sat0.490.009beside0.810.012the0.330.008river4.880.718bank2.280.053and0.000.005watched1.140.017the0.000.005scoreweightOne softmax over the ten sources. Row sum = 1.

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.

In [20]:
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]]

From the drawing to tensors¶

Known inputs → Two heads → One predictionKnown inputsTwo headsOne prediction

Keep B = 2 examples, T = 10 tokens, D = 4 coordinates and n_heads = 2.

Create the tables, projections and prediction MLP¶

Token + position → ScratchMultiHead → Hidden → vocabularyToken + position20 × 4 and 10 × 4ScratchMultiHead4 coordinates, 2 headsHidden → vocabulary4 → 8 → 20

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.

In [21]:
"""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
In [22]:
model = TinyMultiHeadLM(vocab_size=20, context=10,
                        width=4, heads=2, hidden=8)
load_worksheet_weights(model, worksheet)
Out[22]:
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¶

Example; Ten token IDs; Observed next tokenExampleTen token IDsObserved next tokenriver[0, 1, 2, 3, 0, 4, 5, 6, 7, 0]watercheque[8, 9, 0, 10, 11, 0, 5, 6, 7, 0]teller

These are integer IDs, not embeddings. Each example has one observed target.

In [23]:
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¶

Two heads inside the same next-token modelKnown tokensIDs [B,T]Lookup + positionE [B,T,D]Head 1Q¹ · K¹ · V¹Head 2Q² · K² · V²Concatenatemessages [B,T,D]W_O + residualE′ = E + ΔELast row → MLPnext-token logitsKeep E for the residual

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¶

Token IDs [2,10] → Token + position rows → E [2,10,4]Token IDs [2,10]integersToken + position rowslookup, then addE [2,10,4]floating-point vectors

The embedding tables are learned parameters. E contains floating-point vectors with shape [2,10,4].

In [24]:
positions = torch.arange(X.shape[1])
E = model.token_embedding(X) + model.position_embedding(positions)

Project the full input three times¶

E [2,10,4] → Three matrices → Q, K, VE [2,10,4]same inputThree matriceseach 4 × 4Q, K, Veach [2,10,4]

Different learned matrices give matching queries, matching keys and transmitted values.

In [25]:
Q = model.attention.W_Q(E)
K = model.attention.W_K(E)
V = model.attention.W_V(E)

Make the head axis explicit¶

[2,10,4] → [2,10,2,2] → [2,2,10,2][2,10,4]B, T, D[2,10,2,2]B, T, heads, d_head[2,2,10,2]B, heads, T, d_head

reshape groups projected coordinates; transpose puts heads before tokens. Apply the same operation to K and V.

In [26]:
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¶

Q [2,2,10,2] → Kᵀ [2,2,2,10] → Scores [2,2,10,10]Q [2,2,10,2]query rowsKᵀ [2,2,2,10]source columnsScores [2,2,10,10]two 10 × 10 grids

The first two axes keep examples and heads separate. Matrix multiplication contracts only the matching-coordinate axis.

In [27]:
scores = Q @ K.transpose(-2, -1)
scores = scores / math.sqrt(2)

Block future sources in every head¶

Receiving position; Allowed source positionsReceiving positionAllowed source positions1121, 231, 2, 3101 through 10

True marks a blocked entry. The same causal mask broadcasts across both examples and both heads.

In [28]:
future = torch.ones(10, 10, dtype=torch.bool).triu(1)
scores = scores.masked_fill(future, float("-inf"))

Each row becomes a distribution over sources¶

Attention weights for head 1Head 1 · final query at position 10The0.006fisherman0.166sat0.009beside0.012the0.008river0.718bank0.053and0.005watched0.017the0.005weightOne softmax over the ten sources. Row sum = 1.

The final axis indexes source tokens. Every allowed row sums to one; masked entries have weight zero.

In [29]:
A = scores.softmax(dim=-1)
assert A.shape == (2, 2, 10, 10)
In [30]:
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¶

A [2,2,10,10] → V [2,2,10,2] → Messages [2,2,10,2]A [2,2,10,10]source weightsV [2,2,10,2]source contentMessages [2,2,10,2]one row per query/head

The source-token axis is summed out. Each head keeps its own two-coordinate message.

In [31]:
messages = A @ V

Put each token’s head messages side by side¶

[2,2,10,2] → [2,10,2,2] → [2,10,4][2,2,10,2]B, heads, T, d_head[2,10,2,2]B, T, heads, d_head[2,10,4]B, T, D

Transpose before reshaping. A direct reshape of the original layout would mix token rows with head rows.

In [32]:
joined = messages.transpose(1, 2).contiguous()
joined = joined.reshape(2, 10, 4)

Return to the shared prediction path¶

Two heads inside the same next-token modelKnown tokensIDs [B,T]Lookup + positionE [B,T,D]Head 1Q¹ · K¹ · V¹Head 2Q² · K² · V²Concatenatemessages [B,T,D]W_O + residualE′ = E + ΔELast row → MLPnext-token logitsKeep E for the residual

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¶

Joined messages → W_O + residual → Last row → MLPJoined messages[2,10,4]W_O + residualupdated E′ [2,10,4]Last row → MLPlogits [2,20]

W_O combines the heads. The residual keeps E. Only the last updated row feeds this next-token loss.

In [33]:
delta = model.attention.W_O(joined)
updated = E + delta
logits = model.readout(F.relu(model.hidden(updated[:, -1])))
In [34]:
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¶

Known inputs → Two heads → One predictionKnown inputsTwo headsOne prediction

More heads change the context update, not the definition of the next-token target.

Score the two observed targets¶

Input X → Two heads + MLP → Logits [2,20] → LossInput X2 × 10 token IDsTwo heads + MLPone prediction/exampleLogits [2,20]targets: water, tellerLossone scalar

Cross-entropy reads logits and observed token IDs. Do not sample a generated token to make the training target.

In [35]:
logits = model(X)
loss = F.cross_entropy(logits, y)

One loss trains all the projections¶

Loss → Gradients → Optimizer → Updated weightsLossfrom the same targetsGradientsall learned parametersOptimizerone updateUpdated weightsused by the next batch

Heads are not assigned jobs or separate labels. They receive gradients from the same prediction loss.

In [36]:
optimizer = torch.optim.AdamW(model.parameters(), lr=0.001)
optimizer.zero_grad()
loss.backward()
optimizer.step()
In [37]:
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¶

River prefix → Vocabulary lookup → History [1,10]River prefixthe ten known wordsVocabulary lookupsame vocabularyHistory [1,10]no observed next token

We already encoded the river sentence. A real application tokenizes and looks up a new prompt with the same vocabulary.

In [38]:
history = torch.tensor([river_ids])
model.eval()
Out[38]:
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¶

Known IDs → Crop to context → Model logits → Choose + appendKnown IDsno future targetCrop to contextlatest 10 in this toyModel logitsfinal row onlyChoose + appendnew known prefix

At inference, weights stay fixed. The current prefix becomes longer; we keep only the configured context window.

In [39]:
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¶

Two heads inside the same next-token modelKnown tokensIDs [B,T]Lookup + positionE [B,T,D]Head 1Q¹ · K¹ · V¹Head 2Q² · K² · V²Concatenatemessages [B,T,D]W_O + residualE′ = E + ΔELast row → MLPnext-token logitsKeep E for the residual

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.

The same operation, packaged by PyTorch¶

Known inputs → Two heads → One predictionKnown inputsTwo headsOne prediction

Which boxes does nn.MultiheadAttention replace?

Replace the head calculation, not the whole model¶

Two heads inside the same next-token modelKnown tokensIDs [B,T]Lookup + positionE [B,T,D]Head 1Q¹ · K¹ · V¹Head 2Q² · K² · V²Concatenatemessages [B,T,D]W_O + residualE′ = E + ΔELast row → MLPnext-token logitsKeep E for the residual

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¶

Input E → nn.MultiheadAttention → Projected update ΔEInput E[B,T,4]nn.MultiheadAttention2 heads × 2 coordinatesProjected update ΔE[B,T,4]

batch_first=True means [B,T,D]. bias=False and dropout=0 match our scratch implementation.

In [40]:
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¶

E, E, E → nn.MultiheadAttention → ΔE and AE, E, Eunprojected input rowsnn.MultiheadAttentionQ/K/V, softmax, mix, W_OΔE and Anot a residual yet

PyTorch performs the learned projections internally. average_attn_weights=False preserves the head axis in the returned weight tensor.

In [41]:
delta, A = mha(E, E, E, attn_mask=future,
               average_attn_weights=False)
updated = E + delta

Compare the same weights, not two random layers¶

Quantity; Expected shapeQuantityExpected shapeProjected update ΔE[2,10,4]Separate head weights A[2,2,10,10]Residual E + ΔE[2,10,4]

The notebook copies the scratch Q/K/V and W_O parameters into PyTorch, then checks both outputs numerically.

In [42]:
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)
In [43]:
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¶

Token + position → Multi-head attention → Residual → Prediction MLPToken + positionoutside the APIMulti-head attentionreturns projected ΔEResidualadd E ourselvesPrediction MLPoutside the API

nn.MultiheadAttention does not add position embeddings, the residual, a vocabulary head or a training loop.

Do more heads help this experiment?¶

Known inputs → Two heads → One predictionKnown inputsTwo headsOne prediction

Move from the small worksheet to the actual TinyStories checkpoints.

More heads need not mean more parameters¶

Setting; Total width; Heads; Width per headSettingTotal widthHeadsWidth per headOne head64164Four heads64416

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¶

Model; Test cross-entropy ↓; Test perplexity ↓ModelTest cross-entropy ↓Test perplexity ↓MLP3.94351.59One head3.44531.34Four heads3.36028.78

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¶

Model; Parameters; Training per seedModelParametersTraining per seedMLP2,332,83258.2 sOne head1,321,12065.6 sFour heads1,321,12069.7 s

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¶

Training prefix → Held-out prefix → Outside-domain promptTraining prefixfamiliar textHeld-out prefixnew story, same taskOutside-domain promptdifferent kind of text

Open the live browser demo ↗ Compare continuations, vocabulary coverage and generation time. This runs real checkpoints with WebGPU or WASM.

Run the calculation yourself¶

Resource; What to doResourceWhat to doNotebook 7Follow every operation and check PyTorch parityNotebook 6Inspect the trained four-head checkpoint and resultsPart 2B · optionalGradients, normalization, full blocks and context cost

Step-by-step notebook · Measured experiment · Optional reference

Next: read a different sequence¶

French decoder row → Query → English keys + valuesFrench decoder rowthe next output tokenQuerywhat do I need?English keys + valueswhere should I read?

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.