MiniGPT, rewritten to be read

Neil Haddley • October 9, 2026

An experiment: the code that runs my trained MiniGPT model, rewritten for people, with a name for every value, every shape written down, every library call explained, and honest notes where nobody knows why it works

AIPythongpttransformerspytorchreadable-codemachine-learning

Writing MiniGPT (Part 1), I spent a long time explaining lines like these, from Jibin Joseph's notebook:

PYTHON
1x = x + self.attn(self.ln1(x))
2q = q.view(B, T, self.n_head, self.head_dim).transpose(1, 2)

They are typical PyTorch: short, close to the maths in the papers, and quick on a GPU, the graphics chip that does most of this arithmetic. But they are written for readers who already know that maths. x is a different thing after every line, and ln1, attn, view and transpose say nothing about what they do. So I tried an experiment. I rewrote the part of the code that runs my trained model, and only that part, to be as easy for a person to read as I could make it, and then checked that it still gives exactly the same answers.

The program is readable/readable_minigpt.py in the series repo, with a notebook version saved with its outputs: open it in Colab, or read it on GitHub. It runs in about a second on my Mac.

The rules I followed
  1. Every value gets its own name. There is no x = x + something. When a step makes something new, the new thing gets a new name that says what it is.
  2. Every shape is written down. A comment such as # [n, 128] means "one row for each of the n letters of the text, 128 numbers in each row".
  3. Every library call is explained the first time it appears: what it does to the numbers.
  4. No hype. Where nobody really knows why a piece works, I say so.

Here is what that does to the lines above:

The originalRewritten
x = x + self.attn(self.ln1(x))attention_contribution = attention(hidden_states_entering, block) then hidden_states_after_attention = hidden_states_entering + attention_contribution
q.view(B, T, self.n_head, self.head_dim).transpose(1, 2)head_looking_queries = looking_queries[:, first_column:after_last_column], inside a loop over the four head slices
self.ln1normalise(hidden_states, block.stretch_before_attention, block.shift_before_attention)
self.lm_head(x)model.next_letter_rows @ final_hidden_state + model.next_letter_biases
Words I use

Here is what each technical word in this post means, so that none of them is used before it is explained. I kept the real names, rather than inventing friendlier ones, because they are the names you will meet in the code and in other explanations.

WordWhat it means here
vectora list of numbers; here, usually 128 of them
table, shapenumbers in rows and columns; a table's shape is its size, such as 65 × 128. PyTorch calls a table of numbers a tensor
weights, biasestrained numbers. A weight multiplies an input number; a bias is one extra number added at the end
embeddingthe trained vector that stands for something: a letter (its token embedding) or a position in the text (its position embedding)
hidden statea letter's current vector as it moves through the model. It starts as the letter's two embeddings added together, and each block adds to it. "Hidden" only means that nobody outside the model sees it
attentionthe step in which each letter's hidden state takes in information from the letters before it
head (here, head slice)one of four copies of attention that run side by side, each using its own 32 of the 128 numbers: a slice of the columns, which is why I call it a head slice
query, key, value (here, looking query, looked-at key, passed-on value)the three vectors attention makes from each hidden state. The usual names are borrowed from looking things up in a database; I add a word to each that says only what it does in the arithmetic, which section 5 explains
MLPshort for multilayer perceptron: a small network inside each block that works on each letter on its own
blockone round of attention followed by the MLP. The model runs four blocks, one after another, each with its own trained numbers
normalisingrescaling a vector's numbers to a standard size
logita score for one letter that could come next, before it becomes a chance
softmaxthe step that turns scores into chances that add up to 100%
batchseveral texts handled at once
GPUa graphics chip, which does this kind of arithmetic much faster than an ordinary processor
The whole program in one picture

Before the details, here is everything the program does to turn goo into its next letter, using the names in the code. Blue steps use the model's fixed trained numbers; orange steps use settings you choose.

Each box is one line of choose_next_letter, with the shape of its data for gooTap the image to open it full size

Each box is one line of choose_next_letter, with the shape of its data for goo

Step 3 is where most of the work happens, so here it is opened up. run_all_blocks runs run_one_block four times, each time with a different block's trained numbers, and keeps every result in a list:

Box 3 of the picture above, opened up: four blocks, each one attention and an MLP, each adding to the hidden statesTap the image to open it full size

Box 3 of the picture above, opened up: four blocks, each one attention and an MLP, each adding to the hidden states

And here is the code that draws it, line for line: run_all_blocks is the loop down the page, and run_one_block is the inside of each block box.

run_all_blocks and run_one_block, with each note pointing to its part of the picture aboveScroll sideways, or tap the diagram to open it full size

run_all_blocks and run_one_block, with each note pointing to its part of the picture above

One level further in: here is a single block box opened up. Inside run_one_block, attention and the MLP each work out something to add, through several steps of their own, while the hidden states go around them and are only changed at the two + steps. Sections 5 to 7 below explain each of these steps in turn.

One block box, opened up: every step inside attention and the MLP, with the shape of its data for gooTap the image to open it full size

One block box, opened up: every step inside attention and the MLP, with the shape of its data for goo

The top level: one function, eight steps

The picture is not a loose summary: it is this function, line for line. choose_next_letter is the whole model on one screen. Each line is one box of the picture, marked # step 1 to # step 8, and every other function in the program is one of the pieces it calls.

PYTHON
1def choose_next_letter(text: str, model: TrainedModel, temperature: float, keep_biggest: int,
2                       random_generator: torch.Generator) -> str:
3    """The whole model, top to bottom: one line for each step in the post's whole-program picture."""
4    letter_ids = letters_to_ids(text)[-MOST_LETTERS_THE_MODEL_CAN_SEE:]                                 # step 1
5    starting_hidden_states = first_hidden_states(letter_ids, model)                                     # step 2
6    hidden_states_after_each_block = run_all_blocks(starting_hidden_states, model)                     # step 3
7    last_letters_hidden_state = hidden_states_after_each_block[-1][-1]                                 # step 4
8    final_hidden_state = normalise(last_letters_hidden_state, model.final_stretch, model.final_shift)  # step 5
9    scores = model.next_letter_rows @ final_hidden_state + model.next_letter_biases                    # step 6
10    chances = chances_from_scores(scores, temperature, keep_biggest)                                   # step 7
11    next_letter_id = spin_the_wheel(chances, random_generator)                                         # step 8
12    return VOCABULARY[next_letter_id]
The same eight lines, with a note under each oneScroll sideways, or tap the diagram to open it full size

The same eight lines, with a note under each one

The sections below open the pieces in the order that makes them easiest to follow: first the numbers the model knows, then the small tools, and then each step from the letters to the wheel.

1. The static data: every number the model knows

Training produced 826,433 numbers, and they are all the model knows. They live in one file, exhibit.pt, as named tables. The original names are short, so the program gives each table a descriptive one:

Original nameName in the rewriteShapeNumbersWhat it is
token_embedding.weighttoken_embedding_table65 × 1288,320one row for each letter: its token embedding
position_embedding.weightposition_embedding_table128 × 12816,384one row for each position: its position embedding
blocks.N.ln1.weight, .biasstretch_before_attention, shift_before_attention128 each256normalising before attention
blocks.N.attn.query.weight, .biaslooking_query_weights, looking_query_biases128 × 128, 12816,512make each position's looking query
blocks.N.attn.key.weight, .biaslooked_at_key_weights, looked_at_key_biases128 × 128, 12816,512make each position's looked-at key
blocks.N.attn.value.weight, .biaspassed_on_value_weights, passed_on_value_biases128 × 128, 12816,512make each position's passed-on value
blocks.N.attn.proj.weight, .biashead_mixing_weights, head_mixing_biases128 × 128, 12816,512mix the four head slices' findings
blocks.N.ln2.weight, .biasstretch_before_mlp, shift_before_mlp128 each256normalising before the MLP
blocks.N.mlp.fc1.weight, .biaswiden_weights, widen_biases512 × 128, 51266,048widen 128 numbers to 512
blocks.N.mlp.fc2.weight, .biasnarrow_weights, narrow_biases128 × 512, 12865,664narrow 512 numbers back to 128
final_ln.weight, .biasfinal_stretch, final_shift128 each256the final normalisation
lm_head.weight, .biasnext_letter_rows, next_letter_biases65 × 128, 658,385score each letter that could come next

The rows marked blocks.N appear four times, once for each block, with the same shapes but different numbers: 198,272 numbers per block. Four blocks, plus the embeddings, the final normalisation and lm_head, make 826,433.

The file also holds four tables called causal_mask: 128 × 128 True or False values meaning "position i may look at position j". They are not learned, just a rule written out as a table, so the rewrite builds the same rule itself when it needs it. ("Causal" because each letter may only be affected by the letters before it, never by the ones after.)

Nobody wrote any of these numbers. Training started them as small random numbers and nudged them, again and again, towards whatever made the model better at guessing the next letter. None of the 128 numbers in a row has a meaning that anyone assigned to it.

In the program, each block's tables become one record with named fields. @dataclass is Python's way of saying "a record with named fields", and torch.Tensor means "a table of numbers":

PYTHON
1@dataclass
2class TrainedBlock:
3    """The trained numbers of one block. Every block has the same shapes, but its own values."""
4
5    # Normalising before attention: one "stretch" and one "shift" number for each of the 128 positions
6    # in a vector. (The original code calls these ln1.weight and ln1.bias.)
7    stretch_before_attention: torch.Tensor      # [128]
8    shift_before_attention: torch.Tensor        # [128]
9
10    # Attention makes three new vectors from each hidden state: a looking query, a looked-at key, and a
11    # passed-on value (usually just called the query, the key, and the value; section 6 explains the
12    # extra words). Each comes from its own table of weights: 128 rows of 128 numbers, plus 128 biases.
13    # (attn.query, attn.key, attn.value)
14    looking_query_weights: torch.Tensor                 # [128, 128]
15    looking_query_biases: torch.Tensor                  # [128]
16    looked_at_key_weights: torch.Tensor                   # [128, 128]
17    looked_at_key_biases: torch.Tensor                    # [128]
18    passed_on_value_weights: torch.Tensor                 # [128, 128]
19    passed_on_value_biases: torch.Tensor                  # [128]
20
21    # After the four head slices have each collected 32 numbers, this table mixes the 128 numbers they
22    # found into the 128 numbers that are added to the hidden state. (attn.proj)
23    head_mixing_weights: torch.Tensor           # [128, 128]
24    head_mixing_biases: torch.Tensor            # [128]
25
26    # Normalising before the MLP. (ln2.weight, ln2.bias)
27    stretch_before_mlp: torch.Tensor            # [128]
28    shift_before_mlp: torch.Tensor              # [128]
29
30    # The MLP: widen 128 numbers to 512, bend them, narrow them back to 128. (mlp.fc1, mlp.fc2)
31    widen_weights: torch.Tensor                 # [512, 128]
32    widen_biases: torch.Tensor                  # [512]
33    narrow_weights: torch.Tensor                # [128, 512]
34    narrow_biases: torch.Tensor                 # [128]
2. The changing state

Everything above is fixed. What changes is worked out fresh for each text, and thrown away afterwards. For a text of n letters, the model's whole state is:

Name in the rewriteShapeWhat it is
letter_idsneach letter's ID: its place in the 65-letter vocabulary
hidden_statesn × 128one vector per letter: it starts as token embedding + position embedding, and every block adds to it
looking_queries, looked_at_keys, passed_on_valuesn × 128 eachmade inside attention from the hidden states; each head slice uses 32 of the 128 numbers
attention_sharesn × n, for each head slicerow i: how much attention position i gives to each position up to itself
widenedn × 512inside the MLP, between widening and narrowing
final_hidden_state128the last letter's hidden state, after all four blocks and the final normalisation
scores65one score, a logit, for each letter that could come next
chances65the scores after temperature, top-k and softmax: the wheel

Nothing else is remembered. The model has no memory from one letter to the next except the text itself: to write the next letter, it starts again from the letters.

3. Four small tools

The whole model is built from four simple operations, used again and again.

Applying a table of weights. PyTorch calls this a "linear layer". For each input vector, every output number is a dot product: multiply the input's numbers by one row of weights, number by number, add up the results, and add that row's bias.

PYTHON
1def apply_weights(input_vectors: torch.Tensor, weights: torch.Tensor, biases: torch.Tensor) -> torch.Tensor:
2    """input_vectors: [rows, inputs]; weights: [outputs, inputs]; biases: [outputs] -> [rows, outputs]"""
3    dot_products = input_vectors @ weights.T     # [rows, outputs]: each input row · each weight row
4    return dot_products + biases                 # the same biases are added to every row

weights.T is the table turned on its side, so rows become columns. The @ sign is matrix multiplication, which here does every one of those dot products at once: every input row against every weight row.

Normalising. This puts each vector's 128 numbers on a standard scale, an average of 0 and a spread of 1, and then stretches and shifts each number by its own trained amount.

PYTHON
1def normalise(vectors: torch.Tensor, stretch: torch.Tensor, shift: torch.Tensor) -> torch.Tensor:
2    """vectors: [rows, 128] -> [rows, 128], each row rescaled to average 0 and spread 1, then stretched and shifted."""
3    average_of_each_row = vectors.mean(dim=-1, keepdim=True)                       # [rows, 1]
4    distance_from_average = vectors - average_of_each_row                          # [rows, 128]
5    average_squared_distance = (distance_from_average ** 2).mean(dim=-1, keepdim=True)   # [rows, 1]
6    spread_of_each_row = torch.sqrt(average_squared_distance + TINY_NUMBER_TO_AVOID_DIVIDING_BY_ZERO)
7    standardised = distance_from_average / spread_of_each_row                      # [rows, 128]
8    return standardised * stretch + shift                                          # [rows, 128]
normalise again, with a note under each lineScroll sideways, or tap the diagram to open it full size

normalise again, with a note under each line

.mean(dim=-1, keepdim=True) averages along the last direction of the table, across the 128 numbers of each row, and keeps the answer as a one-number column, so that it can be subtracted from every number in its row. torch.sqrt is the square root.

Why does normalising help? It keeps the numbers from growing or shrinking out of control as they pass through many steps, which makes training more stable. That much is well established, by the paper that introduced it and by years of use since. Exactly why it helps is still argued about. A 2019 study, Understanding and Improving Layer Normalization, opens by saying "it is still unclear where the effectiveness stems from", and found that "the derivatives of the mean and variance are more important than forward normalization": the way normalising changes training mattered more than the rescaling itself.

The bend (GELU). Two tables of weights in a row, with nothing between them, are no more powerful than one: a weighted mix of weighted mixes is still a weighted mix. A bend in between fixes that.

PYTHON
1def bend(numbers: torch.Tensor) -> torch.Tensor:
2    """GELU, applied to every number separately. Same shape in, same shape out."""
3    return 0.5 * numbers * (1.0 + torch.erf(numbers / math.sqrt(2.0)))

torch.erf is the "error function", a standard S-shaped curve from statistics, and this is the exact formula PyTorch uses. It lets big positive numbers through almost unchanged and turns big negative ones into almost 0: bend(3) is 2.996 and bend(-3) is −0.004. Why this curve, rather than a simpler one? Mostly because it worked well in experiments. The paper that introduced it justifies it with "an empirical evaluation" against older curves, across vision, language, and speech tasks. A later paper that compared several alternatives, GLU Variants Improve Transformer, ends with unusual candour:

We offer no explanation as to why these architectures seem to work; we attribute their success, as all else, to divine benevolence.

Turning scores into shares (softmax). Used twice: to share out attention, and at the very end, to turn the 65 scores into the wheel.

PYTHON
1def scores_to_shares(scores: torch.Tensor) -> torch.Tensor:
2    """Softmax along the last direction: each row of scores becomes shares that add up to 1."""
3    biggest_score_in_each_row = scores.amax(dim=-1, keepdim=True)
4    scores_minus_biggest = scores - biggest_score_in_each_row      # the biggest becomes 0
5    positive_numbers = torch.exp(scores_minus_biggest)              # all between 0 and 1
6    total_of_each_row = positive_numbers.sum(dim=-1, keepdim=True)
7    return positive_numbers / total_of_each_row

.amax finds the biggest number in each row, torch.exp raises e (about 2.718) to the power of each number, and .sum adds each row up. A score of minus infinity becomes exactly 0. The scores 2, 1 and 0 become shares of 66.5%, 24.5% and 9.0%.

4. From letters to the first hidden states
PYTHON
1def letters_to_ids(text: str) -> list:
2    """Each letter's place in the vocabulary. Letters that are not in it are skipped."""
3    return [LETTER_TO_ID[letter] for letter in text if letter in LETTER_TO_ID]
PYTHON
1def first_hidden_states(letter_ids: list, model: TrainedModel) -> torch.Tensor:
2    """[n] IDs -> [n, 128]: token embedding + position embedding, one row per letter."""
3    number_of_letters = len(letter_ids)
4    token_embeddings = model.token_embedding_table[letter_ids]                   # [n, 128]
5    position_embeddings = model.position_embedding_table[:number_of_letters]     # [n, 128]
6    return token_embeddings + position_embeddings                                # [n, 128]

table[letter_ids] picks one row for each ID in the list, in order, and table[:n] keeps the first n rows. For goo, the IDs are 45, 53 and 53. The two os pick the same token embedding, but their position embeddings differ, so they start differently: position 2 begins −0.006, −0.031, 0.017, and position 3 begins 0.019, −0.080, 0.013.

5. Attention: each letter looks back

This is the only place in the model where one letter's hidden state is affected by another's.

Attention makes three vectors from each hidden state, usually called the query, the key, and the value. Those names suggest a meaning that nobody has shown the numbers have, so in the code I add a word to each that says only what it does in the arithmetic, and keep the usual word so that you can still match it to papers and other code:

Usual nameIn the codeWhat it does in the arithmetic
querylooking_queriesused when this position looks back at the others
keylooked_at_keysused when another position looks at this one
valuepassed_on_valueswhat this position passes on to whoever looks at it
headhead_sliceone of four copies of attention, each working on its own 32 of the 128 numbers: a slice of the columns
PYTHON
1def attention(hidden_states: torch.Tensor, block: TrainedBlock) -> torch.Tensor:
2    """[n, 128] -> [n, 128]: what attention adds to each hidden state."""
3    number_of_positions = hidden_states.shape[0]
4    normalised = normalise(hidden_states, block.stretch_before_attention, block.shift_before_attention)
5
6    looking_queries = apply_weights(normalised, block.looking_query_weights, block.looking_query_biases)      # [n, 128]
7    looked_at_keys = apply_weights(normalised, block.looked_at_key_weights, block.looked_at_key_biases)      # [n, 128]
8    passed_on_values = apply_weights(normalised, block.passed_on_value_weights, block.passed_on_value_biases)  # [n, 128]
9
10    # may_look_at[i][j] is True when position i may look at position j: only j <= i.
11    may_look_at = torch.tril(torch.ones(number_of_positions, number_of_positions, dtype=torch.bool))
12
13    findings_of_each_head_slice = []
14    for head_slice in range(HEADS_PER_BLOCK):
15        first_column = head_slice * NUMBERS_PER_HEAD
16        after_last_column = first_column + NUMBERS_PER_HEAD
17        head_looking_queries = looking_queries[:, first_column:after_last_column]    # [n, 32]
18        head_looked_at_keys = looked_at_keys[:, first_column:after_last_column]      # [n, 32]
19        head_passed_on_values = passed_on_values[:, first_column:after_last_column]  # [n, 32]
20
21        # match_scores[i][j]: position i's looking query · position j's looked-at key
22        match_scores = head_looking_queries @ head_looked_at_keys.T          # [n, n]
23        shrunk_scores = match_scores / math.sqrt(NUMBERS_PER_HEAD)          # [n, n]
24        scores_without_later_positions = shrunk_scores.masked_fill(~may_look_at, float("-inf"))
25        attention_shares = scores_to_shares(scores_without_later_positions)  # [n, n], rows add up to 1
26
27        # Each position collects every position's passed-on values, weighted by its shares of attention.
28        what_this_head_slice_collected = attention_shares @ head_passed_on_values   # [n, 32]
29        findings_of_each_head_slice.append(what_this_head_slice_collected)
30
31    all_head_slices_side_by_side = torch.cat(findings_of_each_head_slice, dim=-1)   # [n, 128]
32    return apply_weights(all_head_slices_side_by_side, block.head_mixing_weights, block.head_mixing_biases)

Reading it from the top:

  1. Normalise the hidden states, then make three new vectors from each one, each with its own table of weights: a looking query, a looked-at key, and a passed-on value.
  2. Build the "only earlier positions" rule. torch.ones(n, n, dtype=torch.bool) makes an n × n table filled with True, and torch.tril keeps its lower-left triangle, including the diagonal, and sets the rest to False. Row i is True for columns 0 to i.
  3. For each head slice, take that slice's 32 columns of the looking queries, looked-at keys, and passed-on values. looking_queries[:, first_column:after_last_column] keeps those columns of every row.
  4. Match: head_looking_queries @ head_looked_at_keys.T gives an n × n table in which row i, column j is position i's looking query multiplied by position j's looked-at key, number by number, and added up.
  5. Shrink the scores by √32. The 2017 authors' reason is, in their own words, a suspicion: "We suspect that for large values of d_k, the dot products grow large in magnitude, pushing the softmax function into regions where it has extremely small gradients" (Attention Is All You Need, section 3.2.1). Here d_k is 32, the numbers per head slice: without the shrink, the scores grow with the length of the vectors, and softmax tends to give nearly all the attention to one position.
  6. Hide later positions: .masked_fill(~may_look_at, float("-inf")) sets every score where the rule is False to minus infinity. (~ turns True into False and back.) Softmax then gives those positions a share of exactly 0, so a letter never looks ahead at the answer it is trying to guess.
  7. Share out and collect: softmax turns each row into shares that add up to 1, and attention_shares @ head_passed_on_values adds up every position's passed-on values, weighted by those shares.
  8. Put the head slices back together: torch.cat lays the four slices' 32-number findings side by side, 128 numbers again, and one more table of weights mixes them.
The head-slice loop again, with a note under each lineScroll sideways, or tap the diagram to open it full size

The head-slice loop again, with a note under each line

One head slice, with the real numbers for gooTap the image to open it full size

One head slice, with the real numbers for goo

For goo, position 3 in block 1, head slice 1, gives 4.0% of its attention to position 1, 92.9% to position 2, and 3.1% to itself: the same numbers as in Part 1.

You can follow one query through all of these steps in Part 1's demo:

6. The MLP: each letter on its own
PYTHON
1def mlp(hidden_states: torch.Tensor, block: TrainedBlock) -> torch.Tensor:
2    """[n, 128] -> [n, 128]: what the MLP adds to each hidden state."""
3    normalised = normalise(hidden_states, block.stretch_before_mlp, block.shift_before_mlp)
4    widened = apply_weights(normalised, block.widen_weights, block.widen_biases)       # [n, 512]
5    bent = bend(widened)                                                               # [n, 512]
6    return apply_weights(bent, block.narrow_weights, block.narrow_biases)              # [n, 128]

The MLP works on each letter's hidden state separately: no position looks at any other. It normalises, widens 128 numbers to 512, bends them, and narrows them back to 128.

What is it for? A common rule of thumb is that attention brings in context from other letters, and the MLP brings in knowledge stored in its weights, such as which letter usually ends a word. There is evidence for that in big models: one study described the MLP layers as acting like "key-value memories". But it is a rough picture, not a proven account. A 2023 study, Does Localization Inform Editing?, found that working out where a fact seems to be stored did "not provide any insight into which model MLP layer would be best to edit" to change it, and in a model this small nobody has worked out what each of the 512 numbers does. Why 512, four times 128? It copies the ratio in the 2017 design, where "the dimensionality of input and output is d_model=512, and the inner-layer has dimensionality d_ff=2048". The paper states those sizes without explaining the choice of four, and later models have kept it because it works, not because anyone derived it.

7. One block, then four

This is where the rewrite differs most from the original. The original writes a block as two lines, both reusing the name x:

PYTHON
1x = x + self.attn(self.ln1(x))
2x = x + self.mlp(self.ln2(x))

The rewrite gives the old and the new hidden states different names, so you can see that nothing is thrown away: each part only adds to what was there.

PYTHON
1def run_one_block(hidden_states_entering: torch.Tensor, block: TrainedBlock) -> torch.Tensor:
2    """[n, 128] -> [n, 128]"""
3    attention_contribution = attention(hidden_states_entering, block)
4    hidden_states_after_attention = hidden_states_entering + attention_contribution
5
6    mlp_contribution = mlp(hidden_states_after_attention, block)
7    hidden_states_leaving = hidden_states_after_attention + mlp_contribution
8    return hidden_states_leaving

The annotated version of these two functions is at the top of the post, next to the picture of step 3 opened up.

PYTHON
1def run_all_blocks(starting_hidden_states: torch.Tensor, model: TrainedModel) -> list:
2    """Returns the hidden states before block 1 and after each block: a list of 5 tables, each [n, 128]."""
3    hidden_states_after_each_block = [starting_hidden_states]
4    for block in model.blocks:
5        hidden_states_entering = hidden_states_after_each_block[-1]   # the latest result
6        hidden_states_after_each_block.append(run_one_block(hidden_states_entering, block))
7    return hidden_states_after_each_block
One block for position 3 of goo, with the first three of each vector's 128 numbersTap the image to open it full size

One block for position 3 of goo, with the first three of each vector's 128 numbers

The four blocks run in order, each starting from the previous block's result. Keeping every result in a list means the whole history can be looked at afterwards. For position 3 of goo, the hidden state's size grows from 0.50 before block 1 to 2.16, 2.72, 2.92 and 3.09 after blocks 1 to 4.

8. From the last hidden state to 65 scores
PYTHON
1def scores_for_next_letter(text: str, model: TrainedModel) -> torch.Tensor:
2    """A text -> [65] scores, one for each letter that could come next."""
3    letter_ids = letters_to_ids(text)[-MOST_LETTERS_THE_MODEL_CAN_SEE:]   # only the last 128 letters fit
4    starting_hidden_states = first_hidden_states(letter_ids, model)
5    hidden_states_after_each_block = run_all_blocks(starting_hidden_states, model)
6    last_letters_hidden_state = hidden_states_after_each_block[-1][-1]   # last block, last position: [128]
7    final_hidden_state = normalise(last_letters_hidden_state, model.final_stretch, model.final_shift)
8    return model.next_letter_rows @ final_hidden_state + model.next_letter_biases   # [65]

Only the last letter's hidden state is needed to guess the next letter. It is normalised one final time, giving the final hidden state, and compared with each of the 65 next-letter rows by a dot product, plus that row's bias. After go, o scores 4.878, the space 4.076, and d 4.062, exactly as in Part 1.

9. Does it give the same answers?

A rewrite for readability must not change the answers. The program builds Jibin Joseph's original model, loads the same trained numbers into it, and compares the 65 scores for four texts, from go to a full line of Coriolanus. The biggest difference is 0.0000029: the kind of difference that comes from adding the same numbers up in a different order, not a change in what the model does.

10. From scores to the next letter
PYTHON
1def chances_from_scores(scores: torch.Tensor, temperature: float, keep_biggest: int) -> torch.Tensor:
2    """[65] scores -> [65] chances that add up to 1."""
3    if temperature <= 0:   # temperature 0 means "always the biggest score": give it the whole wheel
4        whole_wheel_to_the_favourite = torch.zeros_like(scores)
5        whole_wheel_to_the_favourite[scores.argmax()] = 1.0
6        return whole_wheel_to_the_favourite
7    scaled_scores = scores / temperature
8    smallest_score_kept = torch.topk(scaled_scores, min(keep_biggest, len(scaled_scores))).values[-1]
9    trimmed_scores = scaled_scores.masked_fill(scaled_scores < smallest_score_kept, float("-inf"))
10    return scores_to_shares(trimmed_scores)
PYTHON
1def spin_the_wheel(chances: torch.Tensor, random_generator: torch.Generator) -> int:
2    """Pick one letter ID, each with its own chance."""
3    random_point = torch.rand(1, generator=random_generator)     # one number between 0 and 1
4    running_totals = torch.cumsum(chances, dim=0)                # [65], ends at 1
5    landed_on = torch.searchsorted(running_totals, random_point).item()
6    return min(landed_on, len(chances) - 1)                      # a safety net for rounding
PYTHON
1def write(start: str, letters_to_add: int, model: TrainedModel,
2          temperature: float = 0.8, keep_biggest: int = 65, seed: int = 0) -> str:
3    random_generator = torch.Generator().manual_seed(seed)
4    text_so_far = start
5    for _ in range(letters_to_add):
6        next_letter = choose_next_letter(text_so_far, model, temperature, keep_biggest, random_generator)
7        text_so_far = text_so_far + next_letter   # the only name that changes: the text itself grows
8    return text_so_far

torch.topk returns the k biggest numbers, biggest first, so .values[-1] is the smallest one kept. torch.cumsum makes running totals of the 65 chances, ending at 1, and torch.searchsorted finds where a random point between 0 and 1 falls among them: the slice of the wheel it lands in. A random-number generator with a fixed seed gives the same "random" spins every time, so the results can be repeated.

text_so_far = text_so_far + next_letter is the one place where I reused a name for a new value. I left it, because the text so far really is one thing that grows letter by letter, and the loop would be harder to follow with a new name for every length.

write asks choose_next_letter, from the top of this post, for one letter at a time, and adds each one to the text. Here is what it writes from ROMEO:, at a temperature of 0.8:

CODE
1ROMEO:
2Then this of hearts fear of Contic it thy slaught?
3
4PEY:
5Come, what I will one that that that be nother,
6As then, my see own banish, thou banish'd risgue,
7And have for a horsure as in by all death.
8Wh
What I learned from rewriting it
References