Attention heads and the causal mask

Compare hand-built versions of five attention patterns seen in trained transformers, switch the causal mask on and off, and see how heads split the model width.

AdvancedExplained in Many heads and the causal mask

Five heads, one sentence

Each head applies its own scoring rule to the same tokens. Pick a head, then pick a query token to see where it looks.

Choose a head

Each token finds an earlier copy of itself and looks at the token that came next, which helps predict a repeat.

Induction heads are described by Elhage et al. (2021) and Olsson et al. (2022), who tie them to in-context learning.

Tap a token to make it the query. Shading shows where it looks.
Attention weights for the Induction head. Rows are queries, columns are keys. Cells above the diagonal are masked.
Querythecatchasedthedog.thecatchasedthe
1.00maskedmaskedmaskedmaskedmaskedmaskedmaskedmaskedmasked
0.880.12maskedmaskedmaskedmaskedmaskedmaskedmaskedmasked
0.790.110.11maskedmaskedmaskedmaskedmaskedmaskedmasked
0.050.940.010.01maskedmaskedmaskedmaskedmaskedmasked
0.650.090.090.090.09maskedmaskedmaskedmaskedmasked
0.600.080.080.080.080.08maskedmaskedmaskedmasked
0.020.480.000.000.480.000.00maskedmaskedmasked
0.050.010.920.010.010.010.010.01maskedmasked
0.050.010.010.910.010.010.010.010.01masked
0.020.320.000.000.320.000.000.320.000.00
HeadInduction
Querycat (7)
Looks most atchased (2), 92%

Try this

  • For each head, predict where position 9 will look before you tap it.
  • Turn off the causal mask and look at which heads change. The previous-token and first-token heads barely change: their high scores are never on later tokens, so the future gets only a sliver of weight.

Splitting the model width into heads

More heads means narrower heads. The total size of the attention layer does not change.

One token's 512-dimensional vector, cut into 8 heads of 64 dimensions. Each head has its own query, key, and value projections of shape 512 × 64, runs attention independently, and the 8 outputs are concatenated back to width 512 and multiplied by WO (512 × 512).

Model
8
Model width, d_model512
Per-head width, d_k = d_model / h64
Projection weights, 4 × d_model²1,048,576

Try "embedding", "softmax", "overfitting", or "backpropagation".