Inside the model · Embedding

Embedding

token + position · GPT-2

shown
toy scale
vocab
50,257
dmodel
8 GPT-2 768
WE params
38.6M
WP params
0.79M

All steps

  1. 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]

  2. 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 params

  3. What 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)

  4. 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,024

  5. Into 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

Go deeper