Multi-head attention · a worked continuation
Why more than one head?
Which parts of this prefix would help you choose a continuation?
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 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.
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.
Multi-head attention · a worked continuation
The complete multi-head operation
Same sentence, separate readings. Each head learns its own queries, keys and values.
These are possible roles, not assigned jobs or measured explanations of the trained heads. The prefix is “The fisherman sat beside the river bank and watched the ___”. Each head sees every allowed source.
Concatenation joins the coordinates. The output projection learns how both heads contribute to the update.
Both heads read the same E. Their messages combine; the final updated row predicts the next token.
H = AV is the message matrix. This diagram keeps the two-head worksheet dimensions: 10 tokens, width 4, two heads of width 2. Next we switch to the trained TinyStories model: 64 slots, width 64, four heads of width 16.
Optional: every worked calculation and the interactive head explorer · Executable Notebook 7.
Multi-head attention · a worked continuation
TinyStories: text becomes training examples
sample = evidence['sample']
print(sample['text_excerpt'])The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
The saved experiment uses 6,000 TinyStories documents. Each document is a complete story.
The saved experiment uses 6,000 TinyStories documents. Each document is a complete story. This is an excerpt from source row 12400.
Original worked example and executable code
TinyStories by Ronen Eldan and Yuanzhi Li. Source dataset, CDLA-Sharing-1.0. Source revision, row IDs and unchanged texts are recorded in story_examples.json. […] marks omitted text.
short_story = story_examples['examples'][0]
print(short_story['text'])The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
One document is one complete story, not one sentence. This is the shortest story in our saved 6,000-document subset; many others are longer.
One document is one complete story, not one sentence. This is the shortest story in our saved 6,000-document subset; many others are longer.
Original worked example and executable code
TinyStories by Ronen Eldan and Yuanzhi Li. Source dataset, CDLA-Sharing-1.0. Source revision, row IDs and unchanged texts are recorded in story_examples.json. […] marks omitted text.
long_stories = story_examples['examples'][1:]
print([s['token_count'] for s in long_stories])The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
These beginnings and endings come from two other documents in the same subset. […] marks omitted text; the lengths count each whole story.
These beginnings and endings come from two other documents in the same subset. […] marks omitted text; the lengths count each whole story. Words are whitespace-separated items; tokens also separate punctuation and exclude start/end markers.
Original worked example and executable code
TinyStories by Ronen Eldan and Yuanzhi Li. Source dataset, CDLA-Sharing-1.0. Source revision, row IDs and unchanged texts are recorded in story_examples.json. […] marks omitted text.
Is a 300-token story a single 64-token training input?
No. A story is a document. Later, we turn it into many next-token examples, each using at most the previous 64 tokens in the saved experiment.
audit = evidence['audit']
print(audit['documents']) # whole-story splitsThe coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
The split happens before overlapping windows are made. Training fits parameters and the vocabulary.
The split happens before overlapping windows are made. Training fits parameters and the vocabulary. Validation chooses settings and checkpoints. Test measures the frozen choice.
sentence = 'Lily found a red ball.'The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
We will trace this complete authored example using a small vocabulary of 10 items and four input slots. The trained experiment later uses 4,000 items and 64 slots.
We use this authored sentence for the arithmetic. Its small vocabulary and randomly initialized models are separate from the measured TinyStories experiment.
tokenization_text = 'redder!'
print(tokenize(tokenization_text))The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
A token is one item the model reads or predicts. The tokenizer chooses these items before we assign integer IDs and look up learned embeddings.
A token is one item the model reads or predicts. The tokenizer chooses these items before we assign integer IDs and look up learned embeddings.
Original worked example and executable code
Hugging Face tokenizer overview. The subword split is illustrative, not a fitted BPE output. The executable English tokenizer is in wordlm.py.
word_tokens = tokenize('redder!')
character_tokens = list('redder!')
illustrative_subwords = ['red', 'der', '!'] # hand-chosenThe coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
Whole-word vocabularies need entries for many word forms, while characters produce longer sequences. Subword methods such as byte pair encoding (BPE) learn reusable pieces from training text to balance these concerns.
Whole-word vocabularies need entries for many word forms, while characters produce longer sequences. Subword methods such as byte pair encoding (BPE) learn reusable pieces from training text to balance these concerns.
Original worked example and executable code
Hugging Face tokenizer overview. The subword split is illustrative, not a fitted BPE output. The executable English tokenizer is in wordlm.py.
Does a context window of four tokens always hold four words?
No. It holds four items produced by that tokenizer. A word can take several subword or character tokens, and punctuation can use a token too.
print(tokenize("Lily can't find 12 balls!"))The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
All three benchmark models use these same word-and-punctuation rules. Build the vocabulary from training stories only; map a missing item to UNK.
Words plus punctuation keep our calculations easy to inspect. All three models use the same rules and vocabulary during training and generation. We build the vocabulary from training stories only and map missing words to the unknown-token marker, UNK.
Original worked example and executable code
Hugging Face tokenizer overview. The subword split is illustrative, not a fitted BPE output. The executable English tokenizer is in wordlm.py.
pieces = tokenize(sentence)The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
This tokenizer normalizes Unicode, lowercases text and keeps punctuation as tokens. Five words plus the full stop give six tokens.
This tokenizer normalizes Unicode, lowercases text and keeps punctuation as tokens. Five words plus the full stop give six tokens. A tokenizer decides what one prediction unit is.
words = list(SPECIAL_TOKENS) + sorted(set(pieces))
vocab = Vocabulary(words, {t:i for i,t in enumerate(words)}, {}, 10, 1)The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
Six ordinary tokens plus four special tokens give C = 10 in this teaching example. IDs name vocabulary items; they are not embedding coordinates.
Six ordinary tokens + four special tokens = C=10; their IDs are arbitrary labels. Here, BOS and EOS mark one complete story (a sequence), not each sentence. PAD means padding; BOS means beginning of sequence; EOS means end of sequence; UNK means unknown token. The benchmark uses a different, frequency-ranked vocabulary with C=4,000.
print(vocab.encode_tokens(['blue'], boundaries=False)) # UNK: [3]
print(vocab.encode_tokens(['blue'], boundaries=True)) # [1, 3, 2]The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
“blue” is absent from this toy vocabulary, so its ID is 3, UNK. BOS and EOS mark the whole story; PAD fills unused input slots.
blue is missing from our toy vocabulary, so this lookup returns [3], the UNK ID. boundaries=False means “do not add BOS or EOS”. The input ['blue'] is already a list of tokens; encode_tokens looks up their integer IDs, without creating embeddings. With boundaries=True, the same call returns [1, 3, 2]: BOS, UNK, EOS.
ids = vocab.encode_tokens(pieces, boundaries=True)The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
One BOS starts the whole story and one EOS ends it, even if it contains several sentences. Full stops remain ordinary tokens inside the story.
In this notebook, each document is a complete story with one BOS at its start and one EOS at its end. We do not add extra markers between sentences. The Lily sentence is our entire toy document, so its six ordinary tokens become eight IDs including BOS and EOS. BOS provides context, and the remaining seven IDs are prediction targets. encode_tokens wraps the token list we pass it, without detecting sentences. We pass a complete story once before creating its training windows, which never cross document boundaries. Other datasets may use different boundary conventions.
Original worked example and executable code
A story has three sentences. How many BOS and EOS markers do we add?
One BOS before the whole story and one EOS after it. Full stops remain ordinary tokens inside the document. Our one-sentence example follows the same rule.
w = 4
target_position = 4
print(list(enumerate(ids)))The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
Python positions start at 0: position 4 contains red, whose vocabulary ID is 9. With a four-token window, this example reads positions 0 through 3 and predicts the token at position 4.
Python positions start at 0: position 4 contains red, whose vocabulary ID is 9. With a four-token window, this example reads positions 0 through 3 and predicts the token at position 4.
context_ids = ids[target_position-w:target_position]
target_id = ids[target_position]The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
context_ids is the input for one example, often called x, and target_id is its single observed answer, often called y. The slice ids[0:4] stops before position 4, so red is outside this input.
context_ids is the input for one example, often called x, and target_id is its single observed answer, often called y. The slice ids[0:4] stops before position 4, so red is outside this input.
t = 1
visible = ids[max(0, t-w):t]
context_ids, target_id = [0] * (w-len(visible)) + visible, ids[t]The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
At position t=1, the visible prefix is [1], meaning BOS. Three PAD IDs fill the unused slots, giving context_ids=[0, 0, 0, 1] and target_id=8 for lily.
At position t=1, the visible prefix is [1], meaning BOS. Three PAD IDs fill the unused slots, giving context_ids=[0, 0, 0, 1] and target_id=8 for lily.
t = 2
visible = ids[max(0, t-w):t]
context_ids, target_id = [0] * (w-len(visible)) + visible, ids[t]The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
Move one token forward. The observed “lily” joins the input, and “found” becomes the next target. Training uses the story’s words, not the model’s guesses.
At t=2, BOS and lily form the known prefix, and the observed next word is found. Computing this pair changes context_ids and target_id, while the stored lists still contain only the first pair.
contexts, targets = [], []
for t in range(1, len(ids)):
visible = ids[max(0, t-w):t]
contexts.append([0] * (w-len(visible)) + visible); targets.append(ids[t])The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
Repeat that operation for each target. These are the first four input–target pairs, with IDs translated back into tokens.
These are the first four stored rows, with the IDs translated back to tokens. Row 3 is the red-target pair we inspected first, now stored in contexts[3] and targets[3].
print(contexts[4:], targets[4:]) # older tokens fall outside the windowThe coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
A full window keeps the previous four tokens and drops older ones. The last target is EOS; windows never cross into another story.
The remaining three rows use full windows, keeping only the four tokens immediately before each target. The last target is EOS, and no window crosses into a different story.
all_X = torch.tensor(contexts, dtype=torch.long)
all_y = torch.tensor(targets, dtype=torch.long)The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
Store all seven input rows in all_X [7, 4] and their seven answers in all_y [7]. These are integer token IDs, not learned vectors yet.
The conversion preserves every token ID and row pairing. all_X has shape [7, 4], all_y has shape [7], and torch.long means integer IDs.
Original worked example and executable code
Have these token IDs become embeddings yet?
No. This step only organizes the IDs into tensors. The embedding layer will later look up a learned vector for each ID.
examples_in_story = len(pieces) + 1 # one extra target: EOSThe coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
Each ordinary token is a target once, and EOS supplies one more target. BOS starts the history, while PAD fills empty input slots, so neither adds a target.
Each ordinary token is a target once, and EOS supplies one more target. BOS starts the history, while PAD fills empty input slots, so neither adds a target.
window_counts = {s: audit['oov'][s]['tokens'] + n
for s, n in audit['documents'].items()}The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
For each split, add its ordinary-token count and its story count. These are counts of supervised examples, not optimizer steps.
For each split, add its ordinary-token count and its story count. These are counts of supervised examples, not optimizer steps.
Original worked example and executable code
Does w=8 create more targets than w=4 here?
No. With this padding rule, both give seven targets. A wider window changes the visible input history, not the target count.
history = ids[:6]
contexts_by_width = {w: history[-w:] for w in [2, 4, 6]}The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
History is everything already known. The context window is the suffix read for this prediction.
History is everything already known. The context window is the suffix read for this prediction. Increasing w can recover an older clue, but also changes model size or computation.
selected = torch.tensor([2, 3])
X, y = all_X[selected], all_y[selected]The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
Each batch row pairs four input token IDs in X with one observed next-token ID in y. The text columns decode the IDs, and the targets come from the story.
Each batch row pairs four input token IDs in X with one observed next-token ID in y. The text columns decode the IDs, and the targets come from the story.
print(X.shape, X.dtype) # two rows of four IDs
print(y.shape, y.dtype) # two observed next-token IDsThe coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
X contains two examples, each with four token IDs. y contains their two next-token IDs. Embedding lookup transforms X; the loss still uses y as IDs.
Tokenization and vocabulary lookup are complete: X and y contain integers, with no embedding coordinates yet. Next, X goes through embedding lookup, while y stays as target IDs for the loss.
Original worked example and executable code
Are the four numbers in each row of X embedding coordinates?
No. They are IDs for four separate token slots. Embedding lookup later replaces each input ID with a learned vector. X.shape describes the size of the ID tensor, not its contents.
B, N = 2, len(all_y)
batch_sizes = [len(all_y[start:start+B]) for start in range(0, N, B)]The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
Seven examples with B = 2 give batches of 2, 2, 2 and 1 in a sequential pass. The saved experiment instead samples 512 windows per update, for 6,000 updates.
B=2 with drop_last=False gives batches of 2, 2, 2 and 1 per pass. With one update per batch, that is four optimizer steps. The benchmark samples 512 windows per update with replacement.
C, T, D, heads = 4000, 64, 64, 4
training_batch_size = 512The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
The rules stay the same: tokenize, look up IDs, add story boundaries, then create padded or cropped input–target pairs.
The toy vocabulary was chosen to expose every ID. The experimental vocabulary is frequency-ranked from training stories, so ordinary words receive different IDs. The four special IDs remain 0–3. The next slide remaps the same two prefixes into that real vocabulary.
prefixes = [['<BOS>', 'lily', 'found'], ['<BOS>', 'lily', 'found', 'a']]
X = torch.tensor([[0]*(T-len(p)) + [real_stoi[t] for t in p] for p in prefixes])
y = torch.tensor([real_stoi['a'], real_stoi['red']])The coloured, thick outline marks the current operation. The complete map remains visible. The top row prepares data. Four parallel head lanes form the model. The bottom branches separate learning from generation. Dashed V routes bypass score calculation. The dashed E route carries the residual.
The small data example uses C=10 and T=4. We switch to C=4,000 and T=64 before running the model. The diagram previews its width-64, four-head architecture. The original data example follows on the next frame.
Same prefixes and answers; new ordinary-word IDs and 64 slots. We will follow this B = 2 batch through the trained model’s architecture.
Lily’s sentence is an authored teaching example, not a claimed source story. The IDs here use the actual 4,000-item benchmark vocabulary. PAD is ID 0. These two rows illustrate batching; the saved training runs use batches of 512. Position IDs are 0–63 within each padded/cropped window.
Multi-head attention · a worked continuation
TinyStories: the multi-head model, training and generation
from multihead import MultiHeadAttentionLM
model = MultiHeadAttentionLM(C, T, d_model=64, hidden=256, heads=4)
B, T = X.shapeC=4,000 vocabulary items, T=64 input slots, model width 64, four heads of width 16 and a 256-unit prediction MLP. B=2 for our displayed batch. The highlighted boxes contain learned parameters. Other boxes reshape tensors or carry out fixed operations. The class is defined in multihead.py.
token_rows = model.token_embedding(X)The ID for “found” selects one row of 64 learned numbers. All occurrences share that token row. B counts examples, not heads.
The PAD token row is initialized to zero and receives no embedding-lookup gradient here, using padding_idx=0. BOS and UNK can learn when used in inputs. EOS is a target rather than an input in these windows, so its input row receives no task gradient. The separate classifier learns an EOS output score. Padding still needs an attention mask.
# Learned 64 x 64 position table, not sinusoidal.
# Trained together with the token embeddings.
positions = torch.arange(T, device=X.device)
E = token_rows + model.position_embedding(positions)[None]This experiment uses learned absolute position embeddings, not sinusoidal encodings. The 64 × 64 table has one trainable row for each slot 0–63 within the padded or cropped window. Add that row to the token embedding. Both tables learn from the prediction loss.
Q = model.W_Q(E).reshape(B, T, 4, 16).transpose(1, 2)
K = model.W_K(E).reshape(B, T, 4, 16).transpose(1, 2)
V = model.W_V(E).reshape(B, T, 4, 16).transpose(1, 2)Every head reads the full input row through learned projections. Split the projected coordinates into four groups of 16.
For the two examples, each projected tensor is [2,4,64,16]: examples, heads, token slots, coordinates per head. Packed W_Q, W_K and W_V each map 64 input coordinates to 64 output coordinates.
scores = (Q @ K.transpose(-2, -1)) / math.sqrt(16)For “BOS lily found”, follow the query at “found”. Its scores compare that receiver with the keys of the available source tokens.
future = torch.ones(T, T, dtype=torch.bool, device=E.device).triu(1)
real = X.ne(0)
mask = future[None] | ((~real[:, None, :]) & real[:, :, None])
A = torch.softmax(scores.masked_fill(mask[:, None], float("-inf")), dim=-1)At “found”, allow BOS, lily and found. PAD and future slots receive zero weight. Each real receiver’s source weights sum to one.
This is the full causal teaching view for real receiver rows. The implementation keeps PAD-only receiver rows numerically defined; they do not feed the loss. The saved training and browser forward path computes only the final receiver’s query, since this task predicts one next token per window. It produces the same final logits with smaller [B,4,1,64] weight tensors.
messages = A @ VA weight scales a complete 16-number value row. Add those weighted rows to make one message per receiver, per head.
joined = messages.transpose(1, 2).reshape(B, T, 64)
updated = E + model.W_O(joined)Joining four 16-coordinate messages restores width 64. W_O learns how to combine them; the residual adds this update to the original E.
hidden = torch.relu(model.hidden_layer(updated[:, -1]))
logits = model.vocab_head(hidden)
loss = F.cross_entropy(logits, y) # raw logits and target IDs“BOS lily found” → target a. “BOS lily found a” → target red. The model supplies scores; the story supplies the answers.
optimizer.zero_grad()
loss.backward()
optimizer.step()Repeat over batches. The loss trains token rows, position rows, all four heads, W_O and the prediction MLP. Validation selects a checkpoint; test measures that frozen choice.
We do not assign head-specific targets such as “setting” or “person”. All heads learn through the same next-token task. These are one-block teaching models, not a full stack of normalized Transformer blocks.
history = [real_stoi["<BOS>"]] + [real_stoi.get(t, 3) for t in tokenize("Lily found")]
visible = history[-T:]
X = torch.tensor([[0] * (T-len(visible)) + visible])Use the training vocabulary and prepend BOS. Do not append EOS to an unfinished prompt. Load the chosen checkpoint and keep its parameters fixed.
Left-pad this short prefix with 61 PAD IDs. After the history grows past 64 tokens, retain its last 64 IDs. There is no observed next-token target during free generation.
with torch.no_grad():
logits = model(X)
logits[:, [0, 1, 3]] = float("-inf") # exclude PAD, BOS, UNK
next_id = logits.argmax(dim=-1).item() # greedy choiceGreedy decoding takes the highest score; sampling draws from the probability distribution. Exclude PAD, BOS and UNK as outputs. EOS can stop the story.
Attention softmax chooses source positions. Vocabulary softmax supplies next-token probabilities. Cross-entropy uses raw logits during training; we do not first apply this generation softmax before that loss.
finished = next_id == real_stoi["<EOS>"]
if not finished:
history.append(next_id) # crop/pad history, then run the same model againKeep generating until EOS or a token limit. Training used observed previous words; generation now feeds back the model’s own choices.
“a” here is an illustrative choice, not a claimed checkpoint output. The saved continuations later show the actual outputs. An early mistake can alter every later input. The weights stay fixed throughout generation.
Multi-head attention · a worked continuation
What did the trained models achieve?
Same 6,000 stories, vocabulary, context and 6,000-update budget. Three seeds: 11, 29, 47. Select by validation loss; evaluate on test stories.
Both attention models add learned position rows, use the same total width and prediction MLP, and have equal parameter counts. The MLP has no separate position table: concatenation encodes slot order. Four heads reuse the single-head optimizer settings. This is not a position-encoding ablation or a parameter-matched MLP comparison.
Saved measurements from seed 11. Choose the best validation checkpoint—not simply the last update.
These are fixed selection-slice validation measurements, not test losses. The table on the next slide aggregates separately selected test checkpoints across three seeds. All traces, checkpoint choices and protocol.
Perplexity = exp(test loss) measures next-token surprise. If every true next token gets probability ¼, perplexity is 4. Lower is better.
Three training runs use different random starts (seeds 11, 29, 47). We report the average ± run-to-run spread (sample standard deviation).
Four-head perplexities: 28.40, 29.49, 28.46 give 28.78 ± 0.61.
For each run, test loss is the average negative log probability of the observed next token, using natural logarithms. Perplexity is exp(loss), or the reciprocal of the geometric mean of those probabilities. The ¼ example assumes that same probability for every observed token. Perplexity is not an accuracy percentage or a literal count of candidate words.
A seed controls random initialization and training-batch sampling. All three runs use the same data split and settings. For each metric separately, average its three run values and compute their sample standard deviation: square the deviations from their mean, sum them, divide by 3 − 1, then take the square root. The ± describes variation between training runs. It is not a confidence interval or variation between individual tokens.
The displayed run values are rounded. The summaries use the saved full-precision values. Compute exp(loss) for each run before averaging perplexities, rather than exponentiating the average loss. Four heads reduce mean perplexity by 8.2% versus one head in this experiment. Saved runs and summaries.
Mean training time, including validation, on Apple M2 Max/MPS. The attention models have equal parameter counts; four heads train slightly slower.
Mean wall-clock training time, including validation, on Apple M2 Max / MPS. Times are device-specific. For full-sequence attention at fixed width D, matching and mixing cost O(T²D) for either head count; storing all weights costs O(hT²). The optimized next-token path used here computes only the final query. Open measured size, loss and timing results · Optional complexity derivation.
Multi-head attention · a worked continuation
Finish by comparing their stories
Saved seed-11 checkpoints · greedy decoding · 24-token limit. Four heads keep “owl” in view here, but all three outputs still have flaws.
Unedited saved outputs for the app’s held-out preset. It was chosen as the shortest test document, not selected for an attractive output. Aggregate loss is stronger evidence than one sample. The app lets students inspect more prompts and generate fresh continuations.
Open the live three-model comparison ↗
Loss, perplexity, parameter counts and training time · Experiment notebook
Start with the held-out owl story and greedy decoding, then try an outside-story prompt. Compare coherence, repetition and topic consistency as well as local generation timings. Live inference runs on your device with WebGPU or WASM; it is not replaying the saved slide outputs. Optional worked arithmetic · Download code and notebooks. This closes the text-attention series. Vision is a separate continuation.