Inside the model
- shown
- real GPT-2
- layers
- 12
- dmodel
- 768
- heads
- 12
- vocab
- 50,257
- params
- 124M
All steps
Tokenize
GPT-2's byte-level BPE splits the text into tokens, each an id in a 50,257-entry vocabulary; Ġ marks a token that starts with a space. Click the tokens to see the merges.
ids [464, 3797, 3332, 319, 262]Embed
Each id selects its row of the embedding matrix WE, and the row of WP for its position is added; the strips show the first 16 of the 768 real numbers. Click a strip to see the lookup.
[5 × 768]Attention
Each position pulls in information from itself and earlier positions: its query (what it looks for) is scored against their keys (what they offer), and it takes a weighted mix of their values; later tokens are masked. Hover a lane to see its weights, and step through all 144 real heads below. Click attn for every product.
12 heads × [5 × 5]MLP
The MLP works on each position alone: it expands the 768 numbers to 3,072 neurons, applies GELU, and projects back. ln1 and ln2, the thin panes, normalise the stream before attn and mlp read it. Click mlp for the products.
[5 × 768] → [5 × 3072]Blocks 2–12
Eleven more blocks, same shape but each with its own weights, add their results to the stream. Right: the logit lens, what the last position would predict if the model stopped after each block. Lane colour sketches where information came from; the bars show the shares.
12 blocks ≈85M params · embeddings ≈39M (WE, reused as WU)Unembed
Only the last position predicts: lnf normalises it and the unembedding WU (the LM head, which GPT-2 ties to WEᵀ) scores all 50,257 tokens with one dot product each. softmax turns the scores into these probabilities, the real ones from GPT-2 small. Click lnf or WU to see it step by step.
[1 × 768] · WEᵀ → [1 × 50,257]Pick the next token
This page always takes the most likely token (greedy decoding) so every number stays real; chat models sample instead, see Unembed & Sampling.
greedy → Ġfloor · 7.6%
Code
idx = torch.tensor([enc.encode("The cat sat on the")]) # (1, T) token ids
x = wte(idx) + wpe(torch.arange(idx.size(1))) # (1, T, 768)
for block in h: # 12 blocks, own weights
x = x + block.attn(block.ln_1(x))
x = x + block.mlp(block.ln_2(x))
logits = lm_head(ln_f(x)[:, -1, :]) # last position only
idx_next = logits.argmax(-1, keepdim=True) # greedy
idx = torch.cat((idx, idx_next), dim=1) # append, next pass