Inside the model · Attention
- shown
- toy scale
- tokens
- 5
- dmodel
- 8 GPT-2 768
- dhead
- 4 GPT-2 64
- heads
- 2 GPT-2 12
All steps
Projections
X (after ln1) is multiplied by WQ, WK and WV; GPT-2 adds biases, the toy leaves them out. A query is what a token looks for, a key what it offers, a value what it hands over when chosen. GPT-2 does all three in one GEMM, X · Wqkv.
GPT-2 [N×768]·[768×2304] + b · 3.5 MFLOPs / tokenScores
Q times K transposed. Row i, column j is the dot product of token i's query with token j's key: how much i should attend to j. Each head works in its own slice of dmodel / heads numbers (64 in GPT-2, 4 here). The grid is N × N: 25 cells here, over a million per head at 1,024 tokens, so doubling the context quadruples this step.
GPT-2 12 × [N×64]·[64×N]Scale
Divide by √dhead. Dot products grow with dimension; scaling keeps softmax from saturating from the start.
toy ÷ 2 · GPT-2 ÷ 8Causal mask
Position i predicts token i + 1, so it may only look at positions ≤ i: the upper triangle is set to −∞. In training every position is predicted at once, and without the mask each could just read the next word.
causal: j > i → −∞Softmax
Softmax each row, turning scores into weights that sum to 1; −∞ becomes 0 after exp. Each row is one token's attention distribution. These toy weights are random, so the pattern means nothing; the Forward pass shows GPT-2's real heads.
A = softmax(S / √d + mask)Weighted sum
Weight the rows of V by attention and sum them. Output row i blends the values of every visible token, and its color blends with them. The panel on the right breaks down the current row.
GPT-2 12 × [N×N]·[N×64]Output projection
Concatenate the heads' outputs, multiply by WO to mix them, and add the result back to the residual stream. Several small heads can each follow a different relation for the cost of one big one. In code: (B, T, 768) is split into (B, 12, T, 64) for attention and merged back here.
GPT-2 [N×768]·[768×768] + b · 1.2 MFLOPs / token
Code
B, T, C = x.size() # x = ln_1(h)
q, k, v = self.c_attn(x).split(self.n_embd, dim=2) # one GEMM for all three
q = q.view(B, T, self.n_head, C // self.n_head).transpose(1, 2) # (B, 12, T, 64)
k = k.view(B, T, self.n_head, C // self.n_head).transpose(1, 2)
v = v.view(B, T, self.n_head, C // self.n_head).transpose(1, 2)
att = q @ k.transpose(-2, -1) # (B, 12, T, T)
att = att * (1.0 / math.sqrt(k.size(-1))) # ÷ √64
att = att.masked_fill(self.bias[:, :, :T, :T] == 0, float("-inf"))
att = F.softmax(att, dim=-1)
y = att @ v # (B, 12, T, 64)
y = y.transpose(1, 2).contiguous().view(B, T, C) # heads side by side
y = self.c_proj(y) # · W_O