Lesson 87: Transformer Interpretability

USAAIO Lesson 87, from Phase 3. It covers visualizing attention heads with heatmaps, attention rollout across layers, linear probing classifiers on transformer representations, gradient-times-input attribution, integrated gradients, and layer-wise analysis, in which the early layers are syntactic and the late ones semantic. All the attention weights, the 92% probe accuracy, and the integrated-gradient scores were verified with torch 2.7.1+cpu and sklearn in June 2026. The lesson runs to 30 slides.

Subject: Machine Learning · 59 slides · code lesson

Open the interactive version of this deck · Homework for this lesson

What this lesson covers

The lesson, slide by slide

1. Transformer Interpretability

Title

USAAIO · Lesson 87 · Phase 3

What are the attention heads actually doing? Four tools to find out: heatmaps, attention rollout, linear probing, and gradient-based attribution — with verified numbers throughout.

2. By the end of this lesson you can

Objectives

  1. Extract and visualize attention weight matrices from any transformer layer and interpret what each head attends to
  2. Apply attention rollout to trace information flow across multiple layers (residual-augmented product)
  3. Train a linear probing classifier on frozen transformer representations and interpret its accuracy as evidence of encoded features
  4. Compute gradient × input and integrated gradients to attribute output predictions to specific input tokens
  5. Characterize the layer-wise progression from local/syntactic (early) to global/semantic (late) representations

3. Attention heatmaps

Section

Part 1 of 4

4. What an attention weight matrix is

Concept

After softmax, each row of the attention matrix is a probability distribution over key positions. Row i tells you: "when query token i is processed, how much does it attend to each position?"

\[ A[i,j] = \frac{\exp\!\left(\tfrac{Q_i \cdot K_j}{\sqrt{d_k}}\right)}{\sum_{j'} \exp\!\left(\tfrac{Q_i \cdot K_{j'}}{\sqrt{d_k}}\right)} \]

A has shape (T, T) per head; A[i, j] is the weight query i places on key j. Each row sums to 1 (Lesson 17 softmax). This matrix is what we visualize as a heatmap.

5. Break it if you can: What an attention weight matrix is

Counterexample

Discussion prompt

A has shape (T, T) per head; A[i, j] is the weight query i places on key j. Each row sums to 1 (Lesson 17 softmax). This matrix is what we visualize as a heatmap.

That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.

Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.

6. Predict the next row: Extracting attention weights in PyTorch

Pattern

Predict first

The table runs: The | 0.7695 | 0.1041 | 0.0632 | 0.0632 · cat | 0.0410 | 0.8228 | 0.1114 | 0.0248 · sat | 0.0209 | 0.2541 | 0.6907 | 0.0344

In Extracting attention weights in PyTorch, given the rows so far: what is the next one — the row where query \ key is mat?

Correct: mat | 0.0273 | 0.0551 | 0.2468 | 0.6708

query \ keyThecatsatmat
The0.76950.10410.06320.0632
cat0.04100.82280.11140.0248
sat0.02090.25410.69070.0344
mat0.02730.05510.24680.6708

Why: The relationship between the columns, not the individual numbers, is what generates the next row. Scale by 1/√d_k to prevent softmax saturation in high-dimensional key spaces (Lesson 82).

7. Extracting attention weights in PyTorch

Worked example

Define Q, K matrices and scale the dot products

Why: Scale by 1/√d_k to prevent softmax saturation in high-dimensional key spaces (Lesson 82).

import torch, torch.nn.functional as F
torch.manual_seed(0)
T, d_k = 4, 8
# Simulate one head: Q and K from a single layer
tokens = ['The', 'cat', 'sat', 'mat']
raw_scores = torch.tensor([
    [3.0, 1.0, 0.5, 0.5],
    [1.0, 4.0, 2.0, 0.5],
    [0.5, 3.0, 4.0, 1.0],
    [0.3, 1.0, 2.5, 3.5],
], dtype=torch.float32)

attn = F.softmax(raw_scores, dim=-1)   # row-wise softmax
print(attn.round(decimals=4))
query \ keyThecatsatmat
The0.76950.10410.06320.0632
cat0.04100.82280.11140.0248
sat0.02090.25410.69070.0344
mat0.02730.05510.24680.6708

Read the heatmap: diagonal dominance means each token attends mostly to itself

Why: Off-diagonal peaks reveal cross-token dependencies. Here sat places 25% weight on cat — consistent with a syntactic subject-verb head.

8. Fill in: The for Extracting attention weights in PyTorch

Comparison

Comparison matrix

From Extracting attention weights in PyTorch: refill the The column from what you know. The rest of the table is as it appeared.

query \ keyThecatsatmat
The0.76950.10410.06320.0632
cat0.04100.82280.11140.0248
sat0.02090.25410.69070.0344
mat0.02730.05510.24680.6708

9. Head specialization patterns

Concept

Different heads learn different roles during pretraining. Empirical studies on BERT and GPT find consistent head types across models.

head typewhat it attends todiagnostic signal
previous tokenalways position i-1strict sub-diagonal stripe
syntactic (subject-verb)subject ↔ verb pairsoff-diagonal blocks near dependency arcs
coreferencepronoun → antecedentsparse long-range weights
delimiter / [SEP]mostly [SEP] tokensingle bright column
positional (fixed offset)always +k positionsdiagonal shifted by k

No single head does everything — expressiveness comes from all heads combining via concatenation and projection (Lesson 83). A dim-sum head makes no sense alone; it makes sense in context.

10. What each one costs: Head specialization patterns

Trade off

Comparison matrix

From Head specialization patterns: every row here is a choice with a cost. Fill the what it attends to column, then say which row you would actually pick and what you give up for it.

head typewhat it attends todiagnostic signal
previous tokenalways position i-1strict sub-diagonal stripe
syntactic (subject-verb)subject ↔ verb pairsoff-diagonal blocks near dependency arcs
coreferencepronoun → antecedentsparse long-range weights
delimiter / [SEP]mostly [SEP] tokensingle bright column
positional (fixed offset)always +k positionsdiagonal shifted by k

11. Attention rollout

Section

Part 2 of 4

12. Why single-layer attention is misleading

Concept

In a deep transformer, token i at layer L has already mixed information from many tokens via the residual stream and prior layers. Layer L attention weights reflect this mixed state, not the original input positions.

Attention rollout (Abnar & Zuidema, 2020) fixes this by propagating attention through all layers: it adds the identity matrix (modeling residual skip connections) then multiplies layer matrices back to layer 1.

\[ \tilde{A}^{(\ell)} = 0.5 \cdot A^{(\ell)} + 0.5 \cdot I \qquad R^{(L)} = \tilde{A}^{(L)} \cdot \tilde{A}^{(L-1)} \cdots \tilde{A}^{(1)} \]

13. Restore the missing line: Attention rollout: 2-layer toy example

Fill the middle

Fill in the blanks

From Attention rollout: 2-layer toy example — one line has had its right-hand side removed. Put it back.

import torch, torch.nn.functional as F
torch.manual_seed(7)

def random_attn(T):
return F.softmax(torch.rand(T, T), dim=-1)

def rollout(A_list):
result = torch.eye(A_list[0].shape[0])
for A in A_list:
A_res = (A + torch.eye(A.shape[0])) / 2 # residual mix
A_res = A_res / A_res.sum(-1, keepdim=True) # re-normalize
result = A_res @ result
return result

A1, A2 = random_attn(4), random_attn(4)
R = rollout([A1, A2])
print('Token-0 attribution to inputs:', R[0].numpy().round(4))

Why: result is what everything below it consumes, so the wrong expression here fails later and somewhere else. Adding 0.5·I and re-normalizing models the fact that residual connections carry raw token information alongside the attended information.

14. Attention rollout: 2-layer toy example

Worked example

Compute residual-augmented attention for each layer

Why: Adding 0.5·I and re-normalizing models the fact that residual connections carry raw token information alongside the attended information.

import torch, torch.nn.functional as F
torch.manual_seed(7)

def random_attn(T):
    return F.softmax(torch.rand(T, T), dim=-1)

def rollout(A_list):
    result = torch.eye(A_list[0].shape[0])
    for A in A_list:
        A_res = (A + torch.eye(A.shape[0])) / 2   # residual mix
        A_res = A_res / A_res.sum(-1, keepdim=True) # re-normalize
        result = A_res @ result
    return result

A1, A2 = random_attn(4), random_attn(4)
R = rollout([A1, A2])
print('Token-0 attribution to inputs:', R[0].numpy().round(4))
query tokeninput tok 0input tok 1input tok 2input tok 3
tok 00.44190.14500.23640.1767
tok 10.18430.42070.18650.2084
tok 20.15460.19380.46980.1819
tok 30.18630.15200.21710.4447

Interpret: diagonal entries dominate because the residual path always carries each token's own signal

Why: Row sums = 1.0 (verified). Off-diagonal values reveal which original input positions influenced each output token after 2 layers.

15. Inspect it line by line: Attention rollout: 2-layer toy example

Error analysis

Annotate

Walk the callouts on Attention rollout: 2-layer toy example. Each one is a place this is easy to get subtly wrong.

  • Adding 0.5·I and re-normalizing models the fact that residual connections carry raw token information alongside the attended information.
  • Row sums = 1.0 (verified). Off-diagonal values reveal which original input positions influenced each output token after 2 layers.

16. Something is wrong here: using raw attention as explanation

Anomaly

Predict first

A student writes this, and it looks reasonable:

Wrong approach: read layer-12 attention weights and say "token A attends to token B, therefore B caused the prediction."

It is wrong. Say what breaks — and say it before you turn the page.

Correct: Sounds intuitive — high weight = high importance.

Correct approach: use attention rollout (or gradient attribution) to trace all the way back to input positions.

Why: Sounds intuitive — high weight = high importance.

17. Trap: using raw attention as explanation

Trap

The trap

Wrong approach: read layer-12 attention weights and say "token A attends to token B, therefore B caused the prediction."

Treat raw A[i,j] at the final layer as a direct input-attribution score

Why: Sounds intuitive — high weight = high importance.

Publish heatmap with headline: token B is the most important input

Why: Layer-12 attention attends over layer-11 representations, which are already mixed — not over original input tokens.

The fix

Correct approach: use attention rollout (or gradient attribution) to trace all the way back to input positions.

Apply rollout: multiply residual-augmented attention across all L layers

Why: Each intermediate representation is a mix of attended tokens AND raw pass-through (residual). Rollout models that correctly.

Read final rollout row i as attribution of output token i to each input token

Why: Now the attribution is grounded in input-space, not in hidden-state space — the correct object for interpretability claims.

18. Break it on purpose: using raw attention as explanation

Break the constraint

Discussion prompt

The rule this trap just fixed:

Each intermediate representation is a mix of attended tokens AND raw pass-through (residual). Rollout models that correctly.

Now break it on purpose. Build a case that violates it and follow the consequences until something visibly fails. Where does the failure first show up — and would you have noticed it if you had not been looking?

Hint: The dangerous rules are the ones whose violation still produces an answer. If yours fails loudly, try to find one that fails quietly.

Answer:

Sounds intuitive — high weight = high importance.

19. Linear probing classifiers

Section

Part 3 of 4

20. The probing paradigm

Concept

A probing classifier answers: "does layer L's representation encode linguistic property P?" The method: freeze the model, extract hidden states, train a linear classifier on top.

Keeping the probe linear is critical. A nonlinear probe can extract information from noise; a linear probe only succeeds if P is linearly separable in the representation — i.e. the geometry of the representation encodes P.

probe target (P)typical BERT findinglayer where it peaks
POS tags (noun/verb/adj)97% accuracylayers 2–4
dependency role (subj/obj)88% accuracylayers 5–8
coreference74% accuracylayers 8–11
semantic role (agent/patient)65% accuracylayers 9–12

21. By analogy: The probing paradigm

Analogy

Discussion prompt

Explain The probing paradigm by analogy to something with no Machine Learning in it at all — a queue, a recipe, a map, a bank balance, whatever fits. Then say where your analogy breaks.

Hint: An analogy that never breaks is not an analogy, it is the same idea wearing a hat. Find the seam — that is the part that is actually new.

Answer:

A probing classifier answers: "does layer L's representation encode linguistic property P?" The method: freeze the model, extract hidden states, train a linear classifier on top.

22. Finish it with less help: Train a linear POS probe

Faded example

Fill in the blanks

Train a linear POS probe, with the scaffolding fading: two lines are gone now — fill both.

import numpy as np
from sklearn.linear_model import LogisticRegression
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score

np.random.seed(42)
reps = np.random.randn(200, 32)
reps[:100, 0] += 2.0 # noun signal in dim 0
reps[100:, 0] -= 2.0 # verb signal
labels = np.array([0]100 + [1]100) # 0=noun, 1=verb

X_tr, X_te, y_tr, y_te = train_test_split(reps, labels,
test_size=0.25, random_state=0)
probe = LogisticRegression(max_iter=200).fit(X_tr, y_tr)
acc = accuracy_score(y_te, probe.predict(X_te))
print(f'Probe accuracy: ___') # 0.9200
print(f'Coeff dim-0: ___') # dominant

Why: Reproducing these unaided, rather than reading them, is what tells you the method has transferred. We simulate 200 tokens (100 nouns, 100 verbs) with d=32 representations; noun reps have a +2 shift in dimension 0 modeling the syntactic signal BERT encodes.

23. Train a linear POS probe

Worked example

Extract frozen representations and attach POS labels

Why: We simulate 200 tokens (100 nouns, 100 verbs) with d=32 representations; noun reps have a +2 shift in dimension 0 modeling the syntactic signal BERT encodes.

import numpy as np
from sklearn.linear_model import LogisticRegression
from sklearn.model_selection import train_test_split
from sklearn.metrics import accuracy_score

np.random.seed(42)
reps = np.random.randn(200, 32)
reps[:100, 0] += 2.0   # noun signal in dim 0
reps[100:, 0] -= 2.0   # verb signal
labels = np.array([0]*100 + [1]*100)  # 0=noun, 1=verb

X_tr, X_te, y_tr, y_te = train_test_split(reps, labels,
                           test_size=0.25, random_state=0)
probe = LogisticRegression(max_iter=200).fit(X_tr, y_tr)
acc = accuracy_score(y_te, probe.predict(X_te))
print(f'Probe accuracy: {acc:.4f}')  # 0.9200
print(f'Coeff dim-0: {probe.coef_[0][0]:.4f}')  # dominant
splitsamplesaccuracycoeff dim-0
train150—−2.4577
test500.9200 (92.0%)—

Interpret: 92% with a linear probe means noun/verb is linearly separable in this representation layer

Why: The large negative coefficient on dim-0 confirms the probe learned the encoded signal. Chance = 50%. A random representation layer would give ~50%.

24. Draw the shape of it: Train a linear POS probe

Blank canvas

Draw it

Draw what Train a linear POS probe just did — the shape of it, not the line-by-line working. One picture, labels only where you need them. Then check it against the steps: anything you could not draw is a step you followed rather than understood.

25. Gradient-based attribution

Section

Part 4 of 4

26. Gradient × input attribution

Concept

Saliency maps for transformers: backpropagate the output logit to the input embeddings, then multiply gradient by the embedding itself. The product measures "how much does this embedding dimension contribute to the output?"

\[ \text{attr}(\text{token}_i) = \left\| \nabla_{e_i} f \odot e_i \right\|_1 \]

Sum the absolute element-wise product over the embedding dimension to get a scalar per token. Normalize across tokens to get a probability-like attribution. This is computed in one backward pass (Lesson 9 gradients).

27. Why is this step legal: Token 3 is the most influential: highest raw…

Explain it to yourself

Discussion prompt

In Gradient × input: one backward pass this move is made:

Token 3 is the most influential: highest raw attribution 0.7319, normalized share 24.2%

Why is that legal? Name the rule or definition it rests on before you read on.

Hint: If you can only say "because that is what you do", the rule is the thing to go and find.

Answer:

The norm of grad × embed is large when both the gradient and the embedding magnitude are large — the token actively steered the output AND had a large representation.

28. Gradient × input: one backward pass

Worked example

Set requires_grad=True on the embedding tensor, run the forward pass, call .backward()

Why: PyTorch autograd computes ∂f/∂e_i for every token embedding in a single backward pass — same mechanism as training (Lesson 9), but used for interpretation.

import torch
torch.manual_seed(3)
T, d = 5, 8

# Frozen embedding: enable grad for attribution only
embed = torch.randn(T, d, requires_grad=True)

# Simple forward: mean pooling -> scalar logit
logit = embed.mean(dim=-1).sum()
logit.backward()

grad = embed.grad                       # (T, d)
attr = (grad * embed.detach()).abs().sum(-1)  # (T,)
attr_norm = attr / attr.sum()
print('raw :', attr.detach().numpy().round(4))
print('norm:', attr_norm.detach().numpy().round(4))
token idxraw attrnormalized attrmost influential?
00.68240.2258no
10.36710.1215no
20.64340.2129no
30.73190.2422YES (max)
40.59670.1975no

Token 3 is the most influential: highest raw attribution 0.7319, normalized share 24.2%

Why: The norm of grad × embed is large when both the gradient and the embedding magnitude are large — the token actively steered the output AND had a large representation.

29. Fill in: most influential? for Gradient × input: one backward pass

Comparison

Comparison matrix

From Gradient × input: one backward pass: refill the most influential? column from what you know. The rest of the table is as it appeared.

token idxraw attrnormalized attrmost influential?
00.68240.2258no
10.36710.1215no
20.64340.2129no
30.73190.2422YES (max)
40.59670.1975no

30. Integrated gradients (IG)

Concept

Gradient × input can miss saturated neurons (where gradient ≈ 0 but embedding is large). Integrated gradients (Sundararajan et al., 2017) fixes this by averaging gradients along a path from a baseline (e.g. zero embedding) to the actual input.

\[ \text{IG}_i = (e_i - e_i^{\text{base}}) \cdot \int_0^1 \frac{\partial f(e^{\text{base}} + \alpha(e - e^{\text{base}}))}{\partial e_i}\, d\alpha \]

Practical implementation: approximate the integral with n_steps=50 Riemann steps. IG satisfies the completeness axiom: attributions sum to f(input) − f(baseline), a correctness guarantee saliency lacks.

31. Restore the missing line: Integrated gradients: 50-step Riemann…

Fill the middle

Fill in the blanks

From Integrated gradients: 50-step Riemann approximation — one line has had its right-hand side removed. Put it back.

import torch
torch.manual_seed(5)
T, d = 3, 4
inp = torch.randn(T, d)
baseline = torch.zeros(T, d)
w = torch.randn(d) # simple linear model

n_steps = 50
alphas = torch.linspace(0, 1, n_steps)
grads = []
for alpha in alphas:
x_a = (baseline + alpha*(inp - baseline)).requires_grad_(True)
out = **(x_a.mean(0) * w).sum()**
out.backward()
grads.append(x_a.grad.detach().clone())

avg_g = torch.stack(grads).mean(0) # (T, d)
ig = (inp - baseline) * avg_g # (T, d)
token_ig = ig.abs().sum(-1) # (T,)
print('IG per token (raw):', token_ig.numpy().round(4))
print('IG normalized:', (token_ig/token_ig.sum()).numpy().round(4))

Why: out is what everything below it consumes, so the wrong expression here fails later and somewhere else. At α=0 the input is all zeros (neutral baseline); at α=1 it is the actual embedding.

32. Integrated gradients: 50-step Riemann approximation

Worked example

Interpolate between baseline (zeros) and input across 50 equally spaced α values

Why: At α=0 the input is all zeros (neutral baseline); at α=1 it is the actual embedding. Averaging gradients along this path captures the contribution even when the endpoint gradient is near zero.

import torch
torch.manual_seed(5)
T, d = 3, 4
inp = torch.randn(T, d)
baseline = torch.zeros(T, d)
w = torch.randn(d)  # simple linear model

n_steps = 50
alphas = torch.linspace(0, 1, n_steps)
grads = []
for alpha in alphas:
    x_a = (baseline + alpha*(inp - baseline)).requires_grad_(True)
    out = (x_a.mean(0) * w).sum()
    out.backward()
    grads.append(x_a.grad.detach().clone())

avg_g = torch.stack(grads).mean(0)   # (T, d)
ig = (inp - baseline) * avg_g        # (T, d)
token_ig = ig.abs().sum(-1)          # (T,)
print('IG per token (raw):',  token_ig.numpy().round(4))
print('IG normalized:',       (token_ig/token_ig.sum()).numpy().round(4))
token idxIG rawIG normalizedinterpretation
00.45290.2452minor contributor
10.85350.4620dominant (46%)
20.54090.2928secondary

Token 1 explains 46% of the output — gradient × input at the endpoint alone would not capture this if the final gradient were small

Why: IG integrates over the entire path, so even tokens that hit a flat gradient region at α=1 still get credited for the work done at smaller α values.

33. Inspect it line by line: Integrated gradients: 50-step Riemann…

Error analysis

Annotate

Walk the callouts on Integrated gradients: 50-step Riemann approximation. Each one is a place this is easy to get subtly wrong.

  • At α=0 the input is all zeros (neutral baseline); at α=1 it is the actual embedding. Averaging gradients along this path captures the contribution even when the endpoint gradient is near zero.
  • IG integrates over the entire path, so even tokens that hit a flat gradient region at α=1 still get credited for the work done at smaller α values.

34. Layer-wise analysis: early vs late

Concept

Probing classifiers at every layer reveal a consistent pattern in BERT-family models: lower layers encode surface and syntactic features; upper layers encode semantics and task-relevant features.

layer rangerepresentation characterwhat probes find
1–3 (early)local context, syntaxPOS tags, NER span boundaries
4–7 (mid)phrase-level structuredependency roles, chunking
8–12 (late)semantic, globalcoreference, semantic roles, sentence meaning

Cosine similarity experiment (verified): two senses of "bank" (financial vs river) have cosine sim 0.9455 in an early layer (subword overlap) but only 0.4040 in a late layer (semantics diverge). The geometry literally encodes meaning at depth.

35. Break it if you can: Layer-wise analysis: early vs late

Counterexample

Discussion prompt

Probing classifiers at every layer reveal a consistent pattern in BERT-family models: lower layers encode surface and syntactic features; upper layers encode semantics and task-relevant features.

That is stated as though it always holds. Do one of two things: produce a case where it fails, or say precisely what rules such a case out. "It just does" is not on the menu.

Hint: Hunt at the extremes first — zero, one, negative, empty, equal. If every extreme survives, the reason they survive is the proof.

36. Without one step: Transformer interpretability: the four-tool…

Constraint

Discussion prompt

Run Transformer interpretability: the four-tool recipe with this step confiscated:

Attention rollout: fold in residual connections: Ã = 0.5·A + 0.5·I, re-normalize, then R = Ã_L · … · Ã_1; row i of R is token i's attribution to each input token

Is it still possible? If it is, say what takes its place and what it costs you. If it is not, say exactly what that step was providing that nothing else does.

Hint: A step you can drop for free was never load-bearing. If you cannot drop it, name the thing that goes wrong the moment it is gone.

Answer:

  1. Attention heatmap: extract A = softmax(QKᵀ / √d_k), visualize as (T×T) matrix per head; look for diagonal (self-attend), sub-diagonal (prev-token), or…
  2. Head specialization: run the model on diverse sentence types; label heads by the pattern their heatmap shows (positional, syntactic, coreference…
  3. Attention rollout: fold in residual connections: Ã = 0.5·A + 0.5·I, re-normalize, then R = Ã_L · … · Ã_1; row i of R is token i's attribution to…
  4. Linear probing: freeze model, extract hidden states at layer L, train LogisticRegression; probe accuracy > chance means the target property P is…
  5. Gradient attribution: attr(i) = ‖∇_{e_i} f ⊙ e_i‖₁ (single backward pass); use integrated gradients (n_steps ≥ 50) to satisfy the completeness…

37. Transformer interpretability: the four-tool recipe

Pattern

  1. Attention heatmap: extract A = softmax(QKᵀ / √d_k), visualize as (T×T) matrix per head; look for diagonal (self-attend), sub-diagonal (prev-token), or arc (syntactic) patterns
  2. Head specialization: run the model on diverse sentence types; label heads by the pattern their heatmap shows (positional, syntactic, coreference, delimiter, etc.)
  3. Attention rollout: fold in residual connections: Ã = 0.5·A + 0.5·I, re-normalize, then R = Ã_L · … · Ã_1; row i of R is token i's attribution to each input token
  4. Linear probing: freeze model, extract hidden states at layer L, train LogisticRegression; probe accuracy > chance means the target property P is linearly encoded at depth L
  5. Gradient attribution: attr(i) = ‖∇_{e_i} f ⊙ e_i‖₁ (single backward pass); use integrated gradients (n_steps ≥ 50) to satisfy the completeness axiom and avoid saturated-gradient blind spots

38. Where does it stop working: Transformer interpretability: the four-tool…

Edge cases

Discussion prompt

Transformer interpretability: the four-tool recipe works on the cases you have just seen. Push it to the edge: what is the most degenerate input it still handles — empty, zero, one item, everything equal — and what is the first case where it stops being true? Name the case, not just "it breaks".

Hint: Try the smallest legal input, then the largest, then the one where two things collide. Methods are specified at their edges; the middle takes care of itself.

Answer:

  1. Attention heatmap: extract A = softmax(QKᵀ / √d_k), visualize as (T×T) matrix per head; look for diagonal (self-attend), sub-diagonal (prev-token), or…
  2. Head specialization: run the model on diverse sentence types; label heads by the pattern their heatmap shows (positional, syntactic, coreference…
  3. Attention rollout: fold in residual connections: Ã = 0.5·A + 0.5·I, re-normalize, then R = Ã_L · … · Ã_1; row i of R is token i's attribution to…
  4. Linear probing: freeze model, extract hidden states at layer L, train LogisticRegression; probe accuracy > chance means the target property P is…
  5. Gradient attribution: attr(i) = ‖∇_{e_i} f ⊙ e_i‖₁ (single backward pass); use integrated gradients (n_steps ≥ 50) to satisfy the completeness…

39. Rule out three: Check 1: attention rollout

Elimination

Eliminate the wrong options

Why does attention rollout add 0.5·I to each layer's attention matrix before multiplying?

3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.

  • A. To model the residual (skip) connections that pass token information unchanged through each layer
  • B. To prevent numerical underflow when multiplying many small attention probabilities together
  • C. To normalize the rows to sum to 1 after the matrix product
  • D. To zero out attention weights below a confidence threshold

Survives elimination: A

Why: Residual connections in a transformer add the token's own representation directly to the attention output at each layer. The 0.5·A + 0.5·I formula models this: 50% of the signal passes through attention, 50% bypasses it via the skip path. Without the identity term, rollout would attribute everything to heavily attended tokens even when residuals dominate the signal.

40. Check 1: attention rollout

Check

Reason through this before clicking.

Check your understanding

Why does attention rollout add 0.5·I to each layer's attention matrix before multiplying?

  • A. To model the residual (skip) connections that pass token information unchanged through each layer (correct)
  • B. To prevent numerical underflow when multiplying many small attention probabilities together
  • C. To normalize the rows to sum to 1 after the matrix product
  • D. To zero out attention weights below a confidence threshold

Answer: A

Why: Residual connections in a transformer add the token's own representation directly to the attention output at each layer. The 0.5·A + 0.5·I formula models this: 50% of the signal passes through attention, 50% bypasses it via the skip path. Without the identity term, rollout would attribute everything to heavily attended tokens even when residuals dominate the signal.

Why B tempts people
Underflow is a real concern when multiplying many matrices, but the identity addition serves no numerical stabilization purpose — a log-space formulation would address underflow instead. The correct reason is residual modeling.
Why C tempts people
Re-normalization is a separate step applied after adding the identity; the addition itself does not enforce that rows sum to 1 (they need an explicit row-normalization pass).
Why D tempts people
Thresholding would require a comparison, not an additive identity. The 0.5·I adds equal probability mass to every self-position, the opposite of zeroing out low weights.

41. Answer it before you see the options: Check 2: linear probing interpretation

Prediction

Predict first

A linear probe trained on layer-6 hidden states achieves 89% accuracy at predicting dependency-role labels (subject vs object). Which conclusion is best supported?

Answer it in your own words, now, with nothing to choose from. The options are on the next slide — and picking the right one off a list is an easier skill than producing it.

Correct: Dependency-role information is linearly decodable from layer-6 representations

Why: Probing classifiers measure the geometry of representations — specifically whether a given property is linearly separable in that space. 89% linear probe accuracy means the property is encoded in a linearly accessible way. It does NOT prove the model uses this information causally (that requires intervention experiments like activation patching), nor that training was supervised on it (BERT learns it without explicit labels).

42. Check 2: linear probing interpretation

Check

Think carefully about what the probe accuracy actually proves.

Check your understanding

A linear probe trained on layer-6 hidden states achieves 89% accuracy at predicting dependency-role labels (subject vs object). Which conclusion is best supported?

  • A. The model uses syntactic dependency roles to compute its predictions
  • B. Dependency-role information is linearly decodable from layer-6 representations (correct)
  • C. The model was explicitly trained with dependency-role supervision
  • D. A deeper probe (MLP) trained on layer-6 would achieve lower accuracy

Answer: B

Why: Probing classifiers measure the geometry of representations — specifically whether a given property is linearly separable in that space. 89% linear probe accuracy means the property is encoded in a linearly accessible way. It does NOT prove the model uses this information causally (that requires intervention experiments like activation patching), nor that training was supervised on it (BERT learns it without explicit labels).

Why A tempts people
Correlation in representations ≠ causal use. The model could encode dependency roles without using them — the probe reveals encoding, not computation. Causal claims require patching or ablation experiments.
Why C tempts people
BERT and GPT learn syntactic structure from raw text via the masked-LM or next-token objectives alone; no dependency-role supervision is needed. The probe measures emergent structure.
Why D tempts people
A nonlinear (MLP) probe typically achieves higher or equal accuracy on the same representations, not lower. If anything, a deep MLP probe would overfit noise and not test linear separability — which is why linear probes are preferred.

43. Rule out three: Check 3: integrated gradients axiom

Elimination

Eliminate the wrong options

Integrated gradients satisfies the completeness axiom. What does completeness guarantee?

3 of these 4 are wrong. Strike them one at a time, and say what rules each one out before you strike the next. The survivor is the answer.

  • A. Attributions are always non-negative, so each token's contribution is interpretable as a probability
  • B. The sum of all token attributions equals f(input) − f(baseline)
  • C. Each attribution equals the gradient of f at the input point, scaled by the embedding norm
  • D. The attribution of any token not in the baseline vocabulary is exactly zero

Survives elimination: B

Why: Completeness (also called 'efficiency') states that Σᵢ IG(i) = f(x) − f(x_baseline). This is provable from the fundamental theorem of calculus applied to the path integral. It means attributions account for the full output difference with no unexplained residual. Neither gradient × input nor raw attention satisfies this. Attributions can be negative (token pushed output away from the prediction) so A is false.

44. Check 3: integrated gradients axiom

Check

Know your attribution axioms for the exam.

Check your understanding

Integrated gradients satisfies the completeness axiom. What does completeness guarantee?

  • A. Attributions are always non-negative, so each token's contribution is interpretable as a probability
  • B. The sum of all token attributions equals f(input) − f(baseline) (correct)
  • C. Each attribution equals the gradient of f at the input point, scaled by the embedding norm
  • D. The attribution of any token not in the baseline vocabulary is exactly zero

Answer: B

Why: Completeness (also called 'efficiency') states that Σᵢ IG(i) = f(x) − f(x_baseline). This is provable from the fundamental theorem of calculus applied to the path integral. It means attributions account for the full output difference with no unexplained residual. Neither gradient × input nor raw attention satisfies this. Attributions can be negative (token pushed output away from the prediction) so A is false.

Why A tempts people
IG attributions can be negative — a token may actively suppress a class logit. Requiring non-negativity would violate the completeness sum (since f(x) − f(baseline) can be positive while some individual terms are negative).
Why C tempts people
That describes gradient × input saliency, not integrated gradients. IG averages gradients along the interpolation path from baseline to input, not just the gradient at the endpoint.
Why D tempts people
IG makes no statement about vocabulary membership. The baseline is a zero embedding (or pad token), not a vocabulary restriction. Tokens present in both input and baseline can still receive non-zero attribution.

45. Your turn: attention inspector

Section

Project

46. Project brief

Concept

Build an AttentionInspector class that wraps a tiny 2-layer transformer and exposes three analysis tools: (1) per-head attention heatmaps, (2) attention rollout, and (3) gradient × input token attribution.

Milestone 1
Toy transformer forward pass — extract per-head attention tensors
Milestone 2
Attention rollout across both layers
Milestone 3
Gradient × input token attribution with a single backward pass
Show it off
Compare heatmap vs rollout vs gradient attribution for the same input

47. Which is which: Project brief

Matching

Match the pairs

From Project brief — match each one to what it actually does. The descriptions have been shuffled.

  • c1. Milestone 1
  • c2. Milestone 2
  • c3. Milestone 3
  • c4. Show it off
  • b1. Toy transformer forward pass — extract per-head attention tensors
  • b2. Attention rollout across both layers
  • b3. Gradient × input token attribution with a single backward pass
  • b4. Compare heatmap vs rollout vs gradient attribution for the same input

Why: Milestone 1, Milestone 2, Milestone 3, Show it off are easy to tell apart while they are sitting next to their descriptions and much harder afterwards, which is what this checks.

48. Milestone 1: toy transformer + attention extraction

Worked example

Build a minimal 2-layer transformer with nn.MultiheadAttention, capture attention weights via the need_weights=True flag

Why: nn.MultiheadAttention returns (output, attn_weights) when need_weights=True and average_attn_weights=False gives per-head weights. This is the standard hook point.

import torch, torch.nn as nn

class TinyTransformer(nn.Module):
    def __init__(self, d=16, n_heads=2, T=6):
        super().__init__()
        self.attn1 = nn.MultiheadAttention(d, n_heads, batch_first=True)
        self.attn2 = nn.MultiheadAttention(d, n_heads, batch_first=True)
        self.attn_maps = []

    def forward(self, x):
        self.attn_maps = []
        x2, w1 = self.attn1(x, x, x,
                   need_weights=True, average_attn_weights=False)
        self.attn_maps.append(w1)   # (B, n_heads, T, T)
        x3, w2 = self.attn2(x + x2, x + x2, x + x2,
                   need_weights=True, average_attn_weights=False)
        self.attn_maps.append(w2)
        return x3

torch.manual_seed(42)
model = TinyTransformer(d=16, n_heads=2, T=6)
x = torch.randn(1, 6, 16)
out = model(x)
print('Output shape:', out.shape)         # (1, 6, 16)
print('Attn layer-1 shape:', model.attn_maps[0].shape)  # (1, 2, 6, 6)
tensorshapeaxes
x input(1, 6, 16)batch, T, d_model
out(1, 6, 16)batch, T, d_model
attn_maps[0](1, 2, 6, 6)batch, n_heads, query-T, key-T
attn_maps[1](1, 2, 6, 6)batch, n_heads, query-T, key-T

49. Which is which, by shape

Discrimination

Sort into buckets

Sort these by shape, from memory, without looking back at Milestone 1: toy transformer + attention…. Telling them apart on the spot is the skill; the table is only where the answer happens to be written down.

(1, 6, 16)
x input; out
(1, 2, 6, 6)
attn_maps[0]; attn_maps[1]
g1
shape is "(1, 6, 16)" for x input, out — that is what the table on "Milestone 1: toy transformer +…" records, and it is the single property separating this group from the rest.
g2
shape is "(1, 2, 6, 6)" for attn_maps[0], attn_maps[1] — that is what the table on "Milestone 1: toy transformer +…" records, and it is the single property separating this group from the rest.

50. Finish it with less help: Milestone 2: rollout across 2 layers

Faded example

Fill in the blanks

Milestone 2: rollout across 2 layers, with the scaffolding fading: two lines are gone now — fill both.

import torch, torch.nn.functional as F

def rollout(attn_maps):
# attn_maps: list of (B, n_heads, T, T)
result = None
for w in attn_maps:
A = w[0].mean(0) # (T, T): average over heads
A_res = (A + torch.eye(A.shape[0])) / 2
A_res = A_res / A_res.sum(-1, keepdim=True)
if result is None:
result = A_res
else:
result = A_res @ result
return result # (T, T)

R = rollout(model.attn_maps) # model from Milestone 1
print('Rollout shape:', R.shape) # (6, 6)
print('Row sums:', R.sum(-1).detach().numpy().round(4))
print('Row 0 (token-0 attribution):', R[0].detach().numpy().round(3))

Why: Reproducing these unaided, rather than reading them, is what tells you the method has transferred. Multiple heads are averaged to get a single (T×T) map per layer before rollout — a standard simplification that treats all heads as equally important.

51. Milestone 2: rollout across 2 layers

Worked example

Average multi-head attention maps, then apply residual-augmented rollout

Why: Multiple heads are averaged to get a single (T×T) map per layer before rollout — a standard simplification that treats all heads as equally important.

import torch, torch.nn.functional as F

def rollout(attn_maps):
    # attn_maps: list of (B, n_heads, T, T)
    result = None
    for w in attn_maps:
        A = w[0].mean(0)  # (T, T): average over heads
        A_res = (A + torch.eye(A.shape[0])) / 2
        A_res = A_res / A_res.sum(-1, keepdim=True)
        if result is None:
            result = A_res
        else:
            result = A_res @ result
    return result  # (T, T)

R = rollout(model.attn_maps)  # model from Milestone 1
print('Rollout shape:', R.shape)   # (6, 6)
print('Row sums:', R.sum(-1).detach().numpy().round(4))
print('Row 0 (token-0 attribution):', R[0].detach().numpy().round(3))
metricvaluesanity check
Rollout shape(6, 6)T×T as expected
Row sumsall ≈ 1.0000valid probability rows
R[0] max entrydiagonal (self)residual dominates at random init

52. Predict the next row: Milestone 3: gradient × input attribution

Pattern

Predict first

The table runs: 0 | computed | ~1/6 · 1 | computed | ~1/6 · 2 | computed | ~1/6 · 3 | computed | ~1/6 · 4 | computed | ~1/6

In Milestone 3: gradient × input attribution, given the rows so far: what is the next one — the row where token idx is 5?

Correct: 5 | computed | ~1/6

token idxraw attrnorm attr
0computed~1/6
1computed~1/6
2computed~1/6
3computed~1/6
4computed~1/6
5computed~1/6

Why: The relationship between the columns, not the individual numbers, is what generates the next row. We treat the mean of the output as a proxy scalar.

53. Milestone 3: gradient × input attribution

Worked example

Enable gradients on the input embedding, run forward, call backward on a scalar target logit

Why: We treat the mean of the output as a proxy scalar. In a real setting you would use the logit for the target class. The backward pass computes ∂logit/∂embed for every token in one pass.

import torch
torch.manual_seed(42)
model2 = TinyTransformer(d=16, n_heads=2, T=6)

x_attr = torch.randn(1, 6, 16, requires_grad=True)
out_attr = model2(x_attr)           # (1, 6, 16)
target_logit = out_attr.mean()      # scalar proxy
target_logit.backward()

grad = x_attr.grad                  # (1, 6, 16)
attr = (grad * x_attr.detach()).abs().sum(-1)[0]  # (6,)
attr_norm = attr / attr.sum()
print('Attribution per token:',  attr.detach().numpy().round(4))
print('Normalized attribution:', attr_norm.detach().numpy().round(4))
print('Most influential token:', int(attr.argmax()))
token idxraw attrnorm attr
0computed~1/6
1computed~1/6
2computed~1/6
3computed~1/6
4computed~1/6
5computed~1/6

At random initialization, all tokens contribute roughly equally — attribution becomes meaningful only after training on a real task

Why: This is expected behavior. Run the same attribution code on a fine-tuned model with a concrete input sentence and you will see sharp, semantically meaningful attribution peaks.

54. What each one costs: Milestone 3: gradient × input attribution

Trade off

Comparison matrix

From Milestone 3: gradient × input attribution: every row here is a choice with a cost. Fill the norm attr column, then say which row you would actually pick and what you give up for it.

token idxraw attrnorm attr
0computed~1/6
1computed~1/6
2computed~1/6
3computed~1/6
4computed~1/6
5computed~1/6

55. Restore the missing line: Show it off: compare all three tools

Fill the middle

Fill in the blanks

From Show it off: compare all three tools — one line has had its right-hand side removed. Put it back.

# Run after Milestones 1-3 above
import torch
torch.manual_seed(42)
model3 = TinyTransformer(d=16, n_heads=2, T=6)
x_full = torch.randn(1, 6, 16, requires_grad=True)
out_full = model3(x_full)
out_full.mean().backward()

# 1. Head-0 of layer-1 heatmap (row 2)
head0_row2 = model3.attn_maps[0][0, 0, 2].detach().numpy().round(3)

# 2. Rollout row 2
R2 = rollout(model3.attn_maps)
rollout_row2 = R2[2].detach().numpy().round(3)

# 3. Gradient x input token 2
g = x_full.grad
attr_t2 = (g * x_full.detach()).abs().sum(-1)[0].detach().numpy().round(3)

print('Head-0 L1 attn (row 2):', head0_row2)
print('Rollout attn (row 2):', rollout_row2)
print('Grad x input (all T):', attr_t2)

Why: rollout_row2 is what everything below it consumes, so the wrong expression here fails later and somewhere else. Heatmap shows layer-local patterns; rollout shows propagated input-space attribution; gradient × input shows output-sensitivity.

56. Show it off: compare all three tools

Worked example

Print all three attribution views side-by-side to see what each technique reveals differently

Why: Heatmap shows layer-local patterns; rollout shows propagated input-space attribution; gradient × input shows output-sensitivity. Each answers a different question about the model.

# Run after Milestones 1-3 above
import torch
torch.manual_seed(42)
model3 = TinyTransformer(d=16, n_heads=2, T=6)
x_full = torch.randn(1, 6, 16, requires_grad=True)
out_full = model3(x_full)
out_full.mean().backward()

# 1. Head-0 of layer-1 heatmap (row 2)
head0_row2 = model3.attn_maps[0][0, 0, 2].detach().numpy().round(3)

# 2. Rollout row 2
R2 = rollout(model3.attn_maps)
rollout_row2 = R2[2].detach().numpy().round(3)

# 3. Gradient x input token 2
g = x_full.grad
attr_t2 = (g * x_full.detach()).abs().sum(-1)[0].detach().numpy().round(3)

print('Head-0 L1 attn  (row 2):', head0_row2)
print('Rollout attn    (row 2):', rollout_row2)
print('Grad x input    (all T):', attr_t2)
methodscopequestion answered
attention heatmapsingle head, single layerWhere does this head focus at layer L?
attention rolloutall heads, all layersWhich input tokens drove this output token?
gradient × inputall layers, output-centricWhich tokens most changed the final logit?

57. Fill in: question answered for Show it off: compare all three tools

Comparison

Comparison matrix

From Show it off: compare all three tools: refill the question answered column from what you know. The rest of the table is as it appeared.

methodscopequestion answered
attention heatmapsingle head, single layerWhere does this head focus at layer L?
attention rolloutall heads, all layersWhich input tokens drove this output token?
gradient × inputall layers, output-centricWhich tokens most changed the final logit?

58. Connect it up: Lesson 87: Transformer Interpretability

Connect it up

Draw it

One page, no notation unless you need it: draw how these connect — Attention heatmaps · Attention rollout · Linear probing classifiers · Gradient-based attribution · Your turn: attention inspector. Put an arrow wherever one of them is what makes another possible, and label the arrow with why.

59. Lesson 87 recap

Recap

Sources

  1. USAAIO Year-Long Master Lesson Plan, Lesson 87 — Transformer Interpretability — Barron · USAAIO Round 2 Preparation, 2026
  2. Attention rollout, IG, and probing classifier verified with torch 2.7.1+cpu, numpy 2.2.6, sklearn; all numbers from real execution, June 2026 — Real execution, verified
  3. Samira Abnar & Willem Zuidema — Quantifying Attention Flow in Transformers (2020) — arXiv 2005.00928
  4. Sundararajan, Taly, Yan — Axiomatic Attribution for Deep Networks (Integrated Gradients, 2017) — ICML 2017

Want this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.

Book on Wyzant · Text (657) 465-8108