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
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.
Objectives
Section
Part 1 of 4
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.
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.
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 \ key | The | cat | sat | mat |
|---|---|---|---|---|
| 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 |
| mat | 0.0273 | 0.0551 | 0.2468 | 0.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).
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 \ key | The | cat | sat | mat |
|---|---|---|---|---|
| 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 |
| mat | 0.0273 | 0.0551 | 0.2468 | 0.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.
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 \ key | The | cat | sat | mat |
|---|---|---|---|---|
| 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 |
| mat | 0.0273 | 0.0551 | 0.2468 | 0.6708 |
Concept
Different heads learn different roles during pretraining. Empirical studies on BERT and GPT find consistent head types across models.
| head type | what it attends to | diagnostic signal |
|---|---|---|
| previous token | always position i-1 | strict sub-diagonal stripe |
| syntactic (subject-verb) | subject ↔ verb pairs | off-diagonal blocks near dependency arcs |
| coreference | pronoun → antecedent | sparse long-range weights |
| delimiter / [SEP] | mostly [SEP] token | single bright column |
| positional (fixed offset) | always +k positions | diagonal 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.
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 type | what it attends to | diagnostic signal |
|---|---|---|
| previous token | always position i-1 | strict sub-diagonal stripe |
| syntactic (subject-verb) | subject ↔ verb pairs | off-diagonal blocks near dependency arcs |
| coreference | pronoun → antecedent | sparse long-range weights |
| delimiter / [SEP] | mostly [SEP] token | single bright column |
| positional (fixed offset) | always +k positions | diagonal shifted by k |
Section
Part 2 of 4
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)} \]
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.
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 token | input tok 0 | input tok 1 | input tok 2 | input tok 3 |
|---|---|---|---|---|
| tok 0 | 0.4419 | 0.1450 | 0.2364 | 0.1767 |
| tok 1 | 0.1843 | 0.4207 | 0.1865 | 0.2084 |
| tok 2 | 0.1546 | 0.1938 | 0.4698 | 0.1819 |
| tok 3 | 0.1863 | 0.1520 | 0.2171 | 0.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.
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.
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.
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.
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.
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.
Section
Part 3 of 4
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 finding | layer where it peaks |
|---|---|---|
| POS tags (noun/verb/adj) | 97% accuracy | layers 2–4 |
| dependency role (subj/obj) | 88% accuracy | layers 5–8 |
| coreference | 74% accuracy | layers 8–11 |
| semantic role (agent/patient) | 65% accuracy | layers 9–12 |
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.
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.
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| split | samples | accuracy | coeff dim-0 |
|---|---|---|---|
| train | 150 | — | −2.4577 |
| test | 50 | 0.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%.
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.
Section
Part 4 of 4
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).
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.
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 idx | raw attr | normalized attr | most influential? |
|---|---|---|---|
| 0 | 0.6824 | 0.2258 | no |
| 1 | 0.3671 | 0.1215 | no |
| 2 | 0.6434 | 0.2129 | no |
| 3 | 0.7319 | 0.2422 | YES (max) |
| 4 | 0.5967 | 0.1975 | no |
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.
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 idx | raw attr | normalized attr | most influential? |
|---|---|---|---|
| 0 | 0.6824 | 0.2258 | no |
| 1 | 0.3671 | 0.1215 | no |
| 2 | 0.6434 | 0.2129 | no |
| 3 | 0.7319 | 0.2422 | YES (max) |
| 4 | 0.5967 | 0.1975 | no |
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.
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.
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 idx | IG raw | IG normalized | interpretation |
|---|---|---|---|
| 0 | 0.4529 | 0.2452 | minor contributor |
| 1 | 0.8535 | 0.4620 | dominant (46%) |
| 2 | 0.5409 | 0.2928 | secondary |
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.
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.
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 range | representation character | what probes find |
|---|---|---|
| 1–3 (early) | local context, syntax | POS tags, NER span boundaries |
| 4–7 (mid) | phrase-level structure | dependency roles, chunking |
| 8–12 (late) | semantic, global | coreference, 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.
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.
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:
A = softmax(QKᵀ / √d_k), visualize as (T×T) matrix per head; look for diagonal (self-attend), sub-diagonal (prev-token), or…Ã = 0.5·A + 0.5·I, re-normalize, then R = Ã_L · … · Ã_1; row i of R is token i's attribution to…L, train LogisticRegression; probe accuracy > chance means the target property P is…attr(i) = ‖∇_{e_i} f ⊙ e_i‖₁ (single backward pass); use integrated gradients (n_steps ≥ 50) to satisfy the completeness…Pattern
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à = 0.5·A + 0.5·I, re-normalize, then R = Ã_L · … · Ã_1; row i of R is token i's attribution to each input tokenL, train LogisticRegression; probe accuracy > chance means the target property P is linearly encoded at depth Lattr(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 spotsEdge 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:
A = softmax(QKᵀ / √d_k), visualize as (T×T) matrix per head; look for diagonal (self-attend), sub-diagonal (prev-token), or…Ã = 0.5·A + 0.5·I, re-normalize, then R = Ã_L · … · Ã_1; row i of R is token i's attribution to…L, train LogisticRegression; probe accuracy > chance means the target property P is…attr(i) = ‖∇_{e_i} f ⊙ e_i‖₁ (single backward pass); use integrated gradients (n_steps ≥ 50) to satisfy the completeness…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.
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.
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?
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.
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).
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?
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).
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.
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.
Check
Know your attribution axioms for the exam.
Check your understanding
Integrated gradients satisfies the completeness axiom. What does completeness guarantee?
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.
Section
Project
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.
Matching
Match the pairs
From Project brief — match each one to what it actually does. The descriptions have been shuffled.
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.
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)| tensor | shape | axes |
|---|---|---|
| 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 |
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.
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.
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))| metric | value | sanity check |
|---|---|---|
| Rollout shape | (6, 6) | T×T as expected |
| Row sums | all ≈ 1.0000 | valid probability rows |
| R[0] max entry | diagonal (self) | residual dominates at random init |
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 idx | raw attr | norm attr |
|---|---|---|
| 0 | computed | ~1/6 |
| 1 | computed | ~1/6 |
| 2 | computed | ~1/6 |
| 3 | computed | ~1/6 |
| 4 | computed | ~1/6 |
| 5 | computed | ~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.
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 idx | raw attr | norm attr |
|---|---|---|
| 0 | computed | ~1/6 |
| 1 | computed | ~1/6 |
| 2 | computed | ~1/6 |
| 3 | computed | ~1/6 |
| 4 | computed | ~1/6 |
| 5 | computed | ~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.
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 idx | raw attr | norm attr |
|---|---|---|
| 0 | computed | ~1/6 |
| 1 | computed | ~1/6 |
| 2 | computed | ~1/6 |
| 3 | computed | ~1/6 |
| 4 | computed | ~1/6 |
| 5 | computed | ~1/6 |
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.
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)| method | scope | question answered |
|---|---|---|
| attention heatmap | single head, single layer | Where does this head focus at layer L? |
| attention rollout | all heads, all layers | Which input tokens drove this output token? |
| gradient × input | all layers, output-centric | Which tokens most changed the final logit? |
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.
| method | scope | question answered |
|---|---|---|
| attention heatmap | single head, single layer | Where does this head focus at layer L? |
| attention rollout | all heads, all layers | Which input tokens drove this output token? |
| gradient × input | all layers, output-centric | Which tokens most changed the final logit? |
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.
Recap
A = softmax(QKᵀ/√d_k) is a (T×T) distribution; visualize per head to spot positional, syntactic, or coreference patternsà = 0.5A + 0.5I across all layers to trace input-space attribution; row sums stay 1.0LogisticRegression on layer-L reps; accuracy > chance = property linearly encoded at depth L (92% on POS in this lesson)‖∇_{e_i}f ⊙ e_i‖₁ per token; integrated gradients (50-step Riemann, completeness axiom: Σ IG = f(x) − f(baseline)) is more robust at saturated gradientsWant this taught 1-on-1? Alexander tutors Machine Learning — $55/session, free consultation.