Inside the model · Attention

Attention

causal self-attention · block 1 · head 1

shown
toy scale
tokens
5
dmodel
8 GPT-2 768
dhead
4 GPT-2 64
heads
2 GPT-2 12

All steps

  1. 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 / token

  2. Scores

    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]

  3. Scale

    Divide by √dhead. Dot products grow with dimension; scaling keeps softmax from saturating from the start.

    toy ÷ 2 · GPT-2 ÷ 8

  4. Causal 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 → −∞

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

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

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

Go deeper