Visualisera attention-mönster
Kunna visualisera attention-huvuden och förstå varför de inte alltid förklarar beteende.
Förkunskaper
- DAttentionkrävs
Intuition
Attention-vikterna är en matris per huvud och lager: hur mycket varje position tittar på varje annan. De går att rita som en värmekarta, och vissa huvuden har slående tydliga mönster:
| Mönster | Vad huvudet gör |
|---|---|
| Diagonal | tittar på sig självt |
| Subdiagonal | föregående token |
| Vertikal kolumn | alla tittar på en «sänk»-token (ofta BOS) |
| Glest, semantiskt | pronomen → syftning, verb → subjekt |
| Induktion | tidigare förekomst av samma token |
Det är fascinerande att titta på — och lätt att övertolka.
Formellt
Varför attention inte är en förklaring. Jain & Wallace (2019) visade att man ofta kan hitta helt andra attention-fördelningar som ger samma utdata. Om flera olika «förklaringar» ger samma svar kan ingen av dem vara förklaringen.
Wiegreffe & Pinter (2019) nyanserade: attention är inte godtycklig, och den bär information — men den är inte tillräcklig som förklaring.
Tre konkreta skäl:
- Attention viktar värden (V), inte indata. En position kan få hög attention men bidra lite om dess V-vektor är liten.
- Flera lager och huvuden kombineras; ett enskilt huvuds mönster säger inte vad nätet gör.
- Attention sinks: många huvuden lägger stor vikt på den första token utan att den bär information — ett tekniskt fenomen, inte semantik.
Vad man ska göra i stället: använd attention för att generera hypoteser, och verifiera med intervention — ablera huvudet eller patcha dess aktivering och se om beteendet ändras. Attention-vikter är en karta över var modellen tittar, inte över vad som avgör.
Kod
from transformer_lens import HookedTransformer
import numpy as np
m = HookedTransformer.from_pretrained("gpt2-small")
text = "När Anna och Erik kom till affären gav Erik en bok till"
tokens = m.to_tokens(text)
_, cache = m.run_with_cache(tokens)
ord_ = m.to_str_tokens(tokens)
# Hitta huvuden med tydliga mönster i stället för att titta på alla 144
for lager in range(m.cfg.n_layers):
A = cache[f"blocks.{lager}.attn.hook_pattern"][0].cpu().numpy() # (huvud, pos, pos)
for huvud in range(A.shape[0]):
M = A[huvud]
forra = np.mean([M[i, i-1] for i in range(1, len(M))]) # föregående token?
sink = M[:, 0].mean() # attention sink?
if forra > 0.5:
print(f"L{lager}H{huvud}: föregående-token-huvud ({forra:.2f})")
elif sink > 0.6:
print(f"L{lager}H{huvud}: attention sink mot BOS ({sink:.2f})")
# Sista positionens fördelning i ett intressant huvud
A = cache["blocks.9.attn.hook_pattern"][0, 9].cpu().numpy()
for i in np.argsort(-A[-1])[:4]:
print(f" {ord_[i]!r}: {A[-1, i]:.2f}")
Att automatiskt klassificera huvudena (som ovan) är mycket mer användbart än att bläddra i 144 värmekartor.
Behärskning innebär
- Visualiserar attention-huvuden
- Tolkar mönster försiktigt
- Vet varför attention inte är en förklaring
Logga in för att göra övningarna och bygga upp din behärskning.
Källor
- arXiv — Attention is not Explanation — arXiv (öppen åtkomst; licens per artikel)
- arXiv — Attention is not not Explanation — arXiv (öppen åtkomst; licens per artikel)
- TransformerLens (MIT) — MIT