Before you start
The small example uses B=2 examples, w=4 context slots and C=10 vocabulary items. Every number is generated by the supplied code, not drawn by hand. The final measured comparison uses a separate 6,000-story corpus and trained models.
Unzip the download, open a terminal in that folder and run:
python -m venv .venv source .venv/bin/activate python -m pip install -r requirements.txt jupyter lab 05_training_and_inference_maps.ipynb
Then choose Run → Run All Cells. This notebook needs no GPU or data download. Notebooks 1 and 3 explain the real corpus and training runs. Keep wordlm.py and the other helpers beside the notebooks.
The training map: stories
We begin with complete stories. The highlighted box supplies the data for everything that follows.
The data: complete short stories
The saved experiment uses 6,000 TinyStories documents. Each document is a complete story. This is an excerpt from source row 12400.
Where are we on the full MLP map?
sample = evidence['sample']
print(sample['text_excerpt'])
print('Source row:', sample['row_idx'])Once upon a time, there lived a little bunny. Source row: 12400
Read one complete story
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.
Where are we on the full MLP map?
short_story = story_examples['examples'][0]
short_words = len(short_story['text'].split())
short_tokens = len(tokenize(short_story['text']))
assert (short_words, short_tokens) == (52, 60)
print(short_story['text'])
print(f'Full story: {short_words} words; {short_tokens} tokens.')Once upon a time, there was an old hotel. The hotel was very big and had many rooms. One day, a big storm came and lightning struck the hotel. The hotel was very old and could not handle the strike. The hotel fell down and everyone had to go home. The end. Full story: 52 words; 60 tokens.
Longer stories have several paragraphs
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.
Where are we on the full MLP map?
long_stories = story_examples['examples'][1:]
story_lengths = []
for story in long_stories:
counts = (len(story['text'].split()), len(tokenize(story['text'])))
assert counts == (story['word_count'], story['token_count'])
story_lengths.append(counts)
print(f"Source row {story['row_idx']}: {counts[0]} words; {counts[1]} tokens")
print(story['text'] + '\n') # full texts here; excerpts in the figure
train_lengths = evidence['audit']['story_length_tokens']['train']
print('Training-story token lengths:', train_lengths)Source row 12400: 107 words; 133 tokens
Once upon a time, there lived a little bunny. The bunny loved carrots, so every day the bunny went outside to find carrots. One day, the bunny was looking for some carrots and came across something strange. It looked like a weird rock.
The bunny hopped closer and asked, "What are you?"
The rock replied, "I'm a weird carrot!"
The bunny was very surprised and said, "Weird carrots don't exist!"
But the rock said, "Oh yes they do! I'm a cool carrot that always rocks!"
The bunny was amazed. He thanked the rock for the new and weird carrot and hopped away to find more delicious snacks.
Source row 588307: 248 words; 300 tokens
Sara and Tom were playing with their crayons and paper. They liked to draw animals and flowers and cars. Sara was drawing a big red pepper, because she liked to eat peppers. Tom was drawing a blue fish, because he liked to swim.
"Look at my pepper!" Sara said to Tom. "It is very pretty and shiny. Do you like it?"
Tom looked at Sara's pepper. He did not like it. He thought it was boring and ugly. He wanted to make Sara feel bad. He was very rude.
"I don't like your pepper," Tom said to Sara. "It is not pretty and shiny. It is silly and dumb. Watch this!"
Tom took his black crayon and made a big mark on Sara's pepper. He drew a line across it and made a face with a tongue sticking out. He laughed and said, "Now your pepper is funny and silly!"
Sara saw what Tom did to her pepper. She felt very sad and angry. She liked her pepper and she worked hard to draw it. She did not think it was funny or silly. She thought it was mean and nasty.
"Tom, you are very rude!" Sara said to Tom. "You ruined my pepper! You are not my friend! Go away!"
Sara took her paper and crayons and ran to the other side of the room. She did not want to play with Tom anymore. She wanted to find a new friend who would be nice and kind.
Training-story token lengths: {'max': 942, 'median': 176, 'min': 60}
The training map: splitting stories
Keep each story in one split before making overlapping windows. Only the training stories fit the vocabulary and model.
Stories stay together when we split the data
The split happens before overlapping windows are made. Training fits parameters and the vocabulary. Validation chooses settings and checkpoints. Test measures the frozen choice.
Where are we on the full MLP map?
audit = evidence['audit']
print(audit['documents'])
assert sum(audit['documents'].values()) == 6000
assert not any(audit['document_overlap'].values()){'test': 592, 'train': 4822, 'validation': 586}
A sentence we can trace completely
We use this authored sentence for the arithmetic. Its small vocabulary and randomly initialized models are separate from the measured TinyStories experiment.
Where are we on the full MLP map?
sentence = 'Lily found a red ball.'
print(sentence)Lily found a red ball.
Tokenization
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.
Where are we on the full MLP map?
tokenization_text = 'redder!'
print('The same text for three tokenization choices:', tokenization_text)The same text for three tokenization choices: redder!
Words, characters or pieces of words
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.
Where are we on the full MLP map?
tokenization_choices = {
'Word + punctuation': tokenize(tokenization_text),
'Character': list(tokenization_text),
'Subword (illustrative)': ['red', 'der', '!'],
}
tokenization_counts = {name: len(parts) for name, parts in tokenization_choices.items()}
assert list(tokenization_counts.values()) == [2, 7, 3]
for name, parts in tokenization_choices.items():
assert ''.join(parts) == tokenization_text
print(name, parts, 'tokens:', len(parts))
# The subword split is hand-chosen, not the output of a trained BPE tokenizer.Word + punctuation ['redder', '!'] tokens: 2 Character ['r', 'e', 'd', 'd', 'e', 'r', '!'] tokens: 7 Subword (illustrative) ['red', 'der', '!'] tokens: 3
The tokenizer used in this notebook
Words plus punctuation keep our calculations easy to inspect. Both 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.
Where are we on the full MLP map?
tokenizer_probe = "Lily can't find 12 balls!"
probe_tokens = tokenize(tokenizer_probe)
assert probe_tokens == ['lily', "can't", 'find', '12', 'balls', '!']
assert tokenize('LILY found!') == ['lily', 'found', '!']
print(tokenizer_probe)
print(probe_tokens)
# This English teaching tokenizer drops case and spacing.
# Its word pattern is ASCII-based; it is not a general multilingual tokenizer.Lily can't find 12 balls! ['lily', "can't", 'find', '12', 'balls', '!']
The training map: tokens and IDs
Text becomes tokens, then integer vocabulary IDs. Learned embedding vectors come later at token lookup.
Lowercase words and separate punctuation
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.
Where are we on the full MLP map?
pieces = tokenize(sentence)
print(pieces)
assert pieces == ['lily', 'found', 'a', 'red', 'ball', '.']
assert len(pieces) == 6['lily', 'found', 'a', 'red', 'ball', '.']
Every vocabulary item gets an integer ID
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.
Where are we on the full MLP map?
words = list(SPECIAL_TOKENS) + sorted(set(pieces))
vocab = Vocabulary(words, {t:i for i,t in enumerate(words)}, {}, 10, 1)
C = len(words)
assert C == 10
print(vocab.stoi){'<PAD>': 0, '<BOS>': 1, '<EOS>': 2, '<UNK>': 3, '.': 4, 'a': 5, 'ball': 6, 'found': 7, 'lily': 8, 'red': 9}
Four special tokens have different jobs
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.
Where are we on the full MLP map?
assert 'blue' not in vocab.stoi
token_ids = vocab.encode_tokens(['blue'], boundaries=False)
print(token_ids) # [3]: UNK
assert token_ids == [3] == [vocab.unk_id]
with_boundaries = vocab.encode_tokens(['blue'], boundaries=True)
print('With BOS and EOS:', with_boundaries) # [1, 3, 2]
assert with_boundaries == [vocab.bos_id, vocab.unk_id, vocab.eos_id][3] With BOS and EOS: [1, 3, 2]
The training map: story boundaries
Add BOS and EOS to each story’s ID sequence. These boundaries let us create the first input and the final stopping target.
One BOS and EOS per document
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.
Where are we on the full MLP map?
ids = vocab.encode_tokens(pieces, boundaries=True)
print(ids)
assert ids == [1, 8, 7, 5, 9, 6, 4, 2]
# Two sentences still make one document with one pair of markers.
two_sentence_story = 'Lily found a ball. Lily found a red ball.'
two_sentence_ids = vocab.encode_tokens(tokenize(two_sentence_story), boundaries=True)
print(' '.join(vocab.decode_ids(two_sentence_ids, skip_special=False)))
assert two_sentence_ids.count(vocab.bos_id) == 1
assert two_sentence_ids.count(vocab.eos_id) == 1
assert two_sentence_ids.count(vocab.stoi['.']) == 2[1, 8, 7, 5, 9, 6, 4, 2] <BOS> lily found a ball . lily found a red ball . <EOS>
The training map: context and target
A window holds the previous w token IDs. Its target is the next observed ID in the story.
Token positions and vocabulary IDs
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.
Where are we on the full MLP map?
w = 4
target_position = 4
print(list(enumerate(ids)))[(0, 1), (1, 8), (2, 7), (3, 5), (4, 9), (5, 6), (6, 4), (7, 2)]
One input window and its observed target
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.
Where are we on the full MLP map?
context_ids = ids[target_position-w:target_position]
target_id = ids[target_position]
assert context_ids == [1, 8, 7, 5] and target_id == 9
print('context_ids:', context_ids, 'target_id:', target_id)context_ids: [1, 8, 7, 5] target_id: 9
Lists for inputs and targets
We inspected the example with red as its target, but have not stored any pairs yet. To collect every example in order, start at position 1, where lily follows BOS.
Where are we on the full MLP map?
contexts = []
targets = []
assert len(contexts) == len(targets) == 0The first pair has only BOS as history
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.
Where are we on the full MLP map?
t = 1
visible = ids[max(0, t-w):t]
context_ids = [vocab.pad_id] * (w-len(visible)) + visible
target_id = ids[t]
assert context_ids == [0, 0, 0, 1] and target_id == 8Storing an input and its target
contexts.append(context_ids) adds the whole four-ID list as one row. targets.append(target_id) adds one answer at the same row index, so contexts[0] and targets[0] belong together.
Where are we on the full MLP map?
contexts.append(context_ids)
targets.append(target_id)
assert contexts == [[0, 0, 0, 1]] and targets == [8]
print('contexts =', contexts)
print('targets =', targets)contexts = [[0, 0, 0, 1]] targets = [8]
The next pair includes the previous observed target
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.
Where are we on the full MLP map?
t = 2
visible = ids[max(0, t-w):t]
context_ids = [vocab.pad_id] * (w-len(visible)) + visible
target_id = ids[t]
assert context_ids == [0, 0, 1, 8] and target_id == 7
assert contexts == [[0, 0, 0, 1]] and targets == [8]The second pair becomes row 1 in both lists
The same two append calls add the new pair after the first one. Each row in contexts still lines up with exactly one entry in targets.
Where are we on the full MLP map?
contexts.append(context_ids)
targets.append(target_id)
assert contexts == [[0, 0, 0, 1], [0, 0, 1, 8]]
assert targets == [8, 7]
print('contexts =', contexts)
print('targets =', targets)contexts = [[0, 0, 0, 1], [0, 0, 1, 8]] targets = [8, 7]
Building the remaining training pairs
Positions 1 and 2 are already stored, so this loop continues at 3 and stops before len(ids)=8. Each pass builds a fresh input list and appends it with the observed target, producing seven paired rows in total.
Where are we on the full MLP map?
for t in range(3, len(ids)):
visible = ids[max(0, t-w):t]
context_ids = [vocab.pad_id] * (w-len(visible)) + visible
target_id = ids[t]
contexts.append(context_ids)
targets.append(target_id)
assert len(contexts) == len(targets) == 7
assert contexts[3] == ids[0:4] and targets[3] == ids[4]Early prefixes need padding
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].
Where are we on the full MLP map?
for row in range(4):
print(row, contexts[row], targets[row])
assert contexts[3] == [1, 8, 7, 5] and targets[3] == 90 [0, 0, 0, 1] 8 1 [0, 0, 1, 8] 7 2 [0, 1, 8, 7] 5 3 [1, 8, 7, 5] 9
Later prefixes drop their oldest tokens
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.
Where are we on the full MLP map?
for row in range(4, 7):
print(row, contexts[row], targets[row])
assert targets[-1] == vocab.eos_id4 [8, 7, 5, 9] 6 5 [7, 5, 9, 6] 4 6 [5, 9, 6, 4] 2
The paired lists become two integer tensors
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.
Where are we on the full MLP map?
all_X = torch.tensor(contexts, dtype=torch.long)
all_y = torch.tensor(targets, dtype=torch.long)
N = len(all_y)
assert N == 7 and all_X.shape == (7, 4)
assert all_X.tolist() == contexts and all_y.tolist() == targets
assert not all_y.eq(vocab.pad_id).any()Six ordinary tokens give seven training examples
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.
Where are we on the full MLP map?
ordinary_tokens = len(pieces)
examples_in_story = ordinary_tokens + 1
assert ordinary_tokens == 6 and examples_in_story == N == 7Training examples across the corpus
The saved training split has 964,338 ordinary tokens across 4,822 complete stories. Adding one EOS target per story gives 969,160 training examples.
Where are we on the full MLP map?
train_tokens = 964_338
train_stories = 4_822
train_examples = train_tokens + train_stories
assert train_tokens == audit['oov']['train']['tokens']
assert train_stories == audit['documents']['train']
assert train_examples == 969160Example counts for each data split
For each split, add its ordinary-token count and its story count. These are counts of supervised examples, not optimizer steps.
Where are we on the full MLP map?
window_counts = {
split: audit['oov'][split]['tokens'] + count
for split, count in audit['documents'].items()
}
assert window_counts['train'] == 969160
assert window_counts['train'] == train_examples
print(window_counts){'test': 120393, 'train': 969160, 'validation': 120220}
Context length controls what is visible
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.
Where are we on the full MLP map?
history = ids[:6] # BOS lily found a red ball
contexts_by_width = {width: history[-width:] for width in [2, 4, 6]}
print(contexts_by_width){2: [9, 6], 4: [7, 5, 9, 6], 6: [1, 8, 7, 5, 9, 6]}
The training map: one batch
We select B=2 examples, each with w=4 input slots. They form X with shape [2,4] and y with shape [2].
A batch of two examples
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.
Where are we on the full MLP map?
selected = torch.tensor([2, 3])
X, y = all_X[selected], all_y[selected]
B = X.shape[0]
assert X.tolist() == [[0, 1, 8, 7], [1, 8, 7, 5]]
assert y.tolist() == [5, 9] and B == 2X and y contain integer token 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.
Where are we on the full MLP map?
print(X.shape, X.dtype)
print(y.shape, y.dtype)
assert X.shape == (2, 4) and y.shape == (2,)
assert X.dtype == y.dtype == torch.longtorch.Size([2, 4]) torch.int64 torch.Size([2]) torch.int64
Seven examples, four batches
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.
Where are we on the full MLP map?
batch_sizes = [len(all_y[start:start+B]) for start in range(0, N, B)]
assert batch_sizes == [2, 2, 2, 1]
print('One sequential pass:', batch_sizes)
print('Benchmark target presentations:', 6000 * 512)One sequential pass: [2, 2, 2, 1] Benchmark target presentations: 3072000
Batch and model dimensions
C includes the special tokens. Queries and keys use the same width dₖ, while values can use a different width dᵥ.
d, h, d_k, d_v = 4, 8, 3, 2
torch.manual_seed(11)
mlp = FixedWindowMLP(C, w, d, h, vocab.pad_id)
torch.manual_seed(11)
attention = CausalAttentionLM(C, w, d, d_k, d_v, h, vocab.pad_id)
print(dict(B=B, w=w, C=C, d=d, h=h, d_k=d_k, d_v=d_v)){'B': 2, 'w': 4, 'C': 10, 'd': 4, 'h': 8, 'd_k': 3, 'd_v': 2}
The full MLP training map
X has shape [2,4] and y has shape [2]. Input IDs enter token lookup, while the observed target goes directly to the loss. Each optimizer step updates all learned layers for the next batch.
Where are we on the full MLP map?
print('X:', tuple(X.shape), 'y:', tuple(y.shape))X: (2, 4) y: (2,)
An ID selects one embedding row
PAD fills empty slots: padding_idx=0 initializes its row to zero and skips its embedding gradient, so it stays zero here. These are numeric zeros, not nulls, while BOS/EOS/UNK have random, trainable rows that can learn when used as inputs. The table has C=10 rows and d=4 columns. ID 7, found, selects one four-number row. This is the initial table, before training, and the coordinates have no assigned word meanings. Only PAD has this zero-row rule in our model. EOS ends a story and is only a target in these training windows, so its input embedding receives no task gradient here, although the separate output layer learns to predict EOS. UNK can receive an embedding gradient when a missing word maps to it in an input. Zeroing PAD is a model choice, not a rule for all special tokens. In the attention model, we also mask padded source positions because a zero embedding alone does not remove them from softmax. Reference: https://docs.pytorch.org/docs/stable/generated/torch.nn.Embedding.html
Where are we on the full MLP map?
embedding_table = mlp.token_embedding.weight
assert mlp.token_embedding.padding_idx == vocab.pad_id == 0
assert embedding_table[vocab.pad_id].eq(0).all()
found_id = vocab.stoi['found']
found_vector = embedding_table[found_id]
assert found_vector.shape == (4,)Input IDs select rows from the embedding table
The learned table T has shape [C,d] = [10,4]. Each ID in X[1] selects one complete row, giving four vectors in the same slot order. The next slide stacks this result for both batch examples.
Where are we on the full MLP map?
selected_rows = embedding_table[X[1]]
assert selected_rows.shape == (w, d)
assert torch.equal(selected_rows[2], found_vector)The whole batch becomes a 2 × 4 × 4 tensor
Each of the two examples has four token slots. Each slot receives four embedding coordinates. PAD uses the fixed zero row. The same token ID selects the same row wherever it occurs.
Where are we on the full MLP map?
E_mlp = mlp.token_embedding(X)
assert E_mlp.shape == (2, 4, 4)
assert E_mlp[0, 0].eq(0).all()
assert torch.equal(E_mlp[0, 2], E_mlp[1, 1]) # lilyThe MLP map: joining the input vectors
Token lookup produced E with shape [B,w,d] = [2,4,4]. Each example now becomes one row of w×d=16 numbers.
Flatten the context axis, keep the batch axis
Each of the four tokens supplies four input nodes, giving 16 numbers in slot order. Both batch examples use this same 16 → 8 → 10 network separately. Flattening only rearranges the coordinates and learns no weights. We never concatenate one batch example with another. Swapping token positions changes which input connections receive each embedding.
Where are we on the full MLP map?
flat = E_mlp.flatten(start_dim=1)
assert flat.shape == (2, 16)
assert torch.equal(flat[1, :4], E_mlp[1, 0])ReLU keeps positive activations
At each of the same eight hidden neurons, ReLU keeps a positive sum and replaces a negative sum with zero. The labels show each value before and after ReLU, with batch shape still [2,8]. This is an activation within the hidden layer, not another set of eight learned neurons. The nonlinearity lets the network learn more than a single affine mapping.
Where are we on the full MLP map?
hidden_mlp = torch.relu(pre_hidden)
assert hidden_mlp.shape == (2, 8)
assert hidden_mlp.ge(0).all()The MLP map: vocabulary scores
The hidden row produces one score for each of C=10 vocabulary items. With two examples, the logits have shape [2,10].
Four input tokens, one next-token target
A logit is a raw next-token score, and the largest gives the model’s guess. Ground truth is the actual next token in the story, outside the four input slots. Both examples use the same MLP: four token IDs select four embeddings, which flatten to 16 numbers, pass through eight hidden activations, then produce ten vocabulary scores. For example 0, the context is <PAD> <BOS> lily found and the observed next token is a. For example 1, the context is <BOS> lily found a and the observed next token is red. These are initial, untrained scores. A correct guess here can happen by chance. Ground truth comes from the data, not from selecting the highest score. Softmax and the training loss follow later.
Where are we on the full MLP map?
logits_mlp = mlp.vocab_head(hidden_mlp)
assert logits_mlp.shape == (2, 10)
assert torch.allclose(logits_mlp, mlp(X))
predicted_ids = logits_mlp.argmax(dim=-1)
assert y.tolist() == [vocab.stoi['a'], vocab.stoi['red']]The full attention training map
The data, targets and loss stay the same. Attention changes how context reaches the hidden prediction layer. This one-layer implementation computes only the final query for each window.
assert X.shape == (2, 4) and y.shape == (2,)
print('Both architectures predict the same two targets:', y.tolist())Both architectures predict the same two targets: [5, 9]
The attention map: input vectors
Attention adds a position vector to each token vector. E still has shape [2,4,4] before the query, key and value projections.
Add a position vector to each token vector
The position table has four rows, one for each context slot. Token and position vectors both have width four, so addition preserves the width. PAD keys will be masked even though position addition can make their input row nonzero.
Where are we on the full ATTENTION map?
positions = torch.arange(w)
token_rows = attention.token_embedding(X)
position_rows = attention.position_embedding(positions)
E = token_rows + position_rows[None, :, :]
assert E.shape == (2, 4, 4)The attention map: queries, keys and values
Each window supplies one final query and four source keys and values. The next three calculations use the same input tensor E.
The final known token supplies one query
For example 1 the final known token is a. Its four-number input row maps to three query coordinates. The unseen target red does not take part in this multiplication.
Where are we on the full ATTENTION map?
q = attention.W_Q(E[:, -1:, :])
query_terms = E[1, -1] * attention.W_Q.weight[0]
assert q.shape == (2, 1, 3)
assert torch.allclose(query_terms.sum(), q[1, 0, 0])Each source has a matching key
All four source rows map to keys of width three. Query and key widths match because we will take their dot products. The same W_K applies at every source position and in both examples.
Where are we on the full ATTENTION map?
K = attention.W_K(E)
assert K.shape == (2, 4, 3)
print('Row-vector W_K shape:', tuple(attention.W_K.weight.T.shape))Row-vector W_K shape: (4, 3)
Each source also has information to send
W_V maps each source to a two-number value. These coordinates carry information into the weighted message. They do not enter the query-key dot product.
Where are we on the full ATTENTION map?
V = attention.W_V(E)
assert V.shape == (2, 4, 2)
print('Row-vector W_V shape:', tuple(attention.W_V.weight.T.shape))Row-vector W_V shape: (4, 2)
The attention map: source weights
Query-key scores determine where to read. Scaling, padding masks and softmax produce one row of four weights per example.
Compare the query with all four keys
Each score is a three-term dot product divided by √3. The result has one row of four source scores per example. These are source-match scores, not vocabulary logits.
Where are we on the full ATTENTION map?
raw_scores = q @ K.transpose(-2, -1)
scores = raw_scores / math.sqrt(d_k)
dot_terms = q[1, 0] * K[1, 2]
assert torch.allclose(dot_terms.sum(), raw_scores[1, 0, 2])
assert scores.shape == (2, 1, 4)Padding must receive zero attention weight
Example 0 has PAD in its first slot, so that score becomes −∞. All slots of example 1 are real tokens. Since these are final queries, their windows contain no future columns.
Where are we on the full ATTENTION map?
pad_mask = X[:, None, :].eq(vocab.pad_id)
masked_scores = scores.masked_fill(pad_mask, float('-inf'))
assert torch.isneginf(masked_scores[0, 0, 0])Softmax turns four source scores into weights
Subtract the largest score, exponentiate and divide by the row sum. This softmax runs over context slots. Later, a different softmax will run over vocabulary items.
Where are we on the full ATTENTION map?
source_exp = (masked_scores - masked_scores.amax(-1, keepdim=True)).exp()
A = source_exp / source_exp.sum(-1, keepdim=True)
assert torch.allclose(A, masked_scores.softmax(-1))
assert A[0, 0, 0] == 0 and torch.allclose(A.sum(-1), torch.ones(2, 1))The attention map: the message
Each source weight multiplies a two-coordinate value row. Their sum gives one message of width dᵥ=2 per example.
Each weight scales a complete value row
For example 1, multiply each two-number value by its source weight and add the four contributions. The result is one two-number message. Weighting does not change the value width.
Where are we on the full ATTENTION map?
contributions = A[1, 0, :, None] * V[1]
message = A @ V
assert message.shape == (2, 1, 2)
assert torch.allclose(contributions.sum(0), message[1, 0])The attention map: the output projection
Wₒ maps the two-coordinate message into four coordinates. The residual addition then combines it with the original final input row.
Project two message coordinates into four
The message has width two, but the original input has width four. The learned W_O is 2×4 in row-vector notation. Every output coordinate can combine both message coordinates.
Where are we on the full ATTENTION map?
update = attention.W_O(message).squeeze(1)
assert update.shape == (2, 4)
assert torch.allclose(message[1, 0] @ attention.W_O.weight.T, update[1])Add the update to the original final row
The residual addition now combines two vectors with the same width. The updated row enters an eight-unit ReLU layer and then the ten-class vocabulary head, just as in the MLP path.
Where are we on the full ATTENTION map?
final = E[:, -1, :] + update
hidden_att = torch.relu(attention.hidden_layer(final))
logits_att = attention.vocab_head(hidden_att)
assert torch.allclose(logits_att, attention(X), atol=1e-6)
assert torch.allclose(logits_att, attention.forward_details(X)['logits'], atol=1e-6)The training map: predictions and targets
Return to the MLP’s two rows of vocabulary logits. The observed IDs y=[5,9], meaning a and red, meet those logits at the loss.
Ten logits become ten next-token probabilities
For the MLP example with target red, apply vocabulary softmax. The denominator includes all ten classes. These values come from the seeded model before training. Even a correct guess here can occur by chance.
Where are we on the full MLP map?
word_exp = (logits_mlp[1] - logits_mlp[1].max()).exp()
p = word_exp / word_exp.sum()
assert torch.allclose(p, logits_mlp[1].softmax(-1))
assert torch.allclose(p.sum(), torch.tensor(1.0))Observed target and model prediction
Argmax returns the largest-probability vocabulary ID. The target remains red because the corpus contains red after this prefix. Training does not replace that observed target with the model’s guess.
Where are we on the full MLP map?
guess_id = int(p.argmax())
target_id = int(y[1])
print('Guess:', words[guess_id], 'target:', words[target_id])
print('Probability of target:', float(p[target_id].detach()))Guess: red target: red Probability of target: 0.1444295048713684
Per-example loss and batch mean
For each example, cross-entropy is −log of the probability assigned to its observed target. PyTorch accepts raw logits and performs log-softmax internally. The optimizer uses the mean across B=2 examples.
Where are we on the full MLP map?
losses = F.cross_entropy(logits_mlp, y, reduction='none')
loss = losses.mean()
manual_loss = -logits_mlp.log_softmax(-1)[torch.arange(B), y]
assert torch.allclose(losses, manual_loss)The training map: backpropagation
The batch mean loss supplies gradients for the learned parameters. Backpropagation computes these gradients before any weight changes.
Gradients from backpropagation
For vocabulary weight W[j,r], the batch-mean gradient is the mean of (p_r − 1[y=r]) × hidden_j. Here we inspect the weight feeding the red logit from hidden unit 0.
Where are we on the full MLP map?
mlp.zero_grad(set_to_none=True)
loss.backward()
manual_grad = ((logits_mlp.softmax(-1)[:, 9] - y.eq(9).float()) * hidden_mlp[:, 0]).mean()
grad = mlp.vocab_head.weight.grad[9, 0]
assert torch.allclose(grad, manual_grad)The training map: updating parameters
The optimizer uses the gradients to change the learned parameters. The next batch repeats the forward pass with those updated weights.
One optimizer step changes the stored weights
For this arithmetic demonstration we use SGD with learning rate 0.1: new weight = old weight − 0.1×gradient. The measured TinyStories runs use validation-tuned AdamW instead.
Where are we on the full MLP map?
old_weight = mlp.vocab_head.weight[9, 0].detach().clone()
before_loss = float(loss.detach())
optimizer = torch.optim.SGD(mlp.parameters(), lr=0.1)
optimizer.step()
new_weight = mlp.vocab_head.weight[9, 0].detach().clone()
after_loss = float(F.cross_entropy(mlp(X), y).detach())
assert torch.allclose(new_weight, old_weight - 0.1*grad)Many batches, with validation between checkpoints
The actual trainer samples 512 training windows per step, computes the loss, backpropagates and updates parameters. Validation selects the best saved checkpoint. The test set does not choose a checkpoint.
Where are we on the full ATTENTION map?
benchmark = evidence['benchmark']
print('Maximum steps:', 6000, 'batch size:', 512)
print('Maximum target presentations:', 6000*512)
print('Checkpoint rule: best validation loss')Maximum steps: 6000 batch size: 512 Maximum target presentations: 3072000 Checkpoint rule: best validation loss
Held-out scoring and free-running generation
Held-out scoring uses observed X,y pairs and records loss. Generation has a prompt but no observed next target. Both freeze parameters. Neither calls backward or optimizer.step.
Where are we on the full MLP map?
mlp.eval()
with torch.inference_mode():
frozen_loss = F.cross_entropy(mlp(X), y)
print('Illustrative scoring loss:', float(frozen_loss))Illustrative scoring loss: 2.2394907474517822
The complete MLP generation loop
Load the same vocabulary and trained parameters. Prepare the latest window, predict, choose a token and append it. Stop on EOS or the generation budget. The only changing state is the history.
print('Generation uses B=1; the parameters stay fixed.')Generation uses B=1; the parameters stay fixed.
Attention uses the same generation loop
The new final token supplies the next query. This notebook recomputes the current window and has no KV cache. Position indices describe the slots in that current window.
attention.eval()
print('Final query only; all available keys and values; no KV cache.')Final query only; all available keys and values; no KV cache.
The generation map: preparing the prompt
Generation uses the saved tokenizer and the latest w IDs. Our worked example follows the frozen MLP with B=1 and no observed next target.
Prepare a prompt using the same rules
Use the training tokenizer and vocabulary on the unfinished prompt. Prepend BOS once, and leave EOS out because generation has not finished.
Where are we on the full MLP map?
prompt = 'Lily found'
prompt_ids = vocab.encode_tokens(tokenize(prompt), boundaries=False)
history = [vocab.bos_id] + prompt_ids
assert history == [1, 8, 7]The prompt fills the same four input slots
Keep the most recent w=4 IDs from the history. This prompt has only three IDs including BOS, so one PAD ID fills the unused slot on the left.
Where are we on the full MLP map?
kept = history[-w:]
context_ids = [vocab.pad_id] * (w-len(kept)) + kept
inference_X = torch.tensor([context_ids])
assert inference_X.tolist() == [[0, 1, 8, 7]]A frozen model scores the next token
The MLP reads the prepared window and returns ten vocabulary scores. inference_mode disables gradient tracking for this forward pass, and no optimizer updates the parameters.
Where are we on the full MLP map?
with torch.inference_mode():
next_logits = mlp(inference_X)[0]
assert next_logits.shape == (C,)The generation map: choosing a token
Vocabulary softmax turns the scores into probabilities. Greedy decoding or sampling chooses the next token ID.
Special input tokens are excluded from generation
Set the scores of PAD, BOS and UNK to negative infinity before softmax. Their probabilities become zero, while EOS remains available as a stopping token.
Where are we on the full MLP map?
blocked_ids = [vocab.pad_id, vocab.bos_id, vocab.unk_id]
with torch.inference_mode():
next_logits[blocked_ids] = float('-inf')
next_p = next_logits.softmax(-1)
assert next_p[blocked_ids].eq(0).all()
assert torch.allclose(next_p.sum(), torch.tensor(1.0))Greedy decoding and sampling make different choices
Greedy chooses the ID with the highest probability, while sampling uses the cumulative probabilities (CDF). The fixed draw 0.8 selects the first CDF entry at least 0.8, making this example reproducible.
Where are we on the full MLP map?
greedy_id = int(next_p.argmax())
cdf = next_p.cumsum(0)
sample_id = int(torch.searchsorted(cdf, torch.tensor(0.8)))The generation map: extending the history
EOS ends generation. Otherwise append the chosen ID and prepare the newest window for another forward pass.
A selected token extends the history unless it is EOS
Use the greedy ID for this demonstration, and check whether it is EOS before changing the history. If it is EOS, stop generation, otherwise append it and prepare the next window.
Where are we on the full MLP map?
history_before_choice = history.copy()
chosen = greedy_id
if chosen != vocab.eos_id:
history.append(chosen)
assert history_before_choice == [1, 8, 7]The chosen token becomes input on the next step
Repeat window preparation, scoring and token selection, stopping on EOS or after four steps. This table repeats generation from the original prompt, with all parameters frozen.
Where are we on the full MLP map?
history = history_before_choice.copy()
frozen = {n:p.detach().clone() for n,p in mlp.named_parameters()}
generation_trace = []
with torch.inference_mode():
for generation_step in range(4):
kept = history[-w:]
x_next = torch.tensor([[vocab.pad_id]*(w-len(kept)) + kept])
z_next = mlp(x_next)[0]
z_next[[0, 1, 3]] = float('-inf')
chosen = int(z_next.argmax())
generation_trace.append((x_next[0].tolist(), chosen))
if chosen == vocab.eos_id: break
history.append(chosen)
assert all(torch.equal(p, frozen[n]) for n,p in mlp.named_parameters())The trained TinyStories comparison
The saved experiment uses the same documents, vocabulary, targets and 64-token context for both models. Each model has validation-tuned optimizer settings. These are three-seed test results, separate from our ten-token arithmetic example.
comparison = benchmark['aggregate']
for name in ['mlp', 'attention']:
stats = comparison[name]['test_perplexity']
print(name, stats['mean'], '+/-', stats['sample_std'])mlp 51.414434800867234 +/- 0.44307967119606395 attention 31.31463254211675 +/- 0.3047408118077636
The same steps in the runnable notebooks
This notebook executes the small calculation behind every figure. Notebook 1 loads the corpus and trains the MLP. Notebook 3 trains attention. Notebook 4 checks the saved comparison and studies representations.
assert B == 2 and w == 4 and C == 10
print('Toy example: 7 targets; worked batch: 2 examples.')
print('Real experiment: 969,160 training windows; context 64; vocabulary 4,000.')Toy example: 7 targets; worked batch: 2 examples. Real experiment: 969,160 training windows; context 64; vocabulary 4,000.