Inside the model · Embedding
- shown
- toy scale
- vocab
- 50,257
- dmodel
- 8 GPT-2 768
- WE params
- 38.6M
- WP params
- 0.79M
All steps
One-hot ids
Each token id becomes a one-hot row: 50,257 zeros with a single 1 at the id. Stacked, the prompt is a 5 × 50,257 matrix.
[5 × 50,257]Row lookup
Multiplying the one-hot rows by WE selects one row of WE per token: a GEMM on paper, a table lookup in practice. Each row is a learned 768-number description of its token, and tokens used alike get similar rows: in GPT-2, Ġcat is closer to Ġdog (cosine 0.55) and Ġkitten (0.50) than to Ġon (0.21).
GPT-2 WE [50,257 × 768] · 38.6M paramsWhat the rows mean
Each row of WE is a learned description of a token. Squeezed from 768 dimensions to 2, words of a kind land near each other: numbers, days and months, family, function words. These are real GPT-2 rows; hover a word to see its closest ones.
real WE rows · 2 of 768 dims (PCA)Add positions
Attention treats its inputs as an unordered set (the causal mask gives only a weak hint of position), so each slot adds its own row of the position matrix WP. Slot i always gets row i. Try the other prompts: a repeated word, and the same words in two orders.
GPT-2 WP [1,024 × 768] · context 1,024Into the stream
The sum is the residual stream that enters block 1: one 768-wide lane per token, still unmixed.
h₀ [N × 768]
Code
# idx: (B, T) token ids
# wte(idx) equals F.one_hot(idx, 50257).float() @ wte.weight, without the multiply
tok_emb = self.transformer.wte(idx) # (B, T, 768)
pos = torch.arange(0, T, device=idx.device)
pos_emb = self.transformer.wpe(pos) # (T, 768)
x = self.transformer.drop(tok_emb + pos_emb) # the residual stream