Visualising attention patterns
Be able to visualise attention heads and understand why they do not always explain behaviour.
Prerequisites
- DAttentionrequired
Intuition
The attention weights are a matrix per head and layer: how much each position looks at each other one. They can be drawn as a heat map, and some heads have strikingly clear patterns:
| Pattern | What the head does |
|---|---|
| Diagonal | looks at itself |
| Subdiagonal | the previous token |
| A vertical column | everyone looks at a «sink» token (often BOS) |
| Sparse, semantic | pronoun → referent, verb → subject |
| Induction | an earlier occurrence of the same token |
It is fascinating to look at — and easy to over-interpret.
Formal
Why attention is not an explanation. Jain & Wallace (2019) showed that you can often find entirely different attention distributions that give the same output. If several different «explanations» give the same answer, none of them can be the explanation.
Wiegreffe & Pinter (2019) added nuance: attention is not arbitrary, and it does carry information — but it is not sufficient as an explanation.
Three concrete reasons:
- Attention weights the values (V), not the input. A position can get high attention but contribute little if its V vector is small.
- Several layers and heads are combined; a single head's pattern does not say what the network is doing.
- Attention sinks: many heads put a lot of weight on the first token without it carrying information — a technical phenomenon, not semantics.
What to do instead: use attention to generate hypotheses, and verify with an intervention — ablate the head or patch its activation and see whether the behaviour changes. Attention weights are a map of where the model looks, not of what decides.
Code
from transformer_lens import HookedTransformer
import numpy as np
m = HookedTransformer.from_pretrained("gpt2-small")
text = "When Anna and Erik came to the shop, Erik gave a book to"
tokens = m.to_tokens(text)
_, cache = m.run_with_cache(tokens)
words = m.to_str_tokens(tokens)
# Find heads with clear patterns instead of looking at all 144
for layer in range(m.cfg.n_layers):
A = cache[f"blocks.{layer}.attn.hook_pattern"][0].cpu().numpy() # (head, pos, pos)
for head in range(A.shape[0]):
M = A[head]
previous = np.mean([M[i, i-1] for i in range(1, len(M))]) # a previous-token head?
sink = M[:, 0].mean() # an attention sink?
if previous > 0.5:
print(f"L{layer}H{head}: a previous-token head ({previous:.2f})")
elif sink > 0.6:
print(f"L{layer}H{head}: an attention sink on BOS ({sink:.2f})")
# The last position's distribution in an interesting head
A = cache["blocks.9.attn.hook_pattern"][0, 9].cpu().numpy()
for i in np.argsort(-A[-1])[:4]:
print(f" {words[i]!r}: {A[-1, i]:.2f}")
Classifying the heads automatically (as above) is far more useful than paging through 144 heat maps.
Mastery means
- Visualises attention heads
- Interprets the patterns carefully
- Knows why attention is not an explanation
Sign in to do the exercises and build your mastery up.
Sources
- arXiv — Attention is not Explanation — arXiv (open access; licence per article)
- arXiv — Attention is not not Explanation — arXiv (open access; licence per article)
- TransformerLens (MIT) — MIT