primer.ml.interpretability
Looking inside the model
Run: python -m primer.ml.interpretability
New to the notation? primer.notation explains every symbol used here from
zero. This lesson builds on the residual stream of primer.ml.transformer
and on behavioural evaluation in primer.agents.evals.
Level 1: The practitioner's guide
In one sentence. Interpretability is the set of tools that read and change a model's internal activations to find out what it represents and which of those representations cause its answer, as opposed to evaluations, which only measure what it says.
When you need it. Behavioural testing answers "what does the model do?", and for most shipping decisions that is enough. You need to open the hood when the question is "why?": a failure you cannot reproduce from the outside, a model that may be right for the wrong reason (the classifier that detects wolves by the snow behind them), a claim that the model is relying on a protected attribute, or a feature you want to turn up or down without retraining. The tell is that you are about to explain a model's behaviour from its outputs alone and cannot tell two explanations apart. This lesson's toy shows why that is dangerous: a probe reads "tense" out of the hidden state at 98% accuracy on unseen examples, and flipping tense moves the model's output by exactly 0.00, while flipping sentiment moves it by 2.00. What is readable inside a model and what the model uses are different questions, and only an intervention answers the second. One practical limit before anything else: every tool here needs the activations of a model you run yourself. A hosted API gives you tokens, not hidden states, so for a model behind an API your instrument is still the evaluation.
Your options. From the cheapest look to the strongest evidence:
| Option | What it does | What it tells you | What it costs | Where it lives |
|---|---|---|---|---|
| Behavioural evaluation | Runs the model on your tasks and scores the outputs | What the model does, at scale, for any model including hosted ones | A golden set and a scoring rule | primer.agents.evals |
| Linear probe | Trains a tiny classifier on frozen hidden states to read one property | That the property is linearly present at that layer, if it beats a control task on unseen data | Labelled examples (100 sufficed in this lesson), minutes of training | Your code, on an open-weight model |
| Logit lens (and tuned lens) | Applies the model's own output layer to the residual stream after each layer | When and where a prediction forms inside the model; no training needed | Nothing beyond a forward pass; a small trained translator per layer for the tuned version | Any residual-stream model you can hook |
| Activation patching (causal tracing) | Copies one activation from a clean run into a corrupted run and measures how much of the answer returns | Which sites cause the answer: the minimum for a causal claim | Two prompts that differ in one fact, and one forward pass per site tested | Hooking libraries such as TransformerLens |
| Sparse autoencoder (SAE) features | Learns an overcomplete dictionary so each hidden state is a few interpretable directions | What the model's units of meaning are, in a form you can read and steer | A large training run on activations, a penalty to tune, and 11% of the variance unexplained in this lesson's toy | Released SAE suites for open models, or your own training |
| Feature steering | Adds a found direction to the residual stream at run time | Whether a feature causes what its name suggests, and a lever without retraining | The feature must exist first; too strong a push degrades the output | Research demos on production models |
| Circuit analysis | Patches site by site until every step from input to output is accounted for | A complete mechanism for one behaviour | Weeks of expert time for a small model | Research |
How to choose. Start from the question, and reach for the cheapest tool that answers it.
- "Does the model know X?": a probe, scored on held-out examples against a random-label control. In this lesson the control fits 73% of its training set and scores 50% unseen, which is what a probe that only memorised looks like.
- "When does it decide?": the logit lens. In the landmark model the fact appears at the subject word after layer 1 and reaches the last word only after layer 2.
- "Does this part cause the answer?": patching. Reading is not enough: the lens reads "Paris" at 0.995 at a position where patching restores 0%.
- "What are the features, and can I steer one?": an SAE, followed by a steering experiment to check that the feature means what its label says.
- "Is the model relying on the wrong thing in production?": a probe or SAE feature to find the candidate, then patching or steering to confirm it, then a behavioural evaluation to measure the effect at scale.
- Whatever you pick, place the result on the ladder from correlation (probes, the lens) through intervention (patching, steering) to mechanism (a circuit), and claim only the rung you reached.
What it costs. Probes and the logit lens cost almost nothing: frozen activations, a forward pass, and for a probe a few hundred labelled examples and seconds of training. Patching costs one forward pass per site, so a full map is layers times positions; the lesson's is 3 by 3, a real model's is hundreds by thousands, and it is repeated for every prompt pair you try. Sparse autoencoders are the expensive tool: millions of activations, a dictionary far wider than the layer, and a sparsity penalty whose choice is a trade. In this lesson's sweep, a penalty of 0.3 matches every planted feature above 0.99 and leaves 11% of the variance unexplained; a penalty of 0.01 rebuilds everything and matches the worst feature at only 0.80, so a perfect rebuild score is not evidence that the features are real. Circuit-level explanations cost research time measured in weeks per behaviour. And all of it presumes access: none of these tools runs on a model you only reach through an API.
What breaks.
- Decodable is not used. The 98%-readable feature with zero effect on the output. Never claim "the model uses X" from a probe; intervene.
- A probe that is too clever. A deep probe can compute the property itself; a probe with no control task can fit noise. Keep probes linear, score them on unseen data, compare with random labels.
- Reading a site nothing reads. The lens can see an answer at a position no later layer consults. Patching restores 0% there.
- A choice that shapes the answer. Patching results depend on how the corrupted prompt is built; SAE features depend on the dictionary's size and penalty, and a bigger dictionary can split one feature into several.
- Redundancy. Models can partly repair themselves when one component is knocked out, so "patching this restores nothing" does not always mean "this plays no role".
- The unexplained remainder. The variance an SAE does not rebuild is model behaviour no feature describes yet, not noise.
- Labels that are guesses. A feature named "Golden Gate Bridge" is a summary of what makes it fire. The name is checked by steering, not by reading.
- Polysemantic neurons. When features outnumber neurons they cannot line up with them, so a single neuron responds to several unrelated things. Look along directions, not at neurons.
In the wild. TransformerLens exposes the internal activations of thousands of open-weight models and lets you cache, edit and replace them as the model runs, which is the plumbing for the lens, probes and patching. Causal tracing (Meng et al., 2022) located factual recall in middle-layer MLPs at the subject's last token in GPT-style models and edited single facts there; Wang et al. (2022) patched their way to a complete circuit for indirect-object identification in GPT-2 small. The tuned lens (Belrose et al., 2023) fixed the logit lens on early layers. Anthropic's Towards Monosemanticity (2023) trained SAEs on a small transformer, and Scaling Monosemanticity (2024) did it inside Claude 3 Sonnet, where turning up one feature made the model bring up the Golden Gate Bridge in almost every answer. Google DeepMind's Gemma Scope releases trained SAEs for every layer of the Gemma 2 2B and 9B base models, so a practitioner can inspect features without training a dictionary.
Go deeper. Level 2 builds each tool on models small enough to check
by hand: a feature as a direction read back by a dot product, a probe
trained by gradient descent with a control task, the logit lens verified
against the real output on primer.ml.transformer.TinyGPT, activation
patching on a two-layer landmark model with a known circuit, superposition
in a five-features-in-two-neurons toy, and a sparse autoencoder that
recovers the planted features only when its penalty is right. If you only
needed to know what these tools can and cannot tell you about a model in
production, you are done.
Level 2: How it works, from scratch
Everyday picture. A car makes a strange noise. You can take it for a test drive and note when the noise happens: that is testing the car's behaviour. Or you can open the hood, put a stethoscope on the engine, and swap parts until the noise stops: that is looking at the mechanism. Test drives tell you what the car does. Only the open hood tells you why.
Every other lesson in this primer treats a trained model as something to
build, train or call. Evaluations (primer.agents.evals) are test drives:
they measure what the model says. Interpretability opens the hood. It
asks what the model's billions of internal numbers represent, and which of
them cause the answer. Three reasons to care:
- Debugging. When a model gets something wrong, you want to know which step failed, the same way you read a stack trace rather than just the error message.
- Trust. A model can be right for the wrong reason. An image classifier that "detects" wolves by looking for snow in the background scores well until it meets a wolf on grass.
- Safety. The output is only part of what a model computes. Some questions (is it relying on a stereotype? does it know more than it says?) can only be answered by reading the computation itself.
A tiny worked example. This lesson asks four questions of one kind of sentence, "the Eiffel Tower is in …" → "Paris", each with its own tool, on toy models small enough to check by hand:
| Question | Tool | What our toy shows |
|---|---|---|
| Is a property stored in this hidden state? | a probe | sentiment, read at 99% on unseen examples |
| What would the model say if it stopped at layer ℓ? | the logit lens | "London", then "Paris", then "Paris" |
| Which activations cause the answer? | activation patching | the subject word early, the last word late |
| What are the model's units of meaning? | superposition and sparse autoencoders | five features packed into two neurons, then recovered |
flowchart LR M["A trained model<br/>(frozen)"] --> A["Hidden activations<br/>at every layer and word"] A --> P["Probe:<br/>is property X in here?"] A --> L["Logit lens:<br/>what would it predict now?"] A --> AP["Activation patching:<br/>does this activation cause the answer?"] A --> S["Sparse autoencoder:<br/>which features make up this state?"] P & L --> R["Reading<br/>(correlation)"] AP --> C["Intervening<br/>(causation)"] S --> U["Units of analysis<br/>(what to read and intervene on)"]
Reading it: every tool starts from the same place: the numbers a frozen model computes inside itself while it runs. The top two tools only read those numbers, so what they find is a correlation. Patching changes them and watches the output, so what it finds is a cause. Sparse autoencoders answer an earlier question: before you can read or change "a feature", you need to know what the features are. Keep the reading/intervening split in mind; it is the most important distinction in this lesson.
A feature is a direction
Everyday picture. A band with two instruments plays through two speakers. The sound engineer pans the guitar mostly to the right speaker and the piano mostly to the left, but each speaker plays a mix of both. If you unplug one speaker you don't lose "the guitar"; you lose part of everything. The instruments are not the speakers. Each instrument is a setting across all speakers, a direction, and the sound in the room is the sum of every instrument times how loudly it plays.
A model's hidden state works the same way. The speakers are neurons: the individual numbers in a layer's output vector. The instruments are features: the things the model has learned to track, such as "this review is positive" or "this verb is in the past tense". Each feature is stored as a direction in the space of neurons, and the hidden state is the sum of every active feature's direction, scaled by how strongly it is present.
A tiny worked example. Take a hidden layer with two neurons and two features whose directions are "positive" = (0.6, 0.8) and "past tense" = (0.8, −0.6). A sentence that is quite positive (0.9) and a little past-tense (0.4) has the hidden state
0.9 · (0.6, 0.8) + 0.4 · (0.8, −0.6) = (0.54 + 0.32, 0.72 − 0.24) = (0.86, 0.48).
Neuron 1 reads 0.86 and neuron 2 reads 0.48. Neither number is "how
positive" or "how past-tense": each neuron is a blend of both features.
To get the features back, take the dot product (multiply matching
entries and add; see primer.notation) of the hidden state with each
direction:
- positive: 0.6 · 0.86 + 0.8 · 0.48 = 0.516 + 0.384 = 0.9
- past tense: 0.8 · 0.86 − 0.6 · 0.48 = 0.688 − 0.288 = 0.4
Level 3: the formula and its symbols
$$ h = \sum_{i=1}^{k} f_i \, d_i \qquad\qquad \hat{f}_i = d_i \cdot h $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $h$ | the hidden state: one number per neuron | (0.86, 0.48) |
| $k$ | how many features are present | 2 |
| $i$ | which feature | 1 = positive, 2 = past tense |
| $f_i$ | how strongly feature $i$ is present | $f_1 = 0.9$, $f_2 = 0.4$ |
| $d_i$ | feature $i$'s direction: a list with one number per neuron, of length 1 | $d_1 = (0.6, 0.8)$ |
| $\sum_{i=1}^{k}$ | add up one term per feature | two terms |
| $\hat{f}_i$ | the amount of feature $i$ we read back (the hat means "estimated") | 0.9 |
| $\cdot$ | dot product: multiply matching entries, then add |
In words: "the hidden state is the sum of each feature's direction times its strength; to read a feature back, dot the hidden state with that feature's direction."
With the numbers: $h = 0.9 \cdot (0.6, 0.8) + 0.4 \cdot (0.8, -0.6) = (0.86, 0.48)$, and $\hat{f}_1 = (0.6, 0.8) \cdot (0.86, 0.48) = 0.9$. The read-back is exact here because the two directions are perpendicular (their dot product is 0.6 · 0.8 + 0.8 · (−0.6) = 0) and each has length 1. Hold on to that condition: the section on superposition is about what happens when it fails.
Level 3: in Python
In Python:
positive = [0.6, 0.8]
past = [0.8, -0.6]
# h = Σ f_i d_i
h = [0.9 * p + 0.4 * q for p, q in zip(positive, past)]
[round(x, 2) for x in h] # → [0.86, 0.48]
# f̂_i = d_i · h
round(sum(d * x for d, x in zip(positive, h)), 2) # → 0.9
round(sum(d * x for d, x in zip(past, h)), 2) # → 0.4
# the directions are perpendicular: their dot product is 0
round(sum(p * q for p, q in zip(positive, past)), 2) # → 0.0
Reading it: the axes are the two neurons. The blue and green arrows are the two feature directions; the red arrow is the hidden state. The dotted lines drop from the hidden state onto each feature arrow, and where they land (0.9 of the way along blue, 0.4 along green) is the dot product: the amount of that feature. Notice that the red arrow's coordinates on the neuron axes, 0.86 and 0.48, are neither amount. To read a model you have to look along the right directions, not along the neurons.
That features are directions, and that most of what a model tracks can be
read with a dot product, is called the linear representation
hypothesis. It is a hypothesis, not a law, but it holds often enough to
power everything below. You have met it before: word-vector arithmetic
such as king − man + woman ≈ queen (primer.ml.embeddings.word2vec) works
because "royalty" and "gender" are directions.
In code: compose_features builds a hidden state from feature amounts and read_features reads them back with dot products.
Probes: can a straight line read it out?
Everyday picture. A doctor can't ask your liver how it is doing, but a blood test can measure a marker that tells them. A probe is a blood test for a hidden state: a tiny classifier, trained by you, that looks at a layer's activations and answers one yes/no question, such as "is this review positive?". The model itself is frozen; only the probe learns.
A tiny worked example. A probe for "positive" on a two-neuron layer has weights w = (1, 0.5) and bias b = 0. On the hidden state h = (2, −1):
- Score it: 1 · 2 + 0.5 · (−1) + 0 = 1.5.
- Squash the score into a probability with the sigmoid σ(z) = 1 / (1 + e^−z), which maps any number into the range 0 to 1: σ(1.5) = 1 / (1 + 0.223) = 0.818.
The probe is 82% sure this hidden state belongs to a positive review. On h = (−1, 1) the score is −1 + 0.5 = −0.5 and σ(−0.5) = 0.378: probably negative.
Level 3: the formula and its symbols
$$ p = \sigma(w \cdot h + b) = \frac{1}{1 + e^{-(w \cdot h + b)}} $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $h$ | the frozen model's hidden state for one example | (2, −1) |
| $w$ | the probe's weights: one per neuron, learned | (1, 0.5) |
| $b$ | the probe's bias: one number, learned | 0 |
| $w \cdot h + b$ | the probe's raw score (its logit) | 1.5 |
| $\sigma$ | the sigmoid: turns any score into a probability between 0 and 1 | |
| $e$ | Euler's number, about 2.718 | $e^{-1.5} = 0.223$ |
| $p$ | the probe's probability that the property is present | 0.818 |
In words: "dot the hidden state with the probe's weights, add the bias, and squash the result into a probability."
With the numbers: $p = \sigma(1 \cdot 2 + 0.5 \cdot (-1) + 0) = \sigma(1.5) = 1 / (1 + 0.223) = 0.818$.
Level 3: in Python
In Python:
import math
w = [1.0, 0.5]
b = 0.0
h = [2.0, -1.0]
# w · h + b
z = sum(wi * hi for wi, hi in zip(w, h)) + b
z # → 1.5
# σ(z)
round(1 / (1 + math.exp(-z)), 3) # → 0.818
# a second hidden state, (−1, 1), reads as probably negative
round(1 / (1 + math.exp(-(-1 * 1.0 + 1 * 0.5))), 3) # → 0.378
This is exactly logistic regression (primer.ml.neural_net), trained
the usual way: show it examples whose answer you know, measure its
binary cross-entropy (primer.ml.losses), and nudge w and b downhill. The
gradient of that loss with respect to the score is simply (p − y), the gap
between the probe's probability and the true 0/1 label, which makes each
update one line of code.
flowchart LR T["Text with a known label<br/>(positive or negative)"] --> M["Frozen model<br/>(no weights change)"] M --> H["Hidden state h<br/>at the layer under study"] H --> P["Probe: σ(w · h + b)"] P --> Y["Probability the label is 'positive'"] Y --> G["Compare with the true label<br/>(cross-entropy)"] G -. "gradient updates w and b only" .-> P
Reading it: the solid arrows are one forward pass; the dotted arrow is learning. The gradient stops at the probe: the model under study never changes, so whatever the probe finds was already in the hidden state. That is the point of a probe. It measures the model, not a new model trained on top of it.
On our toy model. PlantedModel stores two known features in a
32-neuron hidden layer: sentiment and tense, each as ±1 along its own
random direction, plus random noise of size 0.4 on every neuron. No single
neuron holds either feature. A probe trained on just 100 examples reads
sentiment correctly on 99% of 1,000 examples it never saw.
Why keep probes simple, and check them against a control. A probe can succeed for the wrong reason. Give a probe random labels, a control task with nothing real to find, and it still fits 73% of its 100 training examples, because 33 adjustable numbers can memorize a lot of noise. On unseen examples it scores 50%, a coin flip. Two habits follow: always score a probe on examples it never saw, and compare it with a control task, so that "the probe works" means "the hidden state encodes this" rather than "the probe is clever". This is also why probes are kept linear: a powerful probe (a deep network) could compute the property from raw ingredients by itself, and then its success would tell you about the probe, not the model.
Decodable is not the same as used
Everyday picture. A library holds a book nobody ever borrows. Finding it on the shelf proves the library has it, not that it shaped anything any reader did.
A tiny worked example. Our toy model's output reads sentiment and, by construction, gives tense a weight of exactly zero. A probe still reads tense at 98% on unseen examples. Now intervene: take 200 examples, flip only one feature's label, keep everything else fixed, and measure how much the model's output moves.
| Flip | Probe accuracy for it | Output moves by |
|---|---|---|
| sentiment | 99% | 2.00 |
| tense | 98% | 0.00 |
Reading it: the left panel is what probes can read. Grey bars are accuracy on the probe's own training examples, blue bars on unseen ones, and the red dashed line is chance. Sentiment and tense are both clearly readable; random labels look readable on the training set and collapse to chance on unseen data, which is exactly what the control is for. The right panel is what the model uses: the change in its output when one feature is flipped. Tense is as readable as sentiment and has no effect at all. The two panels answer different questions, and only the right one is about cause.
A probe finding information does not mean the model uses it. To find out what the model uses you have to change something and watch the output, which is what the rest of this lesson does.
In code: train_probe fits a Probe by gradient descent on frozen hidden states from PlantedModel; probe_report trains the three probes, and flip_effect performs the intervention.
The logit lens: reading the model's mind mid-thought
Everyday picture. A writer keeps every draft of an article. Draft 1 says the landmark is "in a European capital", draft 2 says "probably Paris", the final says "Paris". Reading the drafts shows when the writer made up their mind. The logit lens reads a language model's drafts.
It works because of how a transformer is built (primer.ml.transformer).
Each word has a running vector, the residual stream. Every layer reads
it and adds a correction to it, rather than replacing it. At the very
end, the output layer (the unembedding) turns the final vector into
one score per word in the vocabulary. Because every layer writes into the
same stream, you can apply that final step early, after any layer, and ask
"what would the model predict if it stopped here?"
A tiny worked example. A vocabulary of three words, Paris, London and Rome, a residual stream of two numbers, and an unembedding that scores Paris = first number, London = second number, Rome = minus the first:
| After | Residual $h$ | Scores (Paris, London, Rome) | Probabilities | Top guess |
|---|---|---|---|---|
| the embedding | (0.2, 0.3) | (0.2, 0.3, −0.2) | (0.36, 0.40, 0.24) | London |
| layer 1, which adds (0.6, 0) | (0.8, 0.3) | (0.8, 0.3, −0.8) | (0.55, 0.34, 0.11) | Paris |
| layer 2, which adds (1.2, −0.3) | (2.0, 0.0) | (2.0, 0.0, −2.0) | (0.87, 0.12, 0.02) | Paris |
The first draft is a vague "some capital" that happens to lean London; layer 1 tips it to Paris; layer 2 commits.
Level 3: the formula and its symbols
$$ \text{lens}_\ell = \text{softmax}\big(W_U \, \text{LN}(h_\ell)\big) $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $\ell$ | which layer we stop after (0 = straight after the embedding) | 2 |
| $h_\ell$ | the residual stream after layer $\ell$, for one word | (2.0, 0.0) |
| $\text{LN}$ | the model's final normalization (primer.ml.transformer); our 2-number toy has none, so here LN(h) = h |
(2.0, 0.0) |
| $W_U$ | the unembedding: one row per vocabulary word, dotted with the state to give that word's score | rows (1, 0), (0, 1), (−1, 0) |
| $W_U \, \text{LN}(h_\ell)$ | the scores (logits), one per word | (2, 0, −2) |
| softmax | turns scores into probabilities that add up to 1 (primer.ml.attention) |
(0.87, 0.12, 0.02) |
| $\text{lens}_\ell$ | what the model would predict if it stopped after layer $\ell$ | Paris at 0.87 |
In words: "take the residual stream partway up, push it through the model's own final normalization and output layer, and read off the probabilities."
With the numbers: after layer 2, $W_U h_2 = (1 \cdot 2 + 0 \cdot 0,\; 0 \cdot 2 + 1 \cdot 0,\; -1 \cdot 2 + 0 \cdot 0) = (2, 0, -2)$; then $e^2 = 7.39$, $e^0 = 1$, $e^{-2} = 0.14$, total 8.52, so P(Paris) = 7.39 / 8.52 = 0.87.
Level 3: in Python
In Python:
import math
# rows: Paris, London, Rome
W_U = [[1, 0], [0, 1], [-1, 0]]
h = [0.2, 0.3]
writes = [[0.6, 0.0], [1.2, -0.3]]
# a residual layer adds its write to the stream
for w in writes:
h = [a + b for a, b in zip(h, w)]
[round(x, 2) for x in h] # → [2.0, 0.0]
# W_U h: one score per word
scores = [sum(u * x for u, x in zip(row, h)) for row in W_U]
[round(s, 2) for s in scores] # → [2.0, 0.0, -2.0]
# softmax
exps = [math.exp(s) for s in scores]
[round(e / sum(exps), 2) for e in exps] # → [0.87, 0.12, 0.02]
flowchart LR E["Embedding"] --> H0(("h₀")) --> L1["Layer 1<br/>adds its write"] --> H1(("h₁")) --> L2["Layer 2<br/>adds its write"] --> H2(("h₂")) --> OUT["Final norm + W_U<br/>(the real output)"] H0 -.-> LENS0["lens: norm + W_U<br/>London 0.40"] H1 -.-> LENS1["lens: norm + W_U<br/>Paris 0.55"] H2 -.-> LENS2["lens: norm + W_U<br/>Paris 0.87"]
Reading it: the solid line along the middle is the residual stream: it
flows from the embedding to the real output, and each layer adds to it.
The dotted taps hang the model's own output layer off the stream after
every layer. Nothing is trained; the lens borrows weights the model
already has. At the top layer the tap and the real output are the same
computation, so they must agree exactly, which is a good test that the
lens is wired correctly (on primer.ml.transformer.TinyGPT, they do).
On a model where we know the answer. LandmarkModel is a two-layer,
one-head transformer built by hand to complete "Eiffel is in" with
"Paris". Layer 1 is an MLP that looks up a fact at every word (landmark
in, city out); layer 2 is an attention head that lets the last word, "in",
find the landmark and copy its city.
flowchart LR T["Eiffel · is · in"] --> EMB["Embed each word"] EMB --> MLP["Layer 1: MLP at every word<br/>Eiffel → writes 'Paris'"] MLP --> ATT["Layer 2: attention head<br/>'in' looks for the landmark,<br/>copies its city"] ATT --> U["Output at 'in':<br/>Paris 3, Rome 0, London 0"]
Reading it: read left to right as the two steps of the answer. The fact ("the Eiffel Tower is in Paris") is looked up at the landmark's own position in layer 1. It is only moved to the last position, the one that predicts the next word, in layer 2. We built it this way on purpose, so every tool below can be checked against a known truth.
Reading it: the left panel is the worked example: at each layer, three bars for the three words, and the blue Paris bar grows from 0.36 to 0.87. The right panel runs the lens over the landmark model: rows are layers, columns are words, and each cell is the lens's probability of "Paris" (1/3 = 0.33 means no opinion among three cities). The Eiffel column goes dark after layer 1, when the MLP looks up the fact. The "in" column only goes dark after layer 2, when attention copies the city over. The lens has shown where and when the answer appears, without being told.
The logit lens was first described for GPT-2 in 2020, where middle layers already "guess" the next word surprisingly well. In some models the early layers read as nonsense through the final output layer, because they do not yet speak its "language"; the tuned lens fixes this by training a small translator for each layer before applying the output layer.
In code: logit_lens applies an output layer to any residual states, worked_lens runs the table above, residual_stream and tinygpt_logit_lens do the same for primer.ml.transformer.TinyGPT, and lens_map builds the grid for LandmarkModel.
Activation patching: which activations cause the answer?
Everyday picture. Two cars of the same model sit side by side: one starts, one doesn't. You move parts from the good car into the bad one, one part at a time, and try the ignition after each swap. The part whose swap makes the bad car start is the part that mattered. Activation patching (also called causal tracing) does this with a model's internal activations.
A tiny worked example. Run the landmark model twice:
- the clean prompt "Eiffel is in": scores Paris 3, Rome 0, so the logit difference Paris − Rome is +3;
- the corrupted prompt "Colosseum is in": Paris 0, Rome 3, so the difference is −3.
Now run the corrupted prompt again, but at one chosen place overwrite the activation with the one from the clean run, and see how much of the gap between −3 and +3 comes back:
- patch the MLP output at the landmark's position: the difference jumps back to +3, so 100% of the answer is restored;
- patch the state at "is": it stays at −3, 0% restored ("is" is the same word in both prompts, so its state carries nothing about the landmark).
Level 3: the formula and its symbols
$$ \text{restored} = \frac{LD_{\text{patched}} - LD_{\text{corrupt}}}{LD_{\text{clean}} - LD_{\text{corrupt}}} $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $LD$ | logit difference: the correct answer's score minus the wrong answer's score (Paris − Rome) | |
| $LD_{\text{clean}}$ | on the clean prompt, nothing patched | +3 |
| $LD_{\text{corrupt}}$ | on the corrupted prompt, nothing patched | −3 |
| $LD_{\text{patched}}$ | on the corrupted prompt, with one clean activation copied in | +3, −3, or anything between |
| restored | the fraction of the gap that one patch closes: 0 = nothing, 1 = everything | 1, 0 |
In words: "how far the patch moved the answer, as a share of the full distance from the corrupted answer to the clean one."
With the numbers: patching the landmark's MLP output gives (3 − (−3)) / (3 − (−3)) = 6 / 6 = 1. Patching "is" gives (−3 − (−3)) / 6 = 0. A patch that left the difference at 0 would give (0 − (−3)) / 6 = 0.5: half the answer restored.
Level 3: in Python
In Python:
clean, corrupt = 3.0, -3.0
# restored = (LD_patched − LD_corrupt) / (LD_clean − LD_corrupt)
def restored(patched):
return (patched - corrupt) / (clean - corrupt)
# the landmark's MLP output, patched in
restored(3.0) # → 1.0
# the state at "is", patched in
restored(-3.0) # → 0.0
# a patch that only reaches a tie
restored(0.0) # → 0.5
Why a difference of two scores, rather than the probability of Paris? Because the corrupted prompt differs from the clean one in exactly one fact, the difference measures exactly that fact, and it moves smoothly, where a probability can saturate near 0 or 1 and hide a change.
flowchart TB subgraph C["1. Clean run: 'Eiffel is in'"] c1["save every activation"] --> c2["Paris − Rome = +3"] end subgraph K["2. Corrupted run: 'Colosseum is in'"] k1["Paris − Rome = −3"] end subgraph P["3. Patched run: corrupted prompt,<br/>ONE activation taken from the clean run"] p1["Paris − Rome = ?"] end c1 -- "copy one activation" --> P P --> R["fraction restored =<br/>(? − (−3)) / (3 − (−3))"] K --> R C --> R
Reading it: there are three runs. The clean run is a donor: its activations are saved. The corrupted run sets the baseline. The patched run is the corrupted prompt with a single activation transplanted from the donor. Repeat the patched run once per layer and position and you get one "fraction restored" per site: a map of where the answer is carried.
Reading it: rows are the residual stream after each layer, columns are the three positions (clean word / corrupted word), and each cell is the fraction of the answer restored by patching that one state. The answer lives at the landmark early (the Eiffel column is 1 after the embedding and after the MLP) and at the last word late (the "in" column is 1 after attention). The diagonal hand-off between them is the two-step circuit we built, found by intervention alone.
Compare with the lens grid above. After layer 2, the lens reads "Paris" at the Eiffel position with probability 0.995, yet patching that state restores 0.00: after the last layer nothing reads that position any more. The information is there, and it no longer matters. Readable is not the same as used, again.
Why it matters in practice. This is how researchers located where GPT-style models store facts: causal tracing on prompts like ours found factual recall concentrated in middle-layer MLPs at the subject's last token, and then used that to edit single facts. The same method, one site at a time, has traced whole circuits, such as the one GPT-2 small uses to fill in "When Mary and John went to the store, John gave a drink to" → "Mary".
In code: LandmarkModel is the hand-built model and LandmarkModel.run accepts patches; fraction_restored is the formula and patching_map patches every layer and position in turn.
Superposition: more features than neurons
Everyday picture. Back to the band, but now five instruments play through two speakers. The engineer pans each instrument to its own position around the room: hard left, front right, back left, and so on. When one instrument plays alone, you can tell which one by where the sound comes from. When two play at once, the positions blur together and you might mistake the pair for a third instrument. The trick works because in this band, most of the time, only one instrument is playing.
Models face the same squeeze. There are far more concepts in the world than neurons in a layer, but in any one sentence almost all of them are absent: features are sparse. So models store more features than they have dimensions, as directions that are nearly perpendicular rather than exactly. That is superposition.
A tiny worked example. Put five features in two neurons, spread 72° apart like the points of a pentagon. Feature $k$'s direction is (cos 72k°, sin 72k°): feature 0 is (1, 0), feature 1 is (0.309, 0.951), and so on. Switch on feature 0 alone, at strength 1. The hidden state is (1, 0). Reading every feature back with a dot product gives
| Feature | Angle from feature 0 | Read-back (cos of the angle) | After adding −0.31 and ReLU |
|---|---|---|---|
| 0 | 0° | 1.000 | 0.69 |
| 1 | 72° | 0.309 | 0 |
| 2 | 144° | −0.809 | 0 |
| 3 | 216° | −0.809 | 0 |
| 4 | 288° | 0.309 | 0 |
Five directions can't all be perpendicular in two dimensions, so reading feature 0 leaks 0.309 into each neighbour: interference. The fix is a small negative bias and a ReLU (which turns negatives into zero): the leak of 0.309 falls below the bias of 0.31 and is filtered to 0, and the real feature survives at 0.69. The cost comes when two neighbours are on at once: each then reads 1 + 0.309 − 0.31 = 0.999 instead of 0.69, too high. Sparsity is a bet that such collisions are rare.
Level 3: the formula and its symbols
$$ \hat{x} = \text{ReLU}\big(W^\top W x + b\big), \qquad \hat{x}_i = \text{ReLU}\Big(\|w_i\|^2 x_i + \sum_{j \neq i} (w_i \cdot w_j)\, x_j + b_i\Big) $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $x$ | the true features: one strength per feature | (1, 0, 0, 0, 0) |
| $W$ | the squeeze: 2 rows (neurons) by 5 columns (features); column $i$ is feature $i$'s direction $w_i$ | the pentagon |
| $W x$ | the hidden state: 2 numbers holding 5 features | (1, 0) |
| $W^\top$ | $W$ transposed (rows and columns swapped): reads each feature back with a dot product | |
| $b$ | one bias per feature, learned; negative values filter small leaks | −0.31 each |
| ReLU | keep positives, turn negatives into 0 | |
| $\hat{x}$ | the features read back out | (0.69, 0, 0, 0, 0) |
| $|w_i|^2$ | feature $i$'s direction dotted with itself (its squared length) | 1 |
| $w_i \cdot w_j$ | how much feature $j$ leaks into feature $i$: 0 if perpendicular | 0.309 for neighbours |
| $\sum_{j \neq i}$ | add up over every other feature $j$ |
In words: "squeeze the features into a few neurons, read each one back by dotting with its direction, subtract a threshold and drop anything negative. What you read for feature $i$ is its own strength, plus a leak from every other active feature, minus the threshold."
With the numbers: with feature 0 alone on, $\hat{x}_1 = \text{ReLU}(1 \cdot 0 + 0.309 \cdot 1 - 0.31) = \text{ReLU}(-0.001) = 0$ and $\hat{x}_0 = \text{ReLU}(1 \cdot 1 - 0.31) = 0.69$. With features 0 and 1 both on, $\hat{x}_0 = \text{ReLU}(1 + 0.309 - 0.31) = 0.999$.
Level 3: in Python
In Python:
import math
# feature k's direction: (cos 72k°, sin 72k°)
W = [[math.cos(math.radians(72 * k)), math.sin(math.radians(72 * k))] for k in range(5)]
def read_back(x, bias, relu=True):
# h = W x: two numbers
h = [sum(W[k][n] * x[k] for k in range(5)) for n in range(2)]
# Wᵀ h + b: one number per feature
z = [W[k][0] * h[0] + W[k][1] * h[1] + bias for k in range(5)]
# ReLU keeps positives and turns negatives into 0
return [round(max(0.0, v) if relu else v, 3) for v in z]
# the leak: feature 0 alone, no bias, no ReLU
read_back([1, 0, 0, 0, 0], bias=0, relu=False) # → [1.0, 0.309, -0.809, -0.809, 0.309]
# the bias and ReLU filter it
read_back([1, 0, 0, 0, 0], bias=-0.31) # → [0.69, 0.0, 0.0, 0.0, 0.0]
# two neighbours on: each reads too high
read_back([1, 1, 0, 0, 0], bias=-0.31) # → [0.999, 0.999, 0.0, 0.0, 0.0]
flowchart LR X["5 features<br/>x₀ … x₄<br/>(mostly zero)"] --> W["squeeze: W<br/>(2 × 5)"] W --> H["hidden state<br/>2 neurons"] H --> WT["read back: Wᵀ<br/>(5 × 2)"] WT --> B["+ bias b<br/>(negative)"] B --> R["ReLU"] R --> XH["5 features<br/>read back"]
Reading it: follow a feature through the bottleneck. Five numbers go in, are squeezed into two, and must be expanded back into five. There is no way to fit five perpendicular directions into two neurons, so the read-back always leaks; the bias and ReLU at the end are what make the leak survivable, by throwing away small readings. This is the toy model from Anthropic's Toy Models of Superposition (2022), and the next step is to train it and see what it chooses.
Training it. Let the model learn $W$ and $b$ itself, by gradient descent on the reconstruction error, with features that matter less and less (feature $i$ is weighted $0.8^i$). Compare two worlds:
- Dense: every feature is on in every example. The model keeps the two most important features, perpendicular, at length 1.00, and gives the other three length 0: it simply drops them.
- Sparse: each feature is on only 5% of the time. The model keeps all five, 72° apart, at length about 1.1, and learns a negative bias of about −0.23 to filter the leaks.
Reading it: in the two left panels each arrow is one feature's learned direction in the two-neuron space, and its length is how well the model stores it. Dense features (left) get the textbook answer: as many features as neurons, perpendicular, the rest discarded. Sparse features (middle) get a pentagon: five features in two dimensions, accepting a little interference in exchange for storing everything. The right panel reads the ideal pentagon by neuron: neuron 1 moves for feature 0 (by 1.0), feature 1 and feature 4 (by 0.309 each), and moves the other way for features 2 and 3. No neuron belongs to one feature.
That last point is why looking at single neurons so often fails. A neuron that responds to several unrelated features is polysemantic. Vision researchers found neurons that respond to cat faces and to the fronts of cars; language models are full of neurons like that. Superposition explains why: when features outnumber neurons, the features cannot line up with the neurons, so every neuron is a mixture.
In code: pentagon builds the five directions, superposition_readout is the formula, train_superposition learns W and b with gradients from superposition_loss_and_grads (checked against finite_difference_gradient), and features_a_neuron_responds_to lists a neuron's features.
Sparse autoencoders: getting the features back
Everyday picture. A sound engineer receives the two-speaker recording of the five-instrument band, with no notes on who played when. She knows one thing about this band: usually only one instrument plays at a time. Many different scores could produce the same sound, but she writes down the one that uses the fewest instruments. That preference is what lets her recover the real parts rather than some arbitrary mixture.
A sparse autoencoder (SAE) does this for a model's hidden states. It learns a dictionary: many more directions than the layer has neurons (it is overcomplete), trained so that every hidden state can be rebuilt from just a few of them. Each learned direction is a candidate feature, and each one's strength on an input is its latent activation.
A tiny worked example. The hidden state is h = (0.8, 0.6) and the dictionary has three directions, (1, 0), (0, 1) and (0.8, 0.6). Two codes rebuild h perfectly:
| Code (strength of each direction) | Rebuilt | Error | Penalty with λ = 0.1 | Loss |
|---|---|---|---|---|
| A: (0.8, 0.6, 0), two directions | (0.8, 0.6) | 0 | 0.1 × (0.8 + 0.6) = 0.14 | 0.14 |
| B: (0, 0, 1), one direction | (0.8, 0.6) | 0 | 0.1 × 1.0 = 0.10 | 0.10 |
Rebuilding alone can't choose between them. The penalty on the total
strength, the L1 penalty (primer.ml.regularization), prefers the code
that uses one direction. One side effect: the best strength for direction
3 is not quite 1. With strength $a$, the loss is $(1 - a)^2 + 0.1a$, which
is lowest at $a = 0.95$ (loss 0.0975): the penalty always pulls strengths
a little toward zero, known as shrinkage.
Level 3: the formula and its symbols
$$ f = \text{ReLU}\big(W_e (h - b_d) + b_e\big), \qquad \hat{h} = W_d f + b_d, \qquad L = \|h - \hat{h}\|^2 + \lambda \sum_{i=1}^{m} |f_i| $$
Symbols
| Symbol | Meaning here | In the example |
|---|---|---|
| $h$ | a hidden state from the model under study | (0.8, 0.6) |
| $m$ | how many latents (dictionary directions) the SAE has, usually far more than neurons | 3 |
| $W_e$, $b_e$ | the encoder's weights and biases: turn a hidden state into latent strengths | |
| $f$ | the code: one strength per latent, ReLU keeps them ≥ 0 and mostly exactly 0 | (0, 0, 1) |
| $W_d$ | the decoder: column $i$ is latent $i$'s direction, kept at length 1 | (1, 0), (0, 1), (0.8, 0.6) |
| $b_d$ | the decoder bias: the typical hidden state, subtracted before encoding and added back after | (0, 0) |
| $\hat{h}$ | the rebuilt hidden state | (0.8, 0.6) |
| $|h - \hat{h}|^2$ | squared length of the rebuild error: add up the squares of its entries | 0 |
| $\lambda$ | the sparsity penalty's strength, chosen by you | 0.1 |
| $\lvert f_i \rvert$ | the size of latent $i$'s strength | |
| $L$ | the loss to minimize: rebuild error plus penalty | 0.10 |
In words: "encode the hidden state into many non-negative strengths, decode it back as a weighted sum of dictionary directions, and pay for both the rebuild error and the total strength used."
With the numbers: for code B, $\hat{h} = 1 \cdot (0.8, 0.6) = (0.8, 0.6)$, the error is 0, the penalty is $0.1 \times 1 = 0.10$, so $L = 0.10$. For code A the penalty is $0.1 \times 1.4 = 0.14$.
Level 3: in Python
In Python:
dictionary = [[1, 0], [0, 1], [0.8, 0.6]]
h = [0.8, 0.6]
lam = 0.1
def loss(code):
# ĥ = Σ f_i · direction_i
rebuilt = [sum(f * d[n] for f, d in zip(code, dictionary)) for n in range(2)]
# ‖h − ĥ‖² + λ Σ |f_i|
return round(sum((a - b) ** 2 for a, b in zip(h, rebuilt)) + lam * sum(abs(f) for f in code), 4)
loss([0.8, 0.6, 0]) # → 0.14
loss([0, 0, 1.0]) # → 0.1
# shrinkage: a slightly weaker code is cheaper still
loss([0, 0, 0.95]) # → 0.0975
flowchart LR H["hidden state h<br/>(d numbers)"] --> ENC["encoder<br/>ReLU(W_e(h − b_d) + b_e)"] ENC --> F["code f<br/>(m ≫ d numbers,<br/>almost all exactly 0)"] F --> DEC["decoder<br/>W_d f + b_d"] DEC --> HH["rebuilt ĥ"] HH --> LOSS["loss = ‖h − ĥ‖² + λ Σ|f|"] H --> LOSS F --> LOSS
Reading it: the hidden state is blown up into a much wider code and squeezed back down. The loss watches two things at once: the rebuild must match the original (the arrow from h), and the code must be small (the arrow from f). Without the second arrow the SAE could use any directions at all. With it, each input is explained by a handful of latents, and each latent's decoder column is a candidate feature you can inspect: read which inputs make it fire, and add its direction to the model to see what it does.
On the superposition toy. Take 20,000 hidden states from the pentagon model (each feature on 5% of the time), train an SAE with 5 latents and λ = 0.3, and compare the learned decoder columns with the five planted directions. Every planted feature is matched by a latent with cosine similarity above 0.99 (1 would be the same direction), and an example with a feature on lights up about 1.07 latents on average. The price is shrinkage: 11% of the variance goes unexplained. With a tiny penalty, λ = 0.01, the SAE rebuilds the data perfectly (0.0% unexplained), yet its worst match to a real feature is only 0.80: any two directions can rebuild a two-dimensional space, so without sparsity pressure nothing forces the dictionary to find the model's actual features.
Reading it: on the left, each grey dot is one hidden state. Because features are sparse, most dots lie along one of five rays: one feature on, at some strength. The few dots off the rays are the rare examples with two features on at once. The thick grey lines are the directions we planted; the red arrows are what the SAE learned from the dots alone, and they land on the rays. On the right, the x-axis is the penalty λ on a log scale. The blue line is the worst match between a planted feature and its nearest latent: it sits around 0.8 to 0.9 for small penalties and peaks near 0.99 at λ = 0.3. The red line, the variance left unexplained, is near 0 for small penalties and climbs past 0.5 at λ = 1, where the penalty starts crushing the codes and the match falls again. Choosing λ is a trade: too small and the latents are not the features, too large and too much of the model's activity is lost.
Why it matters in practice. This is the tool that turned interpretability from "a few hand-picked neurons" into something that scales. Anthropic's Towards Monosemanticity (2023) trained SAEs on a small transformer and found thousands of features that each fire for one recognizable thing. Scaling Monosemanticity (2024) did the same inside a production model, Claude 3 Sonnet, and found features for concepts such as the Golden Gate Bridge. Turning that one feature up made the model bring up the bridge in almost every answer, which is the causal check that the feature means what it seems to mean.
In code: sae_objective computes the loss for one proposed code, SparseAutoencoder holds the encoder and decoder with its hand-derived SparseAutoencoder.loss_and_grads, train_sae fits it, and match_features compares learned directions with planted ones.
What these tools cannot (yet) show
Everyday picture. A brain scan shows which regions light up while a person reads, but not the sentence they are thinking. Every tool in this lesson is like that: a real instrument that measures something true and partial.
A tiny worked example. Each of our own toys already showed a limit:
| What we saw | The limit it shows |
|---|---|
| A probe read tense at 98%, and the output ignored tense completely | Probes find information, not use |
| The lens read "Paris" at a position where patching restored 0% | Reading an activation doesn't mean anything downstream reads it |
| The weak SAE rebuilt everything perfectly and still found the wrong directions | A good rebuild score doesn't mean the features are real |
| The strong SAE left 11% of the variance unexplained | Some of the model's activity is in no feature the dictionary found |
| We knew the landmark model's circuit because we built it | In a real model, nobody has the answer key |
flowchart LR A["Correlation<br/>probes, the logit lens:<br/>'this information is here'"] --> B["Intervention<br/>patching, steering:<br/>'changing this changes the output'"] B --> C["Mechanism<br/>a circuit explained end to end:<br/>'this is how the output is computed'"]
Reading it: the arrow is the direction of stronger evidence. Most claims about real models sit on the left or in the middle. A complete mechanism, every step from input to output accounted for, exists only for narrow behaviours in small or carefully chosen settings. When you read an interpretability result, ask which box it reached.
The main open problems, in plain words:
- Choices shape answers. Patching results depend on how the corrupted prompt is built (swap one word? add noise?). SAE features depend on the dictionary's size and λ: a bigger dictionary can split one feature into several finer ones.
- Redundancy hides importance. Patching one site at a time can miss parts that back each other up. Models have been observed to partly repair themselves when one component is knocked out, so "patching this restores 0%" does not always mean "this plays no role".
- Coverage. The unexplained part of an SAE's rebuild is not noise to ignore: it is model behaviour no feature describes yet.
- Scale and labour. A circuit that explains one behaviour of a small model can take researchers weeks. No one has a complete account of how a frontier model produces any long answer.
- Labels are human guesses. Naming a feature "Golden Gate Bridge" summarizes the inputs that make it fire; checking that the name is right needs interventions like the one above.
Why it matters in practice. Treat these tools as evidence, stacked
next to behavioural evaluations (primer.agents.evals), never as a
certificate. A probe or a lens is a cheap first look; a patching or
steering experiment is the minimum for a causal claim; and a clean result
on a toy, like every result in this lesson, is where understanding starts,
not where it ends.
In 20 seconds
- Features are directions in activation space, not individual neurons; a dot product with the right direction reads a feature out.
- Probes are small linear classifiers trained on frozen activations. They show information is present, not that the model uses it; score them on unseen data and against a control task.
- The logit lens applies the model's own output layer to intermediate residual states, showing what it would predict after each layer.
- Activation patching copies one activation from a clean run into a corrupted run and measures how much of the right answer returns: the basic causal experiment.
- Superposition: when features are sparse, models store more features than they have neurons, as nearly perpendicular directions, which makes neurons polysemantic.
- Sparse autoencoders learn an overcomplete dictionary with an L1 penalty so each input uses a few latents, recovering features from superposition, at the cost of some unexplained variance.
Self-test questions
What does it mean to say a feature is a "direction" rather than a neuron? The feature's presence is stored as a pattern across many neurons: the hidden state moves along a particular direction in proportion to how strongly the feature is present. You read it with a dot product against that direction. Any single neuron is typically a mix of several features.
A linear probe reads a property from layer 12 at 95% accuracy. What can you conclude, and what can't you? You can conclude the property is linearly decodable from layer 12, as long as the 95% was on held-out examples and clearly beats a control task with random labels. You can't conclude the model uses it. For that you need an intervention: remove or change that information and see whether the output changes.
How does the logit lens work, and why does it agree with the model at the last layer? It takes the residual stream after some layer and applies the model's own final normalization and unembedding, turning it into next-token probabilities. At the last layer that is exactly the computation the model itself performs, so the two must match. Earlier layers are read "as if the model stopped there".
Describe an activation patching experiment, and what the "fraction restored" means. Run a clean prompt and save its activations; run a corrupted prompt that changes one fact; then rerun the corrupted prompt with one activation replaced by its clean value. The fraction restored is how much of the clean-minus-corrupted logit difference the patch brings back: 1 means that activation alone carries the fact, 0 means it carries none of it.
Why do models use superposition, and why does it make neurons hard to interpret? There are more useful features than neurons. When features are rarely active at the same time, the model can store them as nearly perpendicular directions and filter the small interference with a bias and ReLU. The directions can't line up with the neurons, so each neuron responds to several unrelated features: it is polysemantic.
Why does a sparse autoencoder need the L1 penalty? What goes wrong if λ is too small or too large? Many dictionaries rebuild the data equally well; the penalty picks the one where each input uses few latents, which pushes latents onto the real features. Too small and the SAE rebuilds perfectly with meaningless directions; too large and it shrinks activations and leaves much of the model's activity unexplained.
The papers behind this lesson
- Alain & Bengio, Understanding intermediate layers using linear classifier probes (2016): https://arxiv.org/abs/1610.01644. Introduced linear probes as a way to measure what each layer of a network makes linearly available.
- Hewitt & Liang, Designing and Interpreting Probes with Control Tasks (2019): https://arxiv.org/abs/1909.03368. Showed that probes can succeed by memorizing, and introduced control tasks and selectivity to tell the two apart.
- Belrose et al., Eliciting Latent Predictions from Transformers with the Tuned Lens (2023): https://arxiv.org/abs/2303.08112. Formalized the logit lens and fixed its failures on early layers by training a small translator for each layer.
- Meng et al., Locating and Editing Factual Associations in GPT (2022): https://arxiv.org/abs/2202.05262. Introduced causal tracing, found factual recall in middle-layer MLPs at the subject's last token, and edited single facts there. Annotated companion
- Wang et al., Interpretability in the Wild: a Circuit for Indirect Object Identification in GPT-2 small (2022): https://arxiv.org/abs/2211.00593. Used patching to reverse-engineer a complete circuit for one behaviour of a real language model.
- Elhage et al., Toy Models of Superposition (2022): https://arxiv.org/abs/2209.10652. Showed with small ReLU models that sparse features are stored in superposition, and when and how the geometry changes. Annotated companion
- Bricken et al., Towards Monosemanticity: Decomposing Language Models With Dictionary Learning (2023): https://transformer-circuits.pub/2023/monosemantic-features/index.html. Trained sparse autoencoders on a small transformer and found thousands of interpretable features hidden in polysemantic neurons. Annotated companion
- Templeton et al., Scaling Monosemanticity: Extracting Interpretable Features from Claude 3 Sonnet (2024): https://transformer-circuits.pub/2024/scaling-monosemanticity/index.html. Scaled sparse autoencoders to a production model and steered its behaviour through the features they found. Annotated companion
Further reading
- Olah et al., Zoom In: An Introduction to Circuits (Distill, 2020): https://distill.pub/2020/circuits/zoom-in/
- Elhage et al., A Mathematical Framework for Transformer Circuits (2021): https://transformer-circuits.pub/2021/framework/index.html
- Elhage et al., Toy Models of Superposition (2022): https://transformer-circuits.pub/2022/toy_model/index.html
- Bricken et al., Towards Monosemanticity (2023): https://transformer-circuits.pub/2023/monosemantic-features/index.html
- Templeton et al., Scaling Monosemanticity (2024): https://transformer-circuits.pub/2024/scaling-monosemanticity/index.html
- Cunningham et al., Sparse Autoencoders Find Highly Interpretable Features in Language Models (2023): https://arxiv.org/abs/2309.08600
- Belinkov, Probing Classifiers: Promises, Shortcomings, and Advances (2021): https://arxiv.org/abs/2102.12452
- Meng et al., Locating and Editing Factual Associations in GPT (2022): https://arxiv.org/abs/2202.05262
- Belrose et al., Eliciting Latent Predictions from Transformers with the Tuned Lens (2023): https://arxiv.org/abs/2303.08112
1r""" 2# Looking inside the model 3 4Run: `python -m primer.ml.interpretability` 5 6New to the notation? `primer.notation` explains every symbol used here from 7zero. This lesson builds on the residual stream of `primer.ml.transformer` 8and on behavioural evaluation in `primer.agents.evals`. 9 10## Level 1: The practitioner's guide 11 12**In one sentence.** Interpretability is the set of tools that read and 13change a model's internal activations to find out what it represents and 14which of those representations cause its answer, as opposed to 15evaluations, which only measure what it says. 16 17**When you need it.** Behavioural testing answers "what does the model 18do?", and for most shipping decisions that is enough. You need to open the 19hood when the question is "why?": a failure you cannot reproduce from the 20outside, a model that may be right for the wrong reason (the classifier 21that detects wolves by the snow behind them), a claim that the model is 22relying on a protected attribute, or a feature you want to turn up or down 23without retraining. The tell is that you are about to explain a model's 24behaviour from its outputs alone and cannot tell two explanations apart. 25This lesson's toy shows why that is dangerous: a probe reads "tense" out 26of the hidden state at 98% accuracy on unseen examples, and flipping tense 27moves the model's output by exactly 0.00, while flipping sentiment moves 28it by 2.00. What is readable inside a model and what the model uses are 29different questions, and only an intervention answers the second. One 30practical limit before anything else: every tool here needs the 31activations of a model you run yourself. A hosted API gives you tokens, 32not hidden states, so for a model behind an API your instrument is still 33the evaluation. 34 35**Your options.** From the cheapest look to the strongest evidence: 36 37| Option | What it does | What it tells you | What it costs | Where it lives | 38|---|---|---|---|---| 39| Behavioural evaluation | Runs the model on your tasks and scores the outputs | What the model does, at scale, for any model including hosted ones | A golden set and a scoring rule | `primer.agents.evals` | 40| Linear probe | Trains a tiny classifier on frozen hidden states to read one property | That the property is linearly present at that layer, if it beats a control task on unseen data | Labelled examples (100 sufficed in this lesson), minutes of training | Your code, on an open-weight model | 41| Logit lens (and tuned lens) | Applies the model's own output layer to the residual stream after each layer | When and where a prediction forms inside the model; no training needed | Nothing beyond a forward pass; a small trained translator per layer for the tuned version | Any residual-stream model you can hook | 42| Activation patching (causal tracing) | Copies one activation from a clean run into a corrupted run and measures how much of the answer returns | Which sites cause the answer: the minimum for a causal claim | Two prompts that differ in one fact, and one forward pass per site tested | Hooking libraries such as TransformerLens | 43| Sparse autoencoder (SAE) features | Learns an overcomplete dictionary so each hidden state is a few interpretable directions | What the model's units of meaning are, in a form you can read and steer | A large training run on activations, a penalty to tune, and 11% of the variance unexplained in this lesson's toy | Released SAE suites for open models, or your own training | 44| Feature steering | Adds a found direction to the residual stream at run time | Whether a feature causes what its name suggests, and a lever without retraining | The feature must exist first; too strong a push degrades the output | Research demos on production models | 45| Circuit analysis | Patches site by site until every step from input to output is accounted for | A complete mechanism for one behaviour | Weeks of expert time for a small model | Research | 46 47**How to choose.** Start from the question, and reach for the cheapest 48tool that answers it. 49 50- "Does the model know X?": a probe, scored on held-out examples against a 51 random-label control. In this lesson the control fits 73% of its 52 training set and scores 50% unseen, which is what a probe that only 53 memorised looks like. 54- "When does it decide?": the logit lens. In the landmark model the fact 55 appears at the subject word after layer 1 and reaches the last word only 56 after layer 2. 57- "Does this part cause the answer?": patching. Reading is not enough: the 58 lens reads "Paris" at 0.995 at a position where patching restores 0%. 59- "What are the features, and can I steer one?": an SAE, followed by a 60 steering experiment to check that the feature means what its label says. 61- "Is the model relying on the wrong thing in production?": a probe or 62 SAE feature to find the candidate, then patching or steering to confirm 63 it, then a behavioural evaluation to measure the effect at scale. 64- Whatever you pick, place the result on the ladder from correlation 65 (probes, the lens) through intervention (patching, steering) to 66 mechanism (a circuit), and claim only the rung you reached. 67 68**What it costs.** Probes and the logit lens cost almost nothing: frozen 69activations, a forward pass, and for a probe a few hundred labelled 70examples and seconds of training. Patching costs one forward pass per site, 71so a full map is layers times positions; the lesson's is 3 by 3, a real 72model's is hundreds by thousands, and it is repeated for every prompt pair 73you try. Sparse autoencoders are the expensive tool: millions of 74activations, a dictionary far wider than the layer, and a sparsity penalty 75whose choice is a trade. In this lesson's sweep, a penalty of 0.3 matches 76every planted feature above 0.99 and leaves 11% of the variance 77unexplained; a penalty of 0.01 rebuilds everything and matches the worst 78feature at only 0.80, so a perfect rebuild score is not evidence that the 79features are real. Circuit-level explanations cost research time measured 80in weeks per behaviour. And all of it presumes access: none of these tools 81runs on a model you only reach through an API. 82 83**What breaks.** 84 85- **Decodable is not used.** The 98%-readable feature with zero effect on 86 the output. Never claim "the model uses X" from a probe; intervene. 87- **A probe that is too clever.** A deep probe can compute the property 88 itself; a probe with no control task can fit noise. Keep probes linear, 89 score them on unseen data, compare with random labels. 90- **Reading a site nothing reads.** The lens can see an answer at a 91 position no later layer consults. Patching restores 0% there. 92- **A choice that shapes the answer.** Patching results depend on how the 93 corrupted prompt is built; SAE features depend on the dictionary's size 94 and penalty, and a bigger dictionary can split one feature into several. 95- **Redundancy.** Models can partly repair themselves when one component 96 is knocked out, so "patching this restores nothing" does not always mean 97 "this plays no role". 98- **The unexplained remainder.** The variance an SAE does not rebuild is 99 model behaviour no feature describes yet, not noise. 100- **Labels that are guesses.** A feature named "Golden Gate Bridge" is a 101 summary of what makes it fire. The name is checked by steering, not by 102 reading. 103- **Polysemantic neurons.** When features outnumber neurons they cannot 104 line up with them, so a single neuron responds to several unrelated 105 things. Look along directions, not at neurons. 106 107**In the wild.** TransformerLens exposes the internal activations of 108thousands of open-weight models and lets you cache, edit and replace them 109as the model runs, which is the plumbing for the lens, probes and 110patching. Causal tracing (Meng et al., 2022) located factual recall in 111middle-layer MLPs at the subject's last token in GPT-style models and 112edited single facts there; Wang et al. (2022) patched their way to a 113complete circuit for indirect-object identification in GPT-2 small. The 114tuned lens (Belrose et al., 2023) fixed the logit lens on early layers. 115Anthropic's *Towards Monosemanticity* (2023) trained SAEs on a small 116transformer, and *Scaling Monosemanticity* (2024) did it inside Claude 3 117Sonnet, where turning up one feature made the model bring up the Golden 118Gate Bridge in almost every answer. Google DeepMind's Gemma Scope releases 119trained SAEs for every layer of the Gemma 2 2B and 9B base models, so a 120practitioner can inspect features without training a dictionary. 121 122**Go deeper.** Level 2 builds each tool on models small enough to check 123by hand: a feature as a direction read back by a dot product, a probe 124trained by gradient descent with a control task, the logit lens verified 125against the real output on `primer.ml.transformer.TinyGPT`, activation 126patching on a two-layer landmark model with a known circuit, superposition 127in a five-features-in-two-neurons toy, and a sparse autoencoder that 128recovers the planted features only when its penalty is right. If you only 129needed to know what these tools can and cannot tell you about a model in 130production, you are done. 131 132## Level 2: How it works, from scratch 133 134**Everyday picture.** A car makes a strange noise. You can take it for a 135test drive and note when the noise happens: that is testing the car's 136*behaviour*. Or you can open the hood, put a stethoscope on the engine, 137and swap parts until the noise stops: that is looking at the *mechanism*. 138Test drives tell you *what* the car does. Only the open hood tells you 139*why*. 140 141Every other lesson in this primer treats a trained model as something to 142build, train or call. Evaluations (`primer.agents.evals`) are test drives: 143they measure what the model says. **Interpretability** opens the hood. It 144asks what the model's billions of internal numbers represent, and which of 145them cause the answer. Three reasons to care: 146 147- **Debugging.** When a model gets something wrong, you want to know which 148 step failed, the same way you read a stack trace rather than just the 149 error message. 150- **Trust.** A model can be right for the wrong reason. An image classifier 151 that "detects" wolves by looking for snow in the background scores well 152 until it meets a wolf on grass. 153- **Safety.** The output is only part of what a model computes. Some 154 questions (is it relying on a stereotype? does it know more than it 155 says?) can only be answered by reading the computation itself. 156 157**A tiny worked example.** This lesson asks four questions of one kind of 158sentence, "the Eiffel Tower is in …" → "Paris", each with its own tool, on 159toy models small enough to check by hand: 160 161| Question | Tool | What our toy shows | 162|---|---|---| 163| Is a property stored in this hidden state? | a **probe** | sentiment, read at 99% on unseen examples | 164| What would the model say if it stopped at layer ℓ? | the **logit lens** | "London", then "Paris", then "Paris" | 165| Which activations *cause* the answer? | **activation patching** | the subject word early, the last word late | 166| What are the model's units of meaning? | **superposition** and **sparse autoencoders** | five features packed into two neurons, then recovered | 167 168```mermaid 169flowchart LR 170 M["A trained model<br/>(frozen)"] --> A["Hidden activations<br/>at every layer and word"] 171 A --> P["Probe:<br/>is property X in here?"] 172 A --> L["Logit lens:<br/>what would it predict now?"] 173 A --> AP["Activation patching:<br/>does this activation cause the answer?"] 174 A --> S["Sparse autoencoder:<br/>which features make up this state?"] 175 P & L --> R["Reading<br/>(correlation)"] 176 AP --> C["Intervening<br/>(causation)"] 177 S --> U["Units of analysis<br/>(what to read and intervene on)"] 178``` 179 180**Reading it:** every tool starts from the same place: the numbers a frozen 181model computes inside itself while it runs. The top two tools only *read* 182those numbers, so what they find is a correlation. Patching *changes* them 183and watches the output, so what it finds is a cause. Sparse autoencoders 184answer an earlier question: before you can read or change "a feature", you 185need to know what the features are. Keep the reading/intervening split in 186mind; it is the most important distinction in this lesson. 187 188## A feature is a direction 189 190**Everyday picture.** A band with two instruments plays through two 191speakers. The sound engineer pans the guitar mostly to the right speaker 192and the piano mostly to the left, but each speaker plays a *mix* of both. 193If you unplug one speaker you don't lose "the guitar"; you lose part of 194everything. The instruments are not the speakers. Each instrument is a 195*setting across all speakers*, a direction, and the sound in the room is 196the sum of every instrument times how loudly it plays. 197 198A model's hidden state works the same way. The speakers are **neurons**: 199the individual numbers in a layer's output vector. The instruments are 200**features**: the things the model has learned to track, such as "this 201review is positive" or "this verb is in the past tense". Each feature is 202stored as a **direction** in the space of neurons, and the hidden state is 203the sum of every active feature's direction, scaled by how strongly it is 204present. 205 206**A tiny worked example.** Take a hidden layer with two neurons and two 207features whose directions are "positive" = (0.6, 0.8) and "past tense" = 208(0.8, −0.6). A sentence that is quite positive (0.9) and a little 209past-tense (0.4) has the hidden state 210 2110.9 · (0.6, 0.8) + 0.4 · (0.8, −0.6) = (0.54 + 0.32, 0.72 − 0.24) = **(0.86, 0.48)**. 212 213Neuron 1 reads 0.86 and neuron 2 reads 0.48. Neither number is "how 214positive" or "how past-tense": each neuron is a blend of both features. 215To get the features back, take the **dot product** (multiply matching 216entries and add; see `primer.notation`) of the hidden state with each 217direction: 218 219- positive: 0.6 · 0.86 + 0.8 · 0.48 = 0.516 + 0.384 = **0.9** 220- past tense: 0.8 · 0.86 − 0.6 · 0.48 = 0.688 − 0.288 = **0.4** 221 222$$ 223h = \sum_{i=1}^{k} f_i \, d_i 224\qquad\qquad 225\hat{f}_i = d_i \cdot h 226$$ 227 228**Symbols** 229 230| Symbol | Meaning here | In the example | 231|---|---|---| 232| $h$ | the hidden state: one number per neuron | (0.86, 0.48) | 233| $k$ | how many features are present | 2 | 234| $i$ | which feature | 1 = positive, 2 = past tense | 235| $f_i$ | how strongly feature $i$ is present | $f_1 = 0.9$, $f_2 = 0.4$ | 236| $d_i$ | feature $i$'s direction: a list with one number per neuron, of length 1 | $d_1 = (0.6, 0.8)$ | 237| $\sum_{i=1}^{k}$ | add up one term per feature | two terms | 238| $\hat{f}_i$ | the amount of feature $i$ we read back (the hat means "estimated") | 0.9 | 239| $\cdot$ | dot product: multiply matching entries, then add | | 240 241**In words:** "the hidden state is the sum of each feature's direction 242times its strength; to read a feature back, dot the hidden state with that 243feature's direction." 244 245**With the numbers:** $h = 0.9 \cdot (0.6, 0.8) + 0.4 \cdot (0.8, -0.6) = 246(0.86, 0.48)$, and $\hat{f}_1 = (0.6, 0.8) \cdot (0.86, 0.48) = 0.9$. The 247read-back is exact here because the two directions are **perpendicular** 248(their dot product is 0.6 · 0.8 + 0.8 · (−0.6) = 0) and each has length 1. 249Hold on to that condition: the section on superposition is about what 250happens when it fails. 251 252**In Python:** 253 254```python 255positive = [0.6, 0.8] 256past = [0.8, -0.6] 257# h = Σ f_i d_i 258h = [0.9 * p + 0.4 * q for p, q in zip(positive, past)] 259[round(x, 2) for x in h] # → [0.86, 0.48] 260# f̂_i = d_i · h 261round(sum(d * x for d, x in zip(positive, h)), 2) # → 0.9 262round(sum(d * x for d, x in zip(past, h)), 2) # → 0.4 263# the directions are perpendicular: their dot product is 0 264round(sum(p * q for p, q in zip(positive, past)), 2) # → 0.0 265``` 266 267 268 269**Reading it:** the axes are the two neurons. The blue and green arrows are 270the two feature directions; the red arrow is the hidden state. The dotted 271lines drop from the hidden state onto each feature arrow, and where they 272land (0.9 of the way along blue, 0.4 along green) is the dot product: the 273amount of that feature. Notice that the red arrow's coordinates on the 274*neuron* axes, 0.86 and 0.48, are neither amount. To read a model you have 275to look along the right directions, not along the neurons. 276 277That features are directions, and that most of what a model tracks can be 278read with a dot product, is called the **linear representation 279hypothesis**. It is a hypothesis, not a law, but it holds often enough to 280power everything below. You have met it before: word-vector arithmetic 281such as king − man + woman ≈ queen (`primer.ml.embeddings.word2vec`) works 282because "royalty" and "gender" are directions. 283 284**In code:** `compose_features` builds a hidden state from feature amounts and `read_features` reads them back with dot products. 285 286## Probes: can a straight line read it out? 287 288**Everyday picture.** A doctor can't ask your liver how it is doing, but a 289blood test can measure a marker that tells them. A **probe** is a blood 290test for a hidden state: a tiny classifier, trained by you, that looks at a 291layer's activations and answers one yes/no question, such as "is this 292review positive?". The model itself is frozen; only the probe learns. 293 294**A tiny worked example.** A probe for "positive" on a two-neuron layer 295has weights w = (1, 0.5) and bias b = 0. On the hidden state h = (2, −1): 296 2971. Score it: 1 · 2 + 0.5 · (−1) + 0 = 1.5. 2982. Squash the score into a probability with the **sigmoid** σ(z) = 299 1 / (1 + e^−z), which maps any number into the range 0 to 1: 300 σ(1.5) = 1 / (1 + 0.223) = **0.818**. 301 302The probe is 82% sure this hidden state belongs to a positive review. On 303h = (−1, 1) the score is −1 + 0.5 = −0.5 and σ(−0.5) = 0.378: probably 304negative. 305 306$$ 307p = \sigma(w \cdot h + b) = \frac{1}{1 + e^{-(w \cdot h + b)}} 308$$ 309 310**Symbols** 311 312| Symbol | Meaning here | In the example | 313|---|---|---| 314| $h$ | the frozen model's hidden state for one example | (2, −1) | 315| $w$ | the probe's weights: one per neuron, learned | (1, 0.5) | 316| $b$ | the probe's bias: one number, learned | 0 | 317| $w \cdot h + b$ | the probe's raw score (its **logit**) | 1.5 | 318| $\sigma$ | the sigmoid: turns any score into a probability between 0 and 1 | | 319| $e$ | Euler's number, about 2.718 | $e^{-1.5} = 0.223$ | 320| $p$ | the probe's probability that the property is present | 0.818 | 321 322**In words:** "dot the hidden state with the probe's weights, add the 323bias, and squash the result into a probability." 324 325**With the numbers:** $p = \sigma(1 \cdot 2 + 0.5 \cdot (-1) + 0) = 326\sigma(1.5) = 1 / (1 + 0.223) = 0.818$. 327 328**In Python:** 329 330```python 331import math 332w = [1.0, 0.5] 333b = 0.0 334h = [2.0, -1.0] 335# w · h + b 336z = sum(wi * hi for wi, hi in zip(w, h)) + b 337z # → 1.5 338# σ(z) 339round(1 / (1 + math.exp(-z)), 3) # → 0.818 340# a second hidden state, (−1, 1), reads as probably negative 341round(1 / (1 + math.exp(-(-1 * 1.0 + 1 * 0.5))), 3) # → 0.378 342``` 343 344This is exactly **logistic regression** (`primer.ml.neural_net`), trained 345the usual way: show it examples whose answer you know, measure its 346binary cross-entropy (`primer.ml.losses`), and nudge w and b downhill. The 347gradient of that loss with respect to the score is simply (p − y), the gap 348between the probe's probability and the true 0/1 label, which makes each 349update one line of code. 350 351```mermaid 352flowchart LR 353 T["Text with a known label<br/>(positive or negative)"] --> M["Frozen model<br/>(no weights change)"] 354 M --> H["Hidden state h<br/>at the layer under study"] 355 H --> P["Probe: σ(w · h + b)"] 356 P --> Y["Probability the label is 'positive'"] 357 Y --> G["Compare with the true label<br/>(cross-entropy)"] 358 G -. "gradient updates w and b only" .-> P 359``` 360 361**Reading it:** the solid arrows are one forward pass; the dotted arrow is 362learning. The gradient stops at the probe: the model under study never 363changes, so whatever the probe finds was already in the hidden state. That 364is the point of a probe. It measures the model, not a new model trained on 365top of it. 366 367**On our toy model.** `PlantedModel` stores two known features in a 36832-neuron hidden layer: sentiment and tense, each as ±1 along its own 369random direction, plus random noise of size 0.4 on every neuron. No single 370neuron holds either feature. A probe trained on just 100 examples reads 371sentiment correctly on **99%** of 1,000 examples it never saw. 372 373**Why keep probes simple, and check them against a control.** A probe can 374succeed for the wrong reason. Give a probe **random labels**, a *control 375task* with nothing real to find, and it still fits **73%** of its 100 376training examples, because 33 adjustable numbers can memorize a lot of 377noise. On unseen examples it scores **50%**, a coin flip. Two habits 378follow: always score a probe on examples it never saw, and compare it with 379a control task, so that "the probe works" means "the hidden state encodes 380this" rather than "the probe is clever". This is also why probes are kept 381*linear*: a powerful probe (a deep network) could compute the property 382from raw ingredients by itself, and then its success would tell you about 383the probe, not the model. 384 385### Decodable is not the same as used 386 387**Everyday picture.** A library holds a book nobody ever borrows. Finding 388it on the shelf proves the library *has* it, not that it shaped anything 389any reader did. 390 391**A tiny worked example.** Our toy model's output reads sentiment and, by 392construction, gives tense a weight of exactly zero. A probe still reads 393tense at **98%** on unseen examples. Now intervene: take 200 examples, flip 394only one feature's label, keep everything else fixed, and measure how much 395the model's output moves. 396 397| Flip | Probe accuracy for it | Output moves by | 398|---|---|---| 399| sentiment | 99% | **2.00** | 400| tense | 98% | **0.00** | 401 402 403 404**Reading it:** the left panel is what probes can read. Grey bars are 405accuracy on the probe's own training examples, blue bars on unseen ones, 406and the red dashed line is chance. Sentiment and tense are both clearly 407readable; random labels look readable on the training set and collapse to 408chance on unseen data, which is exactly what the control is for. The right 409panel is what the model *uses*: the change in its output when one feature 410is flipped. Tense is as readable as sentiment and has no effect at all. 411The two panels answer different questions, and only the right one is about 412cause. 413 414A probe finding information does not mean the model uses it. To find out 415what the model uses you have to change something and watch the output, 416which is what the rest of this lesson does. 417 418**In code:** `train_probe` fits a `Probe` by gradient descent on frozen hidden states from `PlantedModel`; `probe_report` trains the three probes, and `flip_effect` performs the intervention. 419 420## The logit lens: reading the model's mind mid-thought 421 422**Everyday picture.** A writer keeps every draft of an article. Draft 1 423says the landmark is "in a European capital", draft 2 says "probably 424Paris", the final says "Paris". Reading the drafts shows *when* the writer 425made up their mind. The **logit lens** reads a language model's drafts. 426 427It works because of how a transformer is built (`primer.ml.transformer`). 428Each word has a running vector, the **residual stream**. Every layer reads 429it and *adds* a correction to it, rather than replacing it. At the very 430end, the output layer (the **unembedding**) turns the final vector into 431one score per word in the vocabulary. Because every layer writes into the 432same stream, you can apply that final step early, after any layer, and ask 433"what would the model predict if it stopped here?" 434 435**A tiny worked example.** A vocabulary of three words, Paris, London and 436Rome, a residual stream of two numbers, and an unembedding that scores 437Paris = first number, London = second number, Rome = minus the first: 438 439| After | Residual $h$ | Scores (Paris, London, Rome) | Probabilities | Top guess | 440|---|---|---|---|---| 441| the embedding | (0.2, 0.3) | (0.2, 0.3, −0.2) | (0.36, 0.40, 0.24) | London | 442| layer 1, which adds (0.6, 0) | (0.8, 0.3) | (0.8, 0.3, −0.8) | (0.55, 0.34, 0.11) | Paris | 443| layer 2, which adds (1.2, −0.3) | (2.0, 0.0) | (2.0, 0.0, −2.0) | (0.87, 0.12, 0.02) | Paris | 444 445The first draft is a vague "some capital" that happens to lean London; 446layer 1 tips it to Paris; layer 2 commits. 447 448$$ 449\text{lens}_\ell = \text{softmax}\big(W_U \, \text{LN}(h_\ell)\big) 450$$ 451 452**Symbols** 453 454| Symbol | Meaning here | In the example | 455|---|---|---| 456| $\ell$ | which layer we stop after (0 = straight after the embedding) | 2 | 457| $h_\ell$ | the residual stream after layer $\ell$, for one word | (2.0, 0.0) | 458| $\text{LN}$ | the model's final normalization (`primer.ml.transformer`); our 2-number toy has none, so here LN(h) = h | (2.0, 0.0) | 459| $W_U$ | the unembedding: one row per vocabulary word, dotted with the state to give that word's score | rows (1, 0), (0, 1), (−1, 0) | 460| $W_U \, \text{LN}(h_\ell)$ | the scores (**logits**), one per word | (2, 0, −2) | 461| softmax | turns scores into probabilities that add up to 1 (`primer.ml.attention`) | (0.87, 0.12, 0.02) | 462| $\text{lens}_\ell$ | what the model would predict if it stopped after layer $\ell$ | Paris at 0.87 | 463 464**In words:** "take the residual stream partway up, push it through the 465model's own final normalization and output layer, and read off the 466probabilities." 467 468**With the numbers:** after layer 2, $W_U h_2 = (1 \cdot 2 + 0 \cdot 0,\; 4690 \cdot 2 + 1 \cdot 0,\; -1 \cdot 2 + 0 \cdot 0) = (2, 0, -2)$; then 470$e^2 = 7.39$, $e^0 = 1$, $e^{-2} = 0.14$, total 8.52, so P(Paris) = 4717.39 / 8.52 = 0.87. 472 473**In Python:** 474 475```python 476import math 477# rows: Paris, London, Rome 478W_U = [[1, 0], [0, 1], [-1, 0]] 479h = [0.2, 0.3] 480writes = [[0.6, 0.0], [1.2, -0.3]] 481# a residual layer adds its write to the stream 482for w in writes: 483 h = [a + b for a, b in zip(h, w)] 484[round(x, 2) for x in h] # → [2.0, 0.0] 485# W_U h: one score per word 486scores = [sum(u * x for u, x in zip(row, h)) for row in W_U] 487[round(s, 2) for s in scores] # → [2.0, 0.0, -2.0] 488# softmax 489exps = [math.exp(s) for s in scores] 490[round(e / sum(exps), 2) for e in exps] # → [0.87, 0.12, 0.02] 491``` 492 493```mermaid 494flowchart LR 495 E["Embedding"] --> H0(("h₀")) --> L1["Layer 1<br/>adds its write"] --> H1(("h₁")) --> L2["Layer 2<br/>adds its write"] --> H2(("h₂")) --> OUT["Final norm + W_U<br/>(the real output)"] 496 H0 -.-> LENS0["lens: norm + W_U<br/>London 0.40"] 497 H1 -.-> LENS1["lens: norm + W_U<br/>Paris 0.55"] 498 H2 -.-> LENS2["lens: norm + W_U<br/>Paris 0.87"] 499``` 500 501**Reading it:** the solid line along the middle is the residual stream: it 502flows from the embedding to the real output, and each layer adds to it. 503The dotted taps hang the model's own output layer off the stream after 504every layer. Nothing is trained; the lens borrows weights the model 505already has. At the top layer the tap and the real output are the same 506computation, so they must agree exactly, which is a good test that the 507lens is wired correctly (on `primer.ml.transformer.TinyGPT`, they do). 508 509**On a model where we know the answer.** `LandmarkModel` is a two-layer, 510one-head transformer built by hand to complete "Eiffel is in" with 511"Paris". Layer 1 is an MLP that looks up a fact at every word (landmark 512in, city out); layer 2 is an attention head that lets the last word, "in", 513find the landmark and copy its city. 514 515```mermaid 516flowchart LR 517 T["Eiffel · is · in"] --> EMB["Embed each word"] 518 EMB --> MLP["Layer 1: MLP at every word<br/>Eiffel → writes 'Paris'"] 519 MLP --> ATT["Layer 2: attention head<br/>'in' looks for the landmark,<br/>copies its city"] 520 ATT --> U["Output at 'in':<br/>Paris 3, Rome 0, London 0"] 521``` 522 523**Reading it:** read left to right as the two steps of the answer. The 524fact ("the Eiffel Tower is in Paris") is looked up at the landmark's own 525position in layer 1. It is only moved to the last position, the one that 526predicts the next word, in layer 2. We built it this way on purpose, so 527every tool below can be checked against a known truth. 528 529 530 531**Reading it:** the left panel is the worked example: at each layer, three 532bars for the three words, and the blue Paris bar grows from 0.36 to 0.87. 533The right panel runs the lens over the landmark model: rows are layers, 534columns are words, and each cell is the lens's probability of "Paris" 535(1/3 = 0.33 means no opinion among three cities). The Eiffel column goes 536dark after layer 1, when the MLP looks up the fact. The "in" column only 537goes dark after layer 2, when attention copies the city over. The lens has 538shown *where* and *when* the answer appears, without being told. 539 540The logit lens was first described for GPT-2 in 2020, where middle layers 541already "guess" the next word surprisingly well. In some models the early 542layers read as nonsense through the final output layer, because they do 543not yet speak its "language"; the **tuned lens** fixes this by training a 544small translator for each layer before applying the output layer. 545 546**In code:** `logit_lens` applies an output layer to any residual states, `worked_lens` runs the table above, `residual_stream` and `tinygpt_logit_lens` do the same for `primer.ml.transformer.TinyGPT`, and `lens_map` builds the grid for `LandmarkModel`. 547 548## Activation patching: which activations cause the answer? 549 550**Everyday picture.** Two cars of the same model sit side by side: one 551starts, one doesn't. You move parts from the good car into the bad one, 552one part at a time, and try the ignition after each swap. The part whose 553swap makes the bad car start is the part that mattered. **Activation 554patching** (also called **causal tracing**) does this with a model's 555internal activations. 556 557**A tiny worked example.** Run the landmark model twice: 558 559- the **clean** prompt "Eiffel is in": scores Paris 3, Rome 0, so the 560 **logit difference** Paris − Rome is **+3**; 561- the **corrupted** prompt "Colosseum is in": Paris 0, Rome 3, so the 562 difference is **−3**. 563 564Now run the corrupted prompt again, but at one chosen place overwrite the 565activation with the one from the clean run, and see how much of the gap 566between −3 and +3 comes back: 567 568- patch the MLP output at the landmark's position: the difference jumps 569 back to **+3**, so 100% of the answer is restored; 570- patch the state at "is": it stays at **−3**, 0% restored ("is" is the 571 same word in both prompts, so its state carries nothing about the 572 landmark). 573 574$$ 575\text{restored} = \frac{LD_{\text{patched}} - LD_{\text{corrupt}}}{LD_{\text{clean}} - LD_{\text{corrupt}}} 576$$ 577 578**Symbols** 579 580| Symbol | Meaning here | In the example | 581|---|---|---| 582| $LD$ | **logit difference**: the correct answer's score minus the wrong answer's score (Paris − Rome) | | 583| $LD_{\text{clean}}$ | on the clean prompt, nothing patched | +3 | 584| $LD_{\text{corrupt}}$ | on the corrupted prompt, nothing patched | −3 | 585| $LD_{\text{patched}}$ | on the corrupted prompt, with one clean activation copied in | +3, −3, or anything between | 586| restored | the fraction of the gap that one patch closes: 0 = nothing, 1 = everything | 1, 0 | 587 588**In words:** "how far the patch moved the answer, as a share of the full 589distance from the corrupted answer to the clean one." 590 591**With the numbers:** patching the landmark's MLP output gives 592(3 − (−3)) / (3 − (−3)) = 6 / 6 = 1. Patching "is" gives 593(−3 − (−3)) / 6 = 0. A patch that left the difference at 0 would give 594(0 − (−3)) / 6 = 0.5: half the answer restored. 595 596**In Python:** 597 598```python 599clean, corrupt = 3.0, -3.0 600# restored = (LD_patched − LD_corrupt) / (LD_clean − LD_corrupt) 601def restored(patched): 602 return (patched - corrupt) / (clean - corrupt) 603# the landmark's MLP output, patched in 604restored(3.0) # → 1.0 605# the state at "is", patched in 606restored(-3.0) # → 0.0 607# a patch that only reaches a tie 608restored(0.0) # → 0.5 609``` 610 611Why a *difference* of two scores, rather than the probability of Paris? 612Because the corrupted prompt differs from the clean one in exactly one 613fact, the difference measures exactly that fact, and it moves smoothly, 614where a probability can saturate near 0 or 1 and hide a change. 615 616```mermaid 617flowchart TB 618 subgraph C["1. Clean run: 'Eiffel is in'"] 619 c1["save every activation"] --> c2["Paris − Rome = +3"] 620 end 621 subgraph K["2. Corrupted run: 'Colosseum is in'"] 622 k1["Paris − Rome = −3"] 623 end 624 subgraph P["3. Patched run: corrupted prompt,<br/>ONE activation taken from the clean run"] 625 p1["Paris − Rome = ?"] 626 end 627 c1 -- "copy one activation" --> P 628 P --> R["fraction restored =<br/>(? − (−3)) / (3 − (−3))"] 629 K --> R 630 C --> R 631``` 632 633**Reading it:** there are three runs. The clean run is a donor: its 634activations are saved. The corrupted run sets the baseline. The patched 635run is the corrupted prompt with a single activation transplanted from 636the donor. Repeat the patched run once per layer and position and you get 637one "fraction restored" per site: a map of where the answer is carried. 638 639 640 641**Reading it:** rows are the residual stream after each layer, columns are 642the three positions (clean word / corrupted word), and each cell is the 643fraction of the answer restored by patching that one state. The answer 644lives at the landmark early (the Eiffel column is 1 after the embedding 645and after the MLP) and at the last word late (the "in" column is 1 after 646attention). The diagonal hand-off between them is the two-step circuit we 647built, found by intervention alone. 648 649Compare with the lens grid above. After layer 2, the lens reads "Paris" 650at the Eiffel position with probability 0.995, yet patching that state 651restores **0.00**: after the last layer nothing reads that position any 652more. The information is there, and it no longer matters. Readable is not 653the same as used, again. 654 655**Why it matters in practice.** This is how researchers located where 656GPT-style models store facts: causal tracing on prompts like ours found 657factual recall concentrated in middle-layer MLPs at the subject's last 658token, and then used that to edit single facts. The same method, one site 659at a time, has traced whole circuits, such as the one GPT-2 small uses to 660fill in "When Mary and John went to the store, John gave a drink to" → 661"Mary". 662 663**In code:** `LandmarkModel` is the hand-built model and `LandmarkModel.run` accepts patches; `fraction_restored` is the formula and `patching_map` patches every layer and position in turn. 664 665## Superposition: more features than neurons 666 667**Everyday picture.** Back to the band, but now five instruments play 668through two speakers. The engineer pans each instrument to its own 669position around the room: hard left, front right, back left, and so on. 670When one instrument plays alone, you can tell which one by where the sound 671comes from. When two play at once, the positions blur together and you 672might mistake the pair for a third instrument. The trick works because in 673this band, most of the time, only one instrument is playing. 674 675Models face the same squeeze. There are far more concepts in the world 676than neurons in a layer, but in any one sentence almost all of them are 677absent: features are **sparse**. So models store more features than they 678have dimensions, as directions that are *nearly* perpendicular rather than 679exactly. That is **superposition**. 680 681**A tiny worked example.** Put five features in two neurons, spread 72° 682apart like the points of a pentagon. Feature $k$'s direction is 683(cos 72k°, sin 72k°): feature 0 is (1, 0), feature 1 is (0.309, 0.951), 684and so on. Switch on feature 0 alone, at strength 1. The hidden state is 685(1, 0). Reading every feature back with a dot product gives 686 687| Feature | Angle from feature 0 | Read-back (cos of the angle) | After adding −0.31 and ReLU | 688|---|---|---|---| 689| 0 | 0° | **1.000** | 0.69 | 690| 1 | 72° | 0.309 | 0 | 691| 2 | 144° | −0.809 | 0 | 692| 3 | 216° | −0.809 | 0 | 693| 4 | 288° | 0.309 | 0 | 694 695Five directions can't all be perpendicular in two dimensions, so reading 696feature 0 leaks 0.309 into each neighbour: **interference**. The fix is a 697small negative bias and a **ReLU** (which turns negatives into zero): the 698leak of 0.309 falls below the bias of 0.31 and is filtered to 0, and the 699real feature survives at 0.69. The cost comes when two neighbours are on 700at once: each then reads 1 + 0.309 − 0.31 = 0.999 instead of 0.69, too 701high. Sparsity is a bet that such collisions are rare. 702 703$$ 704\hat{x} = \text{ReLU}\big(W^\top W x + b\big), 705\qquad 706\hat{x}_i = \text{ReLU}\Big(\|w_i\|^2 x_i + \sum_{j \neq i} (w_i \cdot w_j)\, x_j + b_i\Big) 707$$ 708 709**Symbols** 710 711| Symbol | Meaning here | In the example | 712|---|---|---| 713| $x$ | the true features: one strength per feature | (1, 0, 0, 0, 0) | 714| $W$ | the squeeze: 2 rows (neurons) by 5 columns (features); column $i$ is feature $i$'s direction $w_i$ | the pentagon | 715| $W x$ | the hidden state: 2 numbers holding 5 features | (1, 0) | 716| $W^\top$ | $W$ **transposed** (rows and columns swapped): reads each feature back with a dot product | | 717| $b$ | one bias per feature, learned; negative values filter small leaks | −0.31 each | 718| ReLU | keep positives, turn negatives into 0 | | 719| $\hat{x}$ | the features read back out | (0.69, 0, 0, 0, 0) | 720| $\|w_i\|^2$ | feature $i$'s direction dotted with itself (its squared length) | 1 | 721| $w_i \cdot w_j$ | how much feature $j$ leaks into feature $i$: 0 if perpendicular | 0.309 for neighbours | 722| $\sum_{j \neq i}$ | add up over every *other* feature $j$ | | 723 724**In words:** "squeeze the features into a few neurons, read each one back 725by dotting with its direction, subtract a threshold and drop anything 726negative. What you read for feature $i$ is its own strength, plus a leak 727from every other active feature, minus the threshold." 728 729**With the numbers:** with feature 0 alone on, $\hat{x}_1 = 730\text{ReLU}(1 \cdot 0 + 0.309 \cdot 1 - 0.31) = \text{ReLU}(-0.001) = 0$ 731and $\hat{x}_0 = \text{ReLU}(1 \cdot 1 - 0.31) = 0.69$. With features 0 732and 1 both on, $\hat{x}_0 = \text{ReLU}(1 + 0.309 - 0.31) = 0.999$. 733 734**In Python:** 735 736```python 737import math 738# feature k's direction: (cos 72k°, sin 72k°) 739W = [[math.cos(math.radians(72 * k)), math.sin(math.radians(72 * k))] for k in range(5)] 740def read_back(x, bias, relu=True): 741 # h = W x: two numbers 742 h = [sum(W[k][n] * x[k] for k in range(5)) for n in range(2)] 743 # Wᵀ h + b: one number per feature 744 z = [W[k][0] * h[0] + W[k][1] * h[1] + bias for k in range(5)] 745 # ReLU keeps positives and turns negatives into 0 746 return [round(max(0.0, v) if relu else v, 3) for v in z] 747# the leak: feature 0 alone, no bias, no ReLU 748read_back([1, 0, 0, 0, 0], bias=0, relu=False) # → [1.0, 0.309, -0.809, -0.809, 0.309] 749# the bias and ReLU filter it 750read_back([1, 0, 0, 0, 0], bias=-0.31) # → [0.69, 0.0, 0.0, 0.0, 0.0] 751# two neighbours on: each reads too high 752read_back([1, 1, 0, 0, 0], bias=-0.31) # → [0.999, 0.999, 0.0, 0.0, 0.0] 753``` 754 755```mermaid 756flowchart LR 757 X["5 features<br/>x₀ … x₄<br/>(mostly zero)"] --> W["squeeze: W<br/>(2 × 5)"] 758 W --> H["hidden state<br/>2 neurons"] 759 H --> WT["read back: Wᵀ<br/>(5 × 2)"] 760 WT --> B["+ bias b<br/>(negative)"] 761 B --> R["ReLU"] 762 R --> XH["5 features<br/>read back"] 763``` 764 765**Reading it:** follow a feature through the bottleneck. Five numbers go 766in, are squeezed into two, and must be expanded back into five. There is 767no way to fit five perpendicular directions into two neurons, so the 768read-back always leaks; the bias and ReLU at the end are what make the 769leak survivable, by throwing away small readings. This is the toy model 770from Anthropic's *Toy Models of Superposition* (2022), and the next step 771is to train it and see what it chooses. 772 773**Training it.** Let the model learn $W$ and $b$ itself, by gradient 774descent on the reconstruction error, with features that matter less and 775less (feature $i$ is weighted $0.8^i$). Compare two worlds: 776 777- **Dense**: every feature is on in every example. The model keeps the two 778 most important features, perpendicular, at length 1.00, and gives the 779 other three length 0: it simply drops them. 780- **Sparse**: each feature is on only 5% of the time. The model keeps 781 **all five**, 72° apart, at length about 1.1, and learns a negative bias 782 of about −0.23 to filter the leaks. 783 784 785 786**Reading it:** in the two left panels each arrow is one feature's learned 787direction in the two-neuron space, and its length is how well the model 788stores it. Dense features (left) get the textbook answer: as many features 789as neurons, perpendicular, the rest discarded. Sparse features (middle) 790get a pentagon: five features in two dimensions, accepting a little 791interference in exchange for storing everything. The right panel reads the 792ideal pentagon by *neuron*: neuron 1 moves for feature 0 (by 1.0), feature 7931 and feature 4 (by 0.309 each), and moves the other way for features 2 794and 3. No neuron belongs to one feature. 795 796That last point is why looking at single neurons so often fails. A neuron 797that responds to several unrelated features is **polysemantic**. Vision 798researchers found neurons that respond to cat faces and to the fronts of 799cars; language models are full of neurons like that. Superposition 800explains why: when features outnumber neurons, the features cannot line up 801with the neurons, so every neuron is a mixture. 802 803**In code:** `pentagon` builds the five directions, `superposition_readout` is the formula, `train_superposition` learns W and b with gradients from `superposition_loss_and_grads` (checked against `finite_difference_gradient`), and `features_a_neuron_responds_to` lists a neuron's features. 804 805## Sparse autoencoders: getting the features back 806 807**Everyday picture.** A sound engineer receives the two-speaker recording 808of the five-instrument band, with no notes on who played when. She knows 809one thing about this band: usually only one instrument plays at a time. 810Many different scores could produce the same sound, but she writes down 811the one that uses the *fewest instruments*. That preference is what lets 812her recover the real parts rather than some arbitrary mixture. 813 814A **sparse autoencoder** (SAE) does this for a model's hidden states. It 815learns a **dictionary**: many more directions than the layer has neurons 816(it is *overcomplete*), trained so that every hidden state can be rebuilt 817from just a few of them. Each learned direction is a candidate feature, 818and each one's strength on an input is its **latent** activation. 819 820**A tiny worked example.** The hidden state is h = (0.8, 0.6) and the 821dictionary has three directions, (1, 0), (0, 1) and (0.8, 0.6). Two codes 822rebuild h perfectly: 823 824| Code (strength of each direction) | Rebuilt | Error | Penalty with λ = 0.1 | Loss | 825|---|---|---|---|---| 826| A: (0.8, 0.6, 0), two directions | (0.8, 0.6) | 0 | 0.1 × (0.8 + 0.6) = 0.14 | 0.14 | 827| B: (0, 0, 1), one direction | (0.8, 0.6) | 0 | 0.1 × 1.0 = 0.10 | **0.10** | 828 829Rebuilding alone can't choose between them. The penalty on the total 830strength, the **L1 penalty** (`primer.ml.regularization`), prefers the code 831that uses one direction. One side effect: the best strength for direction 8323 is not quite 1. With strength $a$, the loss is $(1 - a)^2 + 0.1a$, which 833is lowest at $a = 0.95$ (loss 0.0975): the penalty always pulls strengths 834a little toward zero, known as **shrinkage**. 835 836$$ 837f = \text{ReLU}\big(W_e (h - b_d) + b_e\big), 838\qquad 839\hat{h} = W_d f + b_d, 840\qquad 841L = \|h - \hat{h}\|^2 + \lambda \sum_{i=1}^{m} |f_i| 842$$ 843 844**Symbols** 845 846| Symbol | Meaning here | In the example | 847|---|---|---| 848| $h$ | a hidden state from the model under study | (0.8, 0.6) | 849| $m$ | how many latents (dictionary directions) the SAE has, usually far more than neurons | 3 | 850| $W_e$, $b_e$ | the encoder's weights and biases: turn a hidden state into latent strengths | | 851| $f$ | the code: one strength per latent, ReLU keeps them ≥ 0 and mostly exactly 0 | (0, 0, 1) | 852| $W_d$ | the decoder: column $i$ is latent $i$'s direction, kept at length 1 | (1, 0), (0, 1), (0.8, 0.6) | 853| $b_d$ | the decoder bias: the typical hidden state, subtracted before encoding and added back after | (0, 0) | 854| $\hat{h}$ | the rebuilt hidden state | (0.8, 0.6) | 855| $\|h - \hat{h}\|^2$ | squared length of the rebuild error: add up the squares of its entries | 0 | 856| $\lambda$ | the sparsity penalty's strength, chosen by you | 0.1 | 857| $\lvert f_i \rvert$ | the size of latent $i$'s strength | | 858| $L$ | the loss to minimize: rebuild error plus penalty | 0.10 | 859 860**In words:** "encode the hidden state into many non-negative strengths, 861decode it back as a weighted sum of dictionary directions, and pay for 862both the rebuild error and the total strength used." 863 864**With the numbers:** for code B, $\hat{h} = 1 \cdot (0.8, 0.6) = (0.8, 8650.6)$, the error is 0, the penalty is $0.1 \times 1 = 0.10$, so $L = 0.10$. 866For code A the penalty is $0.1 \times 1.4 = 0.14$. 867 868**In Python:** 869 870```python 871dictionary = [[1, 0], [0, 1], [0.8, 0.6]] 872h = [0.8, 0.6] 873lam = 0.1 874def loss(code): 875 # ĥ = Σ f_i · direction_i 876 rebuilt = [sum(f * d[n] for f, d in zip(code, dictionary)) for n in range(2)] 877 # ‖h − ĥ‖² + λ Σ |f_i| 878 return round(sum((a - b) ** 2 for a, b in zip(h, rebuilt)) + lam * sum(abs(f) for f in code), 4) 879loss([0.8, 0.6, 0]) # → 0.14 880loss([0, 0, 1.0]) # → 0.1 881# shrinkage: a slightly weaker code is cheaper still 882loss([0, 0, 0.95]) # → 0.0975 883``` 884 885```mermaid 886flowchart LR 887 H["hidden state h<br/>(d numbers)"] --> ENC["encoder<br/>ReLU(W_e(h − b_d) + b_e)"] 888 ENC --> F["code f<br/>(m ≫ d numbers,<br/>almost all exactly 0)"] 889 F --> DEC["decoder<br/>W_d f + b_d"] 890 DEC --> HH["rebuilt ĥ"] 891 HH --> LOSS["loss = ‖h − ĥ‖² + λ Σ|f|"] 892 H --> LOSS 893 F --> LOSS 894``` 895 896**Reading it:** the hidden state is blown up into a much wider code and 897squeezed back down. The loss watches two things at once: the rebuild must 898match the original (the arrow from h), and the code must be small (the 899arrow from f). Without the second arrow the SAE could use any directions 900at all. With it, each input is explained by a handful of latents, and each 901latent's decoder column is a candidate feature you can inspect: read which 902inputs make it fire, and add its direction to the model to see what it 903does. 904 905**On the superposition toy.** Take 20,000 hidden states from the pentagon 906model (each feature on 5% of the time), train an SAE with 5 latents and 907λ = 0.3, and compare the learned decoder columns with the five planted 908directions. Every planted feature is matched by a latent with cosine 909similarity above **0.99** (1 would be the same direction), and an example 910with a feature on lights up about **1.07** latents on average. The price 911is shrinkage: **11%** of the variance goes unexplained. With a tiny 912penalty, λ = 0.01, the SAE rebuilds the data perfectly (0.0% unexplained), 913yet its worst match to a real feature is only **0.80**: any two directions 914can rebuild a two-dimensional space, so without sparsity pressure nothing 915forces the dictionary to find the model's actual features. 916 917 918 919**Reading it:** on the left, each grey dot is one hidden state. Because 920features are sparse, most dots lie along one of five rays: one feature on, 921at some strength. The few dots off the rays are the rare examples with two 922features on at once. The thick grey lines are the directions we planted; 923the red arrows are what the SAE learned from the dots alone, and they land 924on the rays. On the right, the x-axis is the penalty λ on a log scale. The 925blue line is the worst match between a planted feature and its nearest 926latent: it sits around 0.8 to 0.9 for small penalties and peaks near 0.99 927at λ = 0.3. The red line, the variance left unexplained, is near 0 for 928small penalties and climbs past 0.5 at λ = 1, where the penalty starts 929crushing the codes and the match falls again. Choosing λ is a trade: too 930small and the latents are not the features, too large and too much of the 931model's activity is lost. 932 933**Why it matters in practice.** This is the tool that turned 934interpretability from "a few hand-picked neurons" into something that 935scales. Anthropic's *Towards Monosemanticity* (2023) trained SAEs on a 936small transformer and found thousands of features that each fire for one 937recognizable thing. *Scaling Monosemanticity* (2024) did the same inside a 938production model, Claude 3 Sonnet, and found features for concepts such as 939the Golden Gate Bridge. Turning that one feature up made the model bring 940up the bridge in almost every answer, which is the causal check that the 941feature means what it seems to mean. 942 943**In code:** `sae_objective` computes the loss for one proposed code, `SparseAutoencoder` holds the encoder and decoder with its hand-derived `SparseAutoencoder.loss_and_grads`, `train_sae` fits it, and `match_features` compares learned directions with planted ones. 944 945## What these tools cannot (yet) show 946 947**Everyday picture.** A brain scan shows which regions light up while a 948person reads, but not the sentence they are thinking. Every tool in this 949lesson is like that: a real instrument that measures something true and 950partial. 951 952**A tiny worked example.** Each of our own toys already showed a limit: 953 954| What we saw | The limit it shows | 955|---|---| 956| A probe read tense at 98%, and the output ignored tense completely | Probes find information, not use | 957| The lens read "Paris" at a position where patching restored 0% | Reading an activation doesn't mean anything downstream reads it | 958| The weak SAE rebuilt everything perfectly and still found the wrong directions | A good rebuild score doesn't mean the features are real | 959| The strong SAE left 11% of the variance unexplained | Some of the model's activity is in no feature the dictionary found | 960| We knew the landmark model's circuit because we built it | In a real model, nobody has the answer key | 961 962```mermaid 963flowchart LR 964 A["Correlation<br/>probes, the logit lens:<br/>'this information is here'"] --> B["Intervention<br/>patching, steering:<br/>'changing this changes the output'"] 965 B --> C["Mechanism<br/>a circuit explained end to end:<br/>'this is how the output is computed'"] 966``` 967 968**Reading it:** the arrow is the direction of stronger evidence. Most 969claims about real models sit on the left or in the middle. A complete 970mechanism, every step from input to output accounted for, exists only for 971narrow behaviours in small or carefully chosen settings. When you read an 972interpretability result, ask which box it reached. 973 974The main open problems, in plain words: 975 976- **Choices shape answers.** Patching results depend on how the corrupted 977 prompt is built (swap one word? add noise?). SAE features depend on the 978 dictionary's size and λ: a bigger dictionary can split one feature into 979 several finer ones. 980- **Redundancy hides importance.** Patching one site at a time can miss 981 parts that back each other up. Models have been observed to partly 982 repair themselves when one component is knocked out, so "patching this 983 restores 0%" does not always mean "this plays no role". 984- **Coverage.** The unexplained part of an SAE's rebuild is not noise to 985 ignore: it is model behaviour no feature describes yet. 986- **Scale and labour.** A circuit that explains one behaviour of a small 987 model can take researchers weeks. No one has a complete account of how a 988 frontier model produces any long answer. 989- **Labels are human guesses.** Naming a feature "Golden Gate Bridge" 990 summarizes the inputs that make it fire; checking that the name is right 991 needs interventions like the one above. 992 993**Why it matters in practice.** Treat these tools as evidence, stacked 994next to behavioural evaluations (`primer.agents.evals`), never as a 995certificate. A probe or a lens is a cheap first look; a patching or 996steering experiment is the minimum for a causal claim; and a clean result 997on a toy, like every result in this lesson, is where understanding starts, 998not where it ends. 999 1000## In 20 seconds 1001 1002- **Features are directions** in activation space, not individual neurons; 1003 a dot product with the right direction reads a feature out. 1004- **Probes** are small linear classifiers trained on frozen activations. 1005 They show information is present, not that the model uses it; score 1006 them on unseen data and against a control task. 1007- **The logit lens** applies the model's own output layer to intermediate 1008 residual states, showing what it would predict after each layer. 1009- **Activation patching** copies one activation from a clean run into a 1010 corrupted run and measures how much of the right answer returns: the 1011 basic causal experiment. 1012- **Superposition**: when features are sparse, models store more features 1013 than they have neurons, as nearly perpendicular directions, which makes 1014 neurons polysemantic. 1015- **Sparse autoencoders** learn an overcomplete dictionary with an L1 1016 penalty so each input uses a few latents, recovering features from 1017 superposition, at the cost of some unexplained variance. 1018 1019## Self-test questions 1020 1021**What does it mean to say a feature is a "direction" rather than a neuron?** 1022The feature's presence is stored as a pattern across many neurons: the 1023hidden state moves along a particular direction in proportion to how 1024strongly the feature is present. You read it with a dot product against 1025that direction. Any single neuron is typically a mix of several features. 1026 1027**A linear probe reads a property from layer 12 at 95% accuracy. What can you conclude, and what can't you?** 1028You can conclude the property is linearly decodable from layer 12, as long 1029as the 95% was on held-out examples and clearly beats a control task with 1030random labels. You can't conclude the model uses it. For that you need an 1031intervention: remove or change that information and see whether the 1032output changes. 1033 1034**How does the logit lens work, and why does it agree with the model at the last layer?** 1035It takes the residual stream after some layer and applies the model's own 1036final normalization and unembedding, turning it into next-token 1037probabilities. At the last layer that is exactly the computation the model 1038itself performs, so the two must match. Earlier layers are read "as if 1039the model stopped there". 1040 1041**Describe an activation patching experiment, and what the "fraction restored" means.** 1042Run a clean prompt and save its activations; run a corrupted prompt that 1043changes one fact; then rerun the corrupted prompt with one activation 1044replaced by its clean value. The fraction restored is how much of the 1045clean-minus-corrupted logit difference the patch brings back: 1 means that 1046activation alone carries the fact, 0 means it carries none of it. 1047 1048**Why do models use superposition, and why does it make neurons hard to interpret?** 1049There are more useful features than neurons. When features are rarely 1050active at the same time, the model can store them as nearly perpendicular 1051directions and filter the small interference with a bias and ReLU. The 1052directions can't line up with the neurons, so each neuron responds to 1053several unrelated features: it is polysemantic. 1054 1055**Why does a sparse autoencoder need the L1 penalty? What goes wrong if λ is too small or too large?** 1056Many dictionaries rebuild the data equally well; the penalty picks the one 1057where each input uses few latents, which pushes latents onto the real 1058features. Too small and the SAE rebuilds perfectly with meaningless 1059directions; too large and it shrinks activations and leaves much of the 1060model's activity unexplained. 1061 1062## The papers behind this lesson 1063 1064- **Alain & Bengio, *Understanding intermediate layers using linear 1065 classifier probes* (2016)**: https://arxiv.org/abs/1610.01644. 1066 Introduced linear probes as a way to measure what each layer of a 1067 network makes linearly available. 1068- **Hewitt & Liang, *Designing and Interpreting Probes with Control Tasks* 1069 (2019)**: https://arxiv.org/abs/1909.03368. Showed that probes can 1070 succeed by memorizing, and introduced control tasks and selectivity to 1071 tell the two apart. 1072- **Belrose et al., *Eliciting Latent Predictions from Transformers with 1073 the Tuned Lens* (2023)**: https://arxiv.org/abs/2303.08112. Formalized 1074 the logit lens and fixed its failures on early layers by training a 1075 small translator for each layer. 1076- **Meng et al., *Locating and Editing Factual Associations in GPT* 1077 (2022)**: https://arxiv.org/abs/2202.05262. Introduced causal tracing, 1078 found factual recall in middle-layer MLPs at the subject's last token, 1079 and edited single facts there. 1080 [Annotated companion](../../papers/rome-causal-tracing.html) 1081- **Wang et al., *Interpretability in the Wild: a Circuit for Indirect 1082 Object Identification in GPT-2 small* (2022)**: 1083 https://arxiv.org/abs/2211.00593. Used patching to reverse-engineer a 1084 complete circuit for one behaviour of a real language model. 1085- **Elhage et al., *Toy Models of Superposition* (2022)**: 1086 https://arxiv.org/abs/2209.10652. Showed with small ReLU models that 1087 sparse features are stored in superposition, and when and how the 1088 geometry changes. 1089 [Annotated companion](../../papers/toy-models-of-superposition.html) 1090- **Bricken et al., *Towards Monosemanticity: Decomposing Language Models 1091 With Dictionary Learning* (2023)**: 1092 https://transformer-circuits.pub/2023/monosemantic-features/index.html. 1093 Trained sparse autoencoders on a small transformer and found thousands 1094 of interpretable features hidden in polysemantic neurons. 1095 [Annotated companion](../../papers/towards-monosemanticity.html) 1096- **Templeton et al., *Scaling Monosemanticity: Extracting Interpretable 1097 Features from Claude 3 Sonnet* (2024)**: 1098 https://transformer-circuits.pub/2024/scaling-monosemanticity/index.html. 1099 Scaled sparse autoencoders to a production model and steered its 1100 behaviour through the features they found. 1101 [Annotated companion](../../papers/scaling-monosemanticity.html) 1102 1103## Further reading 1104 1105- Olah et al., *Zoom In: An Introduction to Circuits* (Distill, 2020): https://distill.pub/2020/circuits/zoom-in/ 1106- Elhage et al., *A Mathematical Framework for Transformer Circuits* (2021): https://transformer-circuits.pub/2021/framework/index.html 1107- Elhage et al., *Toy Models of Superposition* (2022): https://transformer-circuits.pub/2022/toy_model/index.html 1108- Bricken et al., *Towards Monosemanticity* (2023): https://transformer-circuits.pub/2023/monosemantic-features/index.html 1109- Templeton et al., *Scaling Monosemanticity* (2024): https://transformer-circuits.pub/2024/scaling-monosemanticity/index.html 1110- Cunningham et al., *Sparse Autoencoders Find Highly Interpretable Features in Language Models* (2023): https://arxiv.org/abs/2309.08600 1111- Belinkov, *Probing Classifiers: Promises, Shortcomings, and Advances* (2021): https://arxiv.org/abs/2102.12452 1112- Meng et al., *Locating and Editing Factual Associations in GPT* (2022): https://arxiv.org/abs/2202.05262 1113- Belrose et al., *Eliciting Latent Predictions from Transformers with the Tuned Lens* (2023): https://arxiv.org/abs/2303.08112 1114""" 1115 1116from __future__ import annotations 1117 1118from dataclasses import dataclass 1119 1120import numpy as np 1121 1122from primer._show import banner, matrix, say, table, takeaway 1123from primer.ml.attention import causal_mask, scaled_dot_product_attention, softmax 1124from primer.ml.neural_net import sigmoid 1125from primer.ml.optimizers import Adam 1126from primer.ml.transformer import TinyGPT, layer_norm 1127 1128# --------------------------------------------------------------------------- 1129# 1. Features are directions 1130# --------------------------------------------------------------------------- 1131 1132# Two perpendicular unit directions in a 2-neuron hidden state (0.6² + 0.8² = 1, and 0.6·0.8 − 0.8·0.6 = 0). 1133WORKED_DIRECTIONS: dict[str, np.ndarray] = { 1134 "positive": np.array([0.6, 0.8]), 1135 "past tense": np.array([0.8, -0.6]), 1136} 1137 1138 1139def compose_features(amounts: dict[str, float], directions: dict[str, np.ndarray]) -> np.ndarray: 1140 """Build a hidden state as Σ amount · direction, one term per feature.""" 1141 return sum(amount * directions[name] for name, amount in amounts.items()) 1142 1143 1144def read_features(h: np.ndarray, directions: dict[str, np.ndarray]) -> dict[str, float]: 1145 """Read each feature back with a dot product. Exact only when the directions are perpendicular unit vectors.""" 1146 return {name: float(np.dot(d, h)) for name, d in directions.items()} 1147 1148 1149# --------------------------------------------------------------------------- 1150# 2. Probes: can a straight line read a property out of the hidden state? 1151# --------------------------------------------------------------------------- 1152 1153 1154def probe_probability(h, w, b: float) -> float: 1155 """One probe reading: σ(w · h + b), the probe's probability that the property is present.""" 1156 return float(sigmoid(np.dot(w, h) + b)) 1157 1158 1159@dataclass 1160class Probe: 1161 """A logistic-regression probe: a weight per neuron and one bias.""" 1162 1163 w: np.ndarray 1164 b: float 1165 1166 def probability(self, H: np.ndarray) -> np.ndarray: 1167 return sigmoid(H @ self.w + self.b) 1168 1169 def accuracy(self, H: np.ndarray, y: np.ndarray) -> float: 1170 return float(np.mean((self.probability(H) > 0.5) == (y == 1))) 1171 1172 1173def train_probe(H: np.ndarray, y: np.ndarray, steps: int = 500, lr: float = 0.5) -> Probe: 1174 """Fit a probe by plain gradient descent on binary cross-entropy. 1175 1176 The model under study is frozen: only the probe's d + 1 numbers change. 1177 The gradient of the mean cross-entropy with respect to the logit is 1178 simply (p − y), which is why the update is one line. 1179 """ 1180 w, b = np.zeros(H.shape[1]), 0.0 1181 for _ in range(steps): 1182 error = sigmoid(H @ w + b) - y # (n,): how far each prediction is from its label 1183 w -= lr * H.T @ error / len(y) 1184 b -= lr * float(error.mean()) 1185 return Probe(w, b) 1186 1187 1188class PlantedModel: 1189 """A toy model whose hidden layer stores two known features as directions. 1190 1191 Each example has a sentiment (0 = negative, 1 = positive) and a tense 1192 (0 = present, 1 = past). The hidden state is 1193 1194 h = (±1) · sentiment_direction + (±1) · tense_direction + noise 1195 1196 in `d` neurons, with both directions random, so no single neuron is 1197 either feature. The model's output reads sentiment and, by construction, 1198 gives tense a weight of exactly zero: the tense information is in the 1199 hidden state, but nothing downstream uses it. 1200 """ 1201 1202 FEATURES = ("sentiment", "tense") 1203 1204 def __init__(self, d: int = 32, noise: float = 0.4, seed: int = 0): 1205 rng = np.random.default_rng(seed) 1206 A = rng.standard_normal((d, 2)) 1207 A /= np.linalg.norm(A, axis=0) # unit directions, not perpendicular to each other 1208 self.directions = dict(zip(self.FEATURES, A.T)) 1209 # The pseudo-inverse's first row is the readout with weight 1 on sentiment and 0 on tense. 1210 self.readout = np.linalg.pinv(A)[0] 1211 self.d, self.noise = d, noise 1212 1213 def hidden(self, labels: dict[str, np.ndarray], noise: np.ndarray) -> np.ndarray: 1214 """(n, d) hidden states for 0/1 labels, with the noise passed in so a flip changes only the label.""" 1215 signs = {name: 2 * np.asarray(labels[name]) - 1 for name in self.FEATURES} # 0/1 -> −1/+1 1216 return sum(np.outer(signs[name], self.directions[name]) for name in self.FEATURES) + noise 1217 1218 def output(self, H: np.ndarray) -> np.ndarray: 1219 """The model's own score for each example: positive means "positive review".""" 1220 return H @ self.readout 1221 1222 def sample(self, n: int, seed: int = 0) -> tuple[np.ndarray, dict[str, np.ndarray], np.ndarray]: 1223 rng = np.random.default_rng(seed) 1224 labels = {name: rng.integers(0, 2, n) for name in self.FEATURES} 1225 noise = rng.normal(0, self.noise, (n, self.d)) 1226 return self.hidden(labels, noise), labels, noise 1227 1228 1229def flip_effect(model: PlantedModel, feature: str, n: int = 200, seed: int = 0) -> float: 1230 """Average change in the model's output when only `feature` is flipped, noise held fixed. 1231 1232 This is an intervention, not a reading: it asks whether the model *uses* 1233 the feature, which no probe can answer. 1234 """ 1235 H, labels, noise = model.sample(n, seed) 1236 flipped = dict(labels, **{feature: 1 - labels[feature]}) 1237 return float(np.mean(np.abs(model.output(model.hidden(flipped, noise)) - model.output(H)))) 1238 1239 1240def probe_report(n_train: int = 100, n_test: int = 1000, seed: int = 0) -> dict[str, dict[str, float]]: 1241 """Train three probes on the same frozen hidden states and score each on its own training set and on unseen examples. 1242 1243 * sentiment: stored and used by the model 1244 * tense: stored but never used by the model's output 1245 * random labels: the control task, with nothing real to find 1246 """ 1247 model = PlantedModel() 1248 H_train, y_train, _ = model.sample(n_train, seed) 1249 H_test, y_test, _ = model.sample(n_test, seed + 1) 1250 rng = np.random.default_rng(seed + 2) 1251 y_train["random labels"] = rng.integers(0, 2, n_train) 1252 y_test["random labels"] = rng.integers(0, 2, n_test) 1253 report = {} 1254 for name in ("sentiment", "tense", "random labels"): 1255 probe = train_probe(H_train, y_train[name]) 1256 report[name] = {"train": probe.accuracy(H_train, y_train[name]), "test": probe.accuracy(H_test, y_test[name])} 1257 return report 1258 1259 1260# --------------------------------------------------------------------------- 1261# 3. The logit lens: read the residual stream with the model's own output layer 1262# --------------------------------------------------------------------------- 1263 1264LENS_VOCAB = ["Paris", "London", "Rome"] 1265# The output layer of the worked example: one row per word, dotted with the 2-number residual state. 1266LENS_UNEMBED = np.array([[1.0, 0.0], [0.0, 1.0], [-1.0, 0.0]]) 1267LENS_START = np.array([0.2, 0.3]) # the residual state straight after the embedding 1268LENS_WRITES = [np.array([0.6, 0.0]), np.array([1.2, -0.3])] # what layers 1 and 2 add to it 1269 1270 1271def logit_lens(residuals: np.ndarray, W_U: np.ndarray, final_norm=None) -> np.ndarray: 1272 """Probabilities the model would give if it stopped here: softmax(W_U · norm(h)) for every state in `residuals`.""" 1273 x = final_norm(residuals) if final_norm is not None else residuals 1274 return softmax(x @ W_U.T, axis=-1) 1275 1276 1277def worked_lens() -> list[dict]: 1278 """The worked example: the lens after the embedding, after layer 1 and after layer 2.""" 1279 states = [LENS_START] 1280 for write in LENS_WRITES: 1281 states.append(states[-1] + write) # a residual layer adds; it never overwrites 1282 probs = logit_lens(np.array(states), LENS_UNEMBED) 1283 return [ 1284 dict(layer=i, residual=s, logits=LENS_UNEMBED @ s, probs=p.tolist(), top=LENS_VOCAB[int(np.argmax(p))]) 1285 for i, (s, p) in enumerate(zip(states, probs)) 1286 ] 1287 1288 1289def residual_stream(model: TinyGPT, ids: np.ndarray) -> np.ndarray: 1290 """(n_layers + 1, seq, d): the residual state after the embedding and after every block.""" 1291 ids = np.asarray(ids) 1292 x = model.wte[ids] + model.wpe[: len(ids)] 1293 states = [x] 1294 for block in model.blocks: 1295 x = block(x) 1296 states.append(x) 1297 return np.stack(states) 1298 1299 1300def tinygpt_logit_lens(model: TinyGPT, ids: np.ndarray) -> np.ndarray: 1301 """(n_layers + 1, seq, vocab) logits: every layer's residual state read through the final norm and the output layer. 1302 1303 Returns logits rather than probabilities so the top layer can be 1304 compared number for number with the model's own output. 1305 """ 1306 R = residual_stream(model, ids) 1307 return layer_norm(R, model.lnf_g, model.lnf_b) @ model.wte.T 1308 1309 1310# --------------------------------------------------------------------------- 1311# 4. A hand-built model where we know where the answer flows 1312# --------------------------------------------------------------------------- 1313 1314 1315class LandmarkModel: 1316 """A two-layer, one-head transformer built by hand to answer "<landmark> is in" -> city. 1317 1318 Residual dimensions (a readable basis, so you can check it by eye; no 1319 tool in this lesson relies on it): 1320 1321 0-2 which landmark (Eiffel, Colosseum, Big Ben) 1322 3 "this token is a landmark" 1323 4-5 which filler word ("is", "in") 1324 6-8 which city (Paris, Rome, London) 1325 1326 Layer 1 is an MLP that looks up a fact at every position: a landmark's 1327 identity in, its city out. Layer 2 is an attention head: the word "in" 1328 queries for the landmark flag, finds the landmark and copies its city 1329 dimensions into the last position. The output layer reads the city 1330 dimensions at the last position. 1331 """ 1332 1333 TOKENS = ["Eiffel", "Colosseum", "Big Ben", "is", "in"] 1334 CITIES = ["Paris", "Rome", "London"] 1335 STAGES = ("resid_embed", "resid_mlp", "resid_attn") 1336 D = 9 1337 1338 def __init__(self, gain: float = 3.0, sharpness: float = 4.0): 1339 D = self.D 1340 self.E = np.zeros((len(self.TOKENS), D)) 1341 for i in range(3): 1342 self.E[i, i] = 1.0 # landmark identity 1343 self.E[i, 3] = 1.0 # landmark flag 1344 self.E[3, 4] = self.E[4, 5] = 1.0 # "is", "in" 1345 # MLP: three hidden units, one per landmark; each writes `gain` onto its city. 1346 self.W_in = np.zeros((D, 3)) 1347 self.W_in[:3, :3] = np.eye(3) 1348 self.W_out = np.zeros((3, D)) 1349 self.W_out[:3, 6:9] = gain * np.eye(3) 1350 # Attention, one head of width 1: "in" asks (query), the landmark flag answers (key). 1351 self.W_q = np.zeros((D, 1)) 1352 self.W_q[5, 0] = sharpness 1353 self.W_k = np.zeros((D, 1)) 1354 self.W_k[3, 0] = sharpness 1355 self.W_v = np.zeros((D, D)) 1356 self.W_v[6:9, 6:9] = np.eye(3) # the value is the city part of the state, copied as is 1357 self.W_U = np.zeros((3, D)) 1358 self.W_U[:, 6:9] = np.eye(3) 1359 1360 def run(self, tokens: list[str], patch: dict | None = None, answer: tuple[str, str] = ("Paris", "Rome")) -> dict: 1361 """Run the model, optionally overwriting activations. 1362 1363 `patch` maps (site, position) to a vector that replaces that activation 1364 during this run. Sites: "resid_embed", "mlp_out", "resid_mlp", 1365 "attn_out", "resid_attn". Returns every stage's residual state 1366 (3, seq, D), the component outputs, the last position's logits, and 1367 the logit difference between the two `answer` cities. 1368 """ 1369 patch = patch or {} 1370 1371 def apply(site: str, acts: np.ndarray) -> np.ndarray: 1372 acts = acts.copy() 1373 for (s, pos), vector in patch.items(): 1374 if s == site: 1375 acts[pos] = vector 1376 return acts 1377 1378 ids = [self.TOKENS.index(t) for t in tokens] 1379 x = apply("resid_embed", self.E[ids]) 1380 stages = [x] 1381 mlp_out = apply("mlp_out", np.maximum(x @ self.W_in, 0) @ self.W_out) 1382 x = apply("resid_mlp", x + mlp_out) 1383 stages.append(x) 1384 mixed, weights = scaled_dot_product_attention(x @ self.W_q, x @ self.W_k, x @ self.W_v, mask=causal_mask(len(ids))) 1385 attn_out = apply("attn_out", mixed) 1386 x = apply("resid_attn", x + attn_out) 1387 stages.append(x) 1388 logits = x[-1] @ self.W_U.T # only the last position predicts the next word 1389 good, bad = (self.CITIES.index(c) for c in answer) 1390 return dict( 1391 resid=np.stack(stages), mlp_out=mlp_out, attn_out=attn_out, attn_weights=weights, 1392 logits=logits, logit_diff=float(logits[good] - logits[bad]), 1393 ) 1394 1395 1396def fraction_restored(patched: float, clean: float, corrupt: float) -> float: 1397 """How much of the clean-vs-corrupt gap one patch closes: 0 = nothing, 1 = everything.""" 1398 return (patched - corrupt) / (clean - corrupt) 1399 1400 1401def patching_map(model: LandmarkModel, clean: list[str], corrupt: list[str]) -> np.ndarray: 1402 """(stage, position) grid: patch each residual state from the clean run into the corrupted run, one at a time.""" 1403 clean_run, corrupt_run = model.run(clean), model.run(corrupt) 1404 grid = np.zeros((len(model.STAGES), len(clean))) 1405 for s, stage in enumerate(model.STAGES): 1406 for pos in range(len(clean)): 1407 patched = model.run(corrupt, patch={(stage, pos): clean_run["resid"][s, pos]}) 1408 grid[s, pos] = fraction_restored(patched["logit_diff"], clean_run["logit_diff"], corrupt_run["logit_diff"]) 1409 return grid 1410 1411 1412def lens_map(model: LandmarkModel, tokens: list[str], answer: str = "Paris") -> np.ndarray: 1413 """(stage, position) grid of the logit lens's probability for `answer`.""" 1414 resid = model.run(tokens)["resid"] 1415 return logit_lens(resid, model.W_U)[..., model.CITIES.index(answer)] 1416 1417 1418# --------------------------------------------------------------------------- 1419# 5. Superposition: more features than dimensions 1420# --------------------------------------------------------------------------- 1421 1422 1423def pentagon() -> np.ndarray: 1424 """(2, 5): five unit feature directions spread 72° apart in a 2-neuron space.""" 1425 angles = np.radians(72 * np.arange(5)) 1426 return np.stack([np.cos(angles), np.sin(angles)]) 1427 1428 1429def superposition_readout(W: np.ndarray, b: np.ndarray, x: np.ndarray, relu: bool = True) -> np.ndarray: 1430 """Squeeze features x into h = W x, then read them back out: ReLU(Wᵀ h + b).""" 1431 z = W.T @ (W @ x) + b 1432 return np.maximum(z, 0) if relu else z 1433 1434 1435def sample_sparse_features(n: int, n_features: int, p_active: float, rng: np.random.Generator) -> np.ndarray: 1436 """(n, n_features): each feature is on with probability p_active, at a strength uniform in [0, 1).""" 1437 on = rng.random((n, n_features)) < p_active 1438 return on * rng.random((n, n_features)) 1439 1440 1441def superposition_loss_and_grads(W: np.ndarray, b: np.ndarray, X: np.ndarray, importance: np.ndarray | None = None): 1442 """Importance-weighted reconstruction loss of the toy model, and its gradients, derived by hand. 1443 1444 Forward: H = X Wᵀ (batch, d), Z = H W + b (batch, n), Y = ReLU(Z) 1445 Loss: mean over the batch of Σᵢ Iᵢ (Yᵢ − Xᵢ)² 1446 W appears twice (squeeze and unsqueeze), so its gradient has two terms. 1447 """ 1448 importance = np.ones(X.shape[1]) if importance is None else importance 1449 H = X @ W.T 1450 Z = H @ W + b 1451 Y = np.maximum(Z, 0) 1452 loss = float(np.mean(np.sum(importance * (Y - X) ** 2, axis=1))) 1453 dZ = 2 * importance * (Y - X) / len(X) * (Z > 0) # ReLU passes gradient only where it was on 1454 dH = dZ @ W.T 1455 grads = {"W": H.T @ dZ + dH.T @ X, "b": dZ.sum(axis=0)} 1456 return loss, grads 1457 1458 1459def train_superposition( 1460 p_active: float, n_features: int = 5, d: int = 2, importance_decay: float = 0.8, 1461 steps: int = 2000, batch: int = 1024, lr: float = 0.01, seed: int = 0, 1462) -> dict: 1463 """Train the toy model of Elhage et al. (2022) on features that are on with probability `p_active`. 1464 1465 Feature i matters `importance_decay ** i` as much as feature 0, so when 1466 the model cannot keep everything it has a reason to choose. 1467 """ 1468 rng = np.random.default_rng(seed) 1469 importance = importance_decay ** np.arange(n_features) 1470 params = {"W": rng.normal(0, 0.5, (d, n_features)), "b": np.zeros(n_features)} 1471 opt = {k: Adam(lr=lr) for k in params} 1472 losses = [] 1473 for _ in range(steps): 1474 X = sample_sparse_features(batch, n_features, p_active, rng) 1475 loss, grads = superposition_loss_and_grads(params["W"], params["b"], X, importance) 1476 for k in params: 1477 params[k] = opt[k].step(params[k], grads[k]) 1478 losses.append(loss) 1479 return dict(params, losses=np.array(losses)) 1480 1481 1482def features_a_neuron_responds_to(W: np.ndarray, neuron: int, threshold: float = 0.25) -> list[int]: 1483 """Indices of the features that push this neuron up by more than `threshold` per unit of feature.""" 1484 return [k for k in range(W.shape[1]) if W[neuron, k] > threshold] 1485 1486 1487def finite_difference_gradient(f, w: np.ndarray, eps: float = 1e-5) -> np.ndarray: 1488 """Slow, obviously correct gradient: nudge each entry up and down and watch the loss.""" 1489 grad = np.zeros_like(w) 1490 for idx in np.ndindex(w.shape): 1491 old = w[idx] 1492 w[idx] = old + eps 1493 up = f(w) 1494 w[idx] = old - eps 1495 down = f(w) 1496 w[idx] = old 1497 grad[idx] = (up - down) / (2 * eps) 1498 return grad 1499 1500 1501# --------------------------------------------------------------------------- 1502# 6. Sparse autoencoders: learn the features back 1503# --------------------------------------------------------------------------- 1504 1505 1506def sae_objective(h: np.ndarray, f: np.ndarray, decoder: np.ndarray, lam: float, b_d: np.ndarray | float = 0.0) -> float: 1507 """The sparse autoencoder's loss for one example and one proposed code f: ‖h − (D f + b_d)‖² + λ Σ|fᵢ|.""" 1508 residual = h - (decoder @ f + b_d) 1509 return float(residual @ residual + lam * np.abs(f).sum()) 1510 1511 1512class SparseAutoencoder: 1513 """encode: f = ReLU(W_e (h − b_d) + b_e); decode: ĥ = W_d f + b_d. 1514 1515 `n_latents` is usually much bigger than `d` (overcomplete). Each decoder 1516 column is kept at length 1, so the L1 penalty cannot be dodged by making 1517 codes tiny and decoder columns huge. 1518 """ 1519 1520 def __init__(self, d: int, n_latents: int, seed: int = 0): 1521 rng = np.random.default_rng(seed) 1522 W_d = rng.standard_normal((d, n_latents)) 1523 self.W_d = W_d / np.linalg.norm(W_d, axis=0) # (d, m): one unit direction per latent 1524 self.W_e = self.W_d.T.copy() # (m, d): start each latent listening for its own direction 1525 self.b_e = np.zeros(n_latents) 1526 self.b_d = np.zeros(d) 1527 1528 def encode(self, H: np.ndarray) -> np.ndarray: 1529 return np.maximum((H - self.b_d) @ self.W_e.T + self.b_e, 0) 1530 1531 def decode(self, F: np.ndarray) -> np.ndarray: 1532 return F @ self.W_d.T + self.b_d 1533 1534 def reconstruction_error(self, H: np.ndarray) -> float: 1535 """Fraction of the data's variance the autoencoder fails to rebuild (0 = perfect).""" 1536 residual = self.decode(self.encode(H)) - H 1537 return float(np.sum(residual**2) / np.sum((H - H.mean(axis=0)) ** 2)) 1538 1539 def loss_and_grads(self, H: np.ndarray, lam: float) -> tuple[float, dict[str, np.ndarray]]: 1540 """Mean over the batch of ‖h − ĥ‖² + λ Σ fᵢ, and the gradient for every parameter, by hand.""" 1541 n = len(H) 1542 centred = H - self.b_d 1543 pre = centred @ self.W_e.T + self.b_e # (n, m) 1544 F = np.maximum(pre, 0) 1545 residual = F @ self.W_d.T + self.b_d - H # (n, d) 1546 loss = float(np.sum(residual**2) / n + lam * F.sum() / n) # F ≥ 0, so |F| = F 1547 d_hat = 2 * residual / n 1548 dF = d_hat @ self.W_d + lam / n 1549 d_pre = dF * (pre > 0) 1550 grads = { 1551 "W_d": d_hat.T @ F, 1552 "W_e": d_pre.T @ centred, 1553 "b_e": d_pre.sum(axis=0), 1554 # b_d is added back in the decoder and subtracted in the encoder. 1555 "b_d": d_hat.sum(axis=0) - (d_pre @ self.W_e).sum(axis=0), 1556 } 1557 return loss, grads 1558 1559 1560def train_sae( 1561 H: np.ndarray, n_latents: int = 5, lam: float = 0.3, steps: int = 2000, batch: int = 512, lr: float = 0.01, seed: int = 0 1562) -> SparseAutoencoder: 1563 """Fit a sparse autoencoder to activations H (n, d) with Adam, renormalizing decoder columns after every step.""" 1564 rng = np.random.default_rng(seed) 1565 sae = SparseAutoencoder(H.shape[1], n_latents, seed) 1566 names = ("W_e", "b_e", "W_d", "b_d") 1567 opt = {k: Adam(lr=lr) for k in names} 1568 for _ in range(steps): 1569 _, grads = sae.loss_and_grads(H[rng.integers(0, len(H), batch)], lam) 1570 for k in names: 1571 setattr(sae, k, opt[k].step(getattr(sae, k), grads[k])) 1572 sae.W_d /= np.linalg.norm(sae.W_d, axis=0) 1573 return sae 1574 1575 1576def match_features(true_directions: np.ndarray, learned_directions: np.ndarray) -> np.ndarray: 1577 """For each true feature (a column), the cosine similarity of the learned column that points most nearly the same way.""" 1578 t = true_directions / np.linalg.norm(true_directions, axis=0) 1579 ell = learned_directions / np.linalg.norm(learned_directions, axis=0) 1580 return (t.T @ ell).max(axis=1) 1581 1582 1583# --------------------------------------------------------------------------- 1584# 7. Figures (rendered into the HTML docs by `make figures`) 1585# --------------------------------------------------------------------------- 1586 1587CLEAN = ["Eiffel", "is", "in"] 1588CORRUPT = ["Colosseum", "is", "in"] 1589STAGE_LABELS = ["after embedding", "after layer 1 (MLP)", "after layer 2 (attention)"] 1590 1591 1592def figures() -> dict: 1593 """Plot this lesson's data. matplotlib is imported here, and only here, 1594 so the lesson itself needs nothing beyond NumPy.""" 1595 import matplotlib 1596 1597 matplotlib.use("Agg") 1598 import matplotlib.pyplot as plt 1599 1600 BLUE, RED, GREEN, MUTED = "#2563eb", "#dc2626", "#059669", "#9ca3af" 1601 figs = {} 1602 1603 # --- 1. Features as directions ----------------------------------------- 1604 fig, ax = plt.subplots(figsize=(4.6, 4.2)) 1605 amounts = {"positive": 0.9, "past tense": 0.4} 1606 h = compose_features(amounts, WORKED_DIRECTIONS) 1607 for (name, d), color in zip(WORKED_DIRECTIONS.items(), (BLUE, GREEN)): 1608 ax.annotate("", d, (0, 0), arrowprops=dict(arrowstyle="->", color=color, lw=2)) 1609 # Labels sit beyond the arrow tip, above it for the upper arrow and below it for the lower one. 1610 ax.text(d[0] * 1.1, d[1] * 1.1, f'"{name}"\ndirection', color=color, ha="center", va="bottom" if d[1] > 0 else "top", fontsize=9) 1611 foot = amounts[name] * d 1612 ax.plot([h[0], foot[0]], [h[1], foot[1]], ls=":", color=color) 1613 ax.plot(*foot, "o", color=color, ms=4) 1614 ax.annotate("", h, (0, 0), arrowprops=dict(arrowstyle="->", color=RED, lw=2.5)) 1615 ax.text(h[0] + 0.04, h[1] + 0.03, "h = (0.86, 0.48)", color=RED) 1616 ax.set_xlim(-0.2, 1.15) 1617 ax.set_ylim(-0.9, 1.1) 1618 ax.set_aspect("equal") 1619 ax.axhline(0, color=MUTED, lw=0.8) 1620 ax.axvline(0, color=MUTED, lw=0.8) 1621 ax.set_xlabel("neuron 1") 1622 ax.set_ylabel("neuron 2") 1623 ax.set_title("Two features, two directions, one hidden state") 1624 figs["directions"] = fig 1625 1626 # --- 2. Probes: what they can read, and what the model uses ------------ 1627 report = probe_report() 1628 model = PlantedModel() 1629 names = ["sentiment", "tense", "random labels"] 1630 fig, (a1, a2) = plt.subplots(1, 2, figsize=(9, 3.4)) 1631 x = np.arange(len(names)) 1632 a1.bar(x - 0.2, [report[n]["train"] for n in names], 0.4, color=MUTED, label="training examples") 1633 a1.bar(x + 0.2, [report[n]["test"] for n in names], 0.4, color=BLUE, label="unseen examples") 1634 for xi, n in zip(x, names): 1635 a1.text(xi + 0.2, report[n]["test"] + 0.02, f"{report[n]['test']:.2f}", ha="center") 1636 a1.axhline(0.5, color=RED, ls="--", lw=1, label="chance") 1637 a1.set_xticks(x, names) 1638 a1.set_ylim(0, 1.3) 1639 a1.set_ylabel("probe accuracy") 1640 a1.set_title("What a probe can read") 1641 a1.legend(frameon=False, loc="upper right", fontsize=8) 1642 effects = [flip_effect(model, f) for f in ("sentiment", "tense")] 1643 a2.bar([0, 1], effects, 0.5, color=[BLUE, MUTED]) 1644 for xi, e in enumerate(effects): 1645 a2.text(xi, e + 0.05, f"{e:.2f}", ha="center") 1646 a2.set_xticks([0, 1], ["flip sentiment", "flip tense"]) 1647 a2.set_ylim(0, 2.5) 1648 a2.set_ylabel("change in the model's output") 1649 a2.set_title("What the model actually uses") 1650 fig.tight_layout() 1651 figs["probe_vs_use"] = fig 1652 1653 # --- 3. Logit lens: worked example and the landmark model ------------- 1654 fig, (a1, a2) = plt.subplots(1, 2, figsize=(9.5, 3.6), gridspec_kw={"width_ratios": [1, 1.2]}) 1655 lens = worked_lens() 1656 layers = np.arange(len(lens)) 1657 for k, (word, color) in enumerate(zip(LENS_VOCAB, (BLUE, GREEN, MUTED))): 1658 a1.bar(layers + (k - 1) * 0.27, [row["probs"][k] for row in lens], 0.27, color=color, label=word) 1659 a1.set_xticks(layers, ["embedding", "layer 1", "layer 2"]) 1660 a1.set_ylabel("lens probability") 1661 a1.set_ylim(0, 1) 1662 a1.set_title("Worked example: the guess firms up") 1663 a1.legend(frameon=False, fontsize=8) 1664 grid = lens_map(LandmarkModel(), CLEAN) 1665 im = a2.imshow(grid, cmap="Blues", vmin=0, vmax=1) 1666 for (r, c), v in np.ndenumerate(grid): 1667 a2.text(c, r, f"{v:.2f}", ha="center", va="center", color="white" if v > 0.6 else "black") 1668 a2.set_xticks(range(3), CLEAN) 1669 a2.set_yticks(range(3), STAGE_LABELS) 1670 a2.set_title('Landmark model: P("Paris") by layer and word') 1671 a2.grid(False) 1672 fig.colorbar(im, ax=a2, fraction=0.046) 1673 fig.tight_layout() 1674 figs["logit_lens"] = fig 1675 1676 # --- 4. Activation patching map --------------------------------------- 1677 grid = patching_map(LandmarkModel(), CLEAN, CORRUPT) 1678 fig, ax = plt.subplots(figsize=(5.6, 3.4)) 1679 im = ax.imshow(grid, cmap="Greens", vmin=0, vmax=1) 1680 for (r, c), v in np.ndenumerate(grid): 1681 ax.text(c, r, f"{v:.2f}", ha="center", va="center", color="white" if v > 0.6 else "black") 1682 ax.set_xticks(range(3), [f"{a} / {b}" if a != b else a for a, b in zip(CLEAN, CORRUPT)]) 1683 ax.set_yticks(range(3), STAGE_LABELS) 1684 ax.set_xlabel("position patched (clean / corrupted word)") 1685 ax.set_title("Fraction of the answer restored by one patch") 1686 ax.grid(False) 1687 fig.colorbar(im, ax=ax, fraction=0.046) 1688 figs["patching"] = fig 1689 1690 # --- 5. Superposition: dense keeps two, sparse keeps five ------------- 1691 fig, axes = plt.subplots(1, 3, figsize=(11, 3.7), gridspec_kw={"width_ratios": [1, 1, 1.2]}) 1692 colors = [BLUE, GREEN, "#d97706", "#7c3aed", RED] 1693 for ax, p, title in ((axes[0], 1.0, "Dense features: keeps 2"), (axes[1], 0.05, "Sparse features: keeps all 5")): 1694 W = train_superposition(p_active=p)["W"] 1695 for k in range(W.shape[1]): 1696 if np.linalg.norm(W[:, k]) > 0.1: 1697 ax.annotate("", W[:, k], (0, 0), arrowprops=dict(arrowstyle="->", color=colors[k], lw=2)) 1698 ax.text(*(W[:, k] * 1.18), f"f{k}", color=colors[k], ha="center", va="center") 1699 else: # a dropped feature: an arrow this short would be all arrowhead 1700 ax.plot(*W[:, k], "o", color=colors[k], ms=4) 1701 ax.set_xlim(-1.45, 1.45) 1702 ax.set_ylim(-1.45, 1.45) 1703 ax.set_aspect("equal") 1704 ax.set_title(f"{title}\n(each feature on {p:.0%} of the time)") 1705 ax.set_xlabel("neuron 1") 1706 ax.set_ylabel("neuron 2") 1707 W = pentagon() 1708 x = np.arange(5) 1709 axes[2].bar(x - 0.2, W[0], 0.4, color=BLUE, label="neuron 1") 1710 axes[2].bar(x + 0.2, W[1], 0.4, color=MUTED, label="neuron 2") 1711 axes[2].axhline(0, color="#4b5563", lw=0.8) 1712 axes[2].set_xticks(x, [f"f{k}" for k in x]) 1713 axes[2].set_ylabel("how much the feature moves the neuron") 1714 axes[2].set_title("Each neuron answers to several features") 1715 axes[2].legend(frameon=False, fontsize=8) 1716 fig.tight_layout() 1717 figs["superposition"] = fig 1718 1719 # --- 6. Sparse autoencoder: recovering the planted features ------------ 1720 rng = np.random.default_rng(1) 1721 H = sample_sparse_features(20_000, 5, 0.05, rng) @ pentagon().T 1722 fig, (a1, a2) = plt.subplots(1, 2, figsize=(9.5, 3.9)) 1723 shown = H[np.abs(H).sum(axis=1) > 0][:600] 1724 a1.scatter(shown[:, 0], shown[:, 1], s=3, color=MUTED, alpha=0.5, label="hidden states") 1725 sae = train_sae(H, n_latents=5, lam=0.3) 1726 for k in range(5): 1727 t = pentagon()[:, k] 1728 a1.plot([0, 1.25 * t[0]], [0, 1.25 * t[1]], color=MUTED, lw=6, alpha=0.4, label="planted features" if k == 0 else None) 1729 d = sae.W_d[:, k] 1730 a1.annotate("", 1.1 * d, (0, 0), arrowprops=dict(arrowstyle="->", color=RED, lw=1.8)) 1731 a1.plot([], [], color=RED, label="learned latents") 1732 a1.set_aspect("equal") 1733 a1.set_xlim(-1.4, 1.4) 1734 a1.set_ylim(-1.4, 1.4) 1735 a1.set_title("λ = 0.3: each latent lands on a feature") 1736 a1.legend(frameon=False, fontsize=8, loc="lower left") 1737 lams = [0.003, 0.01, 0.03, 0.1, 0.2, 0.3, 0.6, 1.0] 1738 worst, error = [], [] 1739 for lam in lams: 1740 s = train_sae(H, n_latents=5, lam=lam) 1741 worst.append(match_features(pentagon(), s.W_d).min()) 1742 error.append(s.reconstruction_error(H)) 1743 a2.plot(lams, worst, "o-", color=BLUE, label="worst feature match (cosine)") 1744 a2.plot(lams, error, "s-", color=RED, label="variance left unexplained") 1745 a2.set_xscale("log") 1746 a2.set_xlabel("sparsity penalty λ") 1747 a2.set_ylim(0, 1.05) 1748 a2.set_title("The penalty trades rebuild for meaning") 1749 a2.legend(frameon=False, fontsize=8) 1750 fig.tight_layout() 1751 figs["sae"] = fig 1752 1753 return figs 1754 1755 1756# --------------------------------------------------------------------------- 1757# 8. Narrated walkthrough 1758# --------------------------------------------------------------------------- 1759 1760 1761def demo() -> None: 1762 banner("1. A feature is a direction") 1763 h = compose_features({"positive": 0.9, "past tense": 0.4}, WORKED_DIRECTIONS) 1764 say( 1765 f""" 1766 Mix 0.9 of the "positive" direction (0.6, 0.8) with 0.4 of the "past 1767 tense" direction (0.8, -0.6) and the hidden state is h = ({h[0]:.2f}, 1768 {h[1]:.2f}). Neither neuron equals either feature. A dot product with 1769 each direction reads the amounts back: 1770 """ 1771 ) 1772 table(["feature", "direction · h"], [(k, v) for k, v in read_features(h, WORKED_DIRECTIONS).items()], floatfmt=".2f") 1773 takeaway("Features live in directions, not in individual neurons.") 1774 1775 banner("2. Probes: reading a property out of the hidden state") 1776 say( 1777 f""" 1778 A probe is a tiny classifier trained on frozen hidden states. The 1779 worked probe w = (1, 0.5), b = 0 reads h = (2, -1) as 1780 σ(1.5) = {probe_probability([2.0, -1.0], [1.0, 0.5], 0.0):.3f}. Now three 1781 probes on a toy model with 32 neurons, trained on 100 examples and 1782 scored on 1,000 unseen ones: 1783 """ 1784 ) 1785 report = probe_report() 1786 table(["probe target", "train accuracy", "unseen accuracy"], [(k, v["train"], v["test"]) for k, v in report.items()], floatfmt=".3f") 1787 model = PlantedModel() 1788 say( 1789 f""" 1790 Random labels fit the training set to {report['random labels']['train']:.0%} and then 1791 score a coin flip on unseen examples: accuracy only counts on held-out 1792 data. Tense is read at {report['tense']['test']:.0%}, yet flipping it moves the 1793 model's output by {flip_effect(model, 'tense'):.2f}, while flipping sentiment moves 1794 it by {flip_effect(model, 'sentiment'):.2f}. 1795 """ 1796 ) 1797 takeaway("A probe shows information is present, not that the model uses it.") 1798 1799 banner("3. The logit lens: what would the model say if it stopped here?") 1800 table( 1801 ["after", "residual", "P(Paris)", "P(London)", "P(Rome)", "top"], 1802 [(["embedding", "layer 1", "layer 2"][r["layer"]], str(np.round(r["residual"], 2)), *r["probs"], r["top"]) for r in worked_lens()], 1803 floatfmt=".2f", 1804 ) 1805 tiny = TinyGPT(vocab_size=11, d_model=16, n_layers=3, n_heads=2, max_len=8, seed=4) 1806 ids = np.array([3, 1, 4, 1, 5]) 1807 agree = np.allclose(tinygpt_logit_lens(tiny, ids)[-1], tiny(ids)) 1808 say(f"On primer.ml.transformer.TinyGPT, the lens at the top layer equals the model's own logits: {agree}.") 1809 matrix('landmark model, lens P("Paris"): rows = layers, cols = Eiffel, is, in', lens_map(LandmarkModel(), CLEAN)) 1810 takeaway("The Eiffel position knows Paris after layer 1; the last word only learns it in layer 2.") 1811 1812 banner("4. Activation patching: which activations carry the answer?") 1813 lm = LandmarkModel() 1814 clean, corrupt = lm.run(CLEAN)["logit_diff"], lm.run(CORRUPT)["logit_diff"] 1815 say( 1816 f""" 1817 Logit difference Paris minus Rome: clean "Eiffel is in" gives {clean:+.2f}, 1818 corrupted "Colosseum is in" gives {corrupt:+.2f}. Copy one clean 1819 residual state into the corrupted run at a time and measure the 1820 fraction of the gap that comes back: 1821 """ 1822 ) 1823 matrix("fraction restored: rows = layers, cols = positions", patching_map(lm, CLEAN, CORRUPT), precision=2) 1824 takeaway("The answer sits at the subject early and at the last word late: a two-step circuit, found by intervention.") 1825 1826 banner("5. Superposition: five features in two neurons") 1827 table( 1828 ["feature", "readout when only f0 is on (no ReLU)", "with bias -0.31 and ReLU"], 1829 [ 1830 (f"f{k}", a, c) 1831 for k, (a, c) in enumerate( 1832 zip( 1833 superposition_readout(pentagon(), np.zeros(5), np.eye(5)[0], relu=False), 1834 superposition_readout(pentagon(), np.full(5, -0.31), np.eye(5)[0]), 1835 ) 1836 ) 1837 ], 1838 floatfmt=".3f", 1839 ) 1840 for p in (1.0, 0.05): 1841 W = train_superposition(p_active=p)["W"] 1842 print(f"trained with each feature on {p:>4.0%} of the time: feature lengths {np.round(np.linalg.norm(W, axis=0), 2)}") 1843 print() 1844 say( 1845 """ 1846 Dense features: the model keeps the two most important and drops the 1847 rest. Sparse features: it keeps all five, 72 degrees apart, and a 1848 negative bias filters the small leaks. In the ideal pentagon, neuron 1 answers to 1849 features 0, 1 and 4: it is polysemantic. 1850 """ 1851 ) 1852 1853 banner("6. Sparse autoencoders: getting the features back") 1854 rng = np.random.default_rng(1) 1855 H = sample_sparse_features(20_000, 5, 0.05, rng) @ pentagon().T 1856 rows = [] 1857 for lam in (0.01, 0.3): 1858 sae = train_sae(H, n_latents=5, lam=lam) 1859 rows.append((lam, sae.reconstruction_error(H), match_features(pentagon(), sae.W_d).min())) 1860 table(["λ", "variance unexplained", "worst feature match (cosine)"], rows, floatfmt=".3f") 1861 takeaway( 1862 "Any two directions can rebuild a 2-D space; only the sparsity penalty makes the " 1863 "autoencoder find the five directions the model really used." 1864 ) 1865 1866 1867if __name__ == "__main__": 1868 demo()
1140def compose_features(amounts: dict[str, float], directions: dict[str, np.ndarray]) -> np.ndarray: 1141 """Build a hidden state as Σ amount · direction, one term per feature.""" 1142 return sum(amount * directions[name] for name, amount in amounts.items())
Build a hidden state as Σ amount · direction, one term per feature.
1145def read_features(h: np.ndarray, directions: dict[str, np.ndarray]) -> dict[str, float]: 1146 """Read each feature back with a dot product. Exact only when the directions are perpendicular unit vectors.""" 1147 return {name: float(np.dot(d, h)) for name, d in directions.items()}
Read each feature back with a dot product. Exact only when the directions are perpendicular unit vectors.
1155def probe_probability(h, w, b: float) -> float: 1156 """One probe reading: σ(w · h + b), the probe's probability that the property is present.""" 1157 return float(sigmoid(np.dot(w, h) + b))
One probe reading: σ(w · h + b), the probe's probability that the property is present.
1160@dataclass 1161class Probe: 1162 """A logistic-regression probe: a weight per neuron and one bias.""" 1163 1164 w: np.ndarray 1165 b: float 1166 1167 def probability(self, H: np.ndarray) -> np.ndarray: 1168 return sigmoid(H @ self.w + self.b) 1169 1170 def accuracy(self, H: np.ndarray, y: np.ndarray) -> float: 1171 return float(np.mean((self.probability(H) > 0.5) == (y == 1)))
A logistic-regression probe: a weight per neuron and one bias.
1174def train_probe(H: np.ndarray, y: np.ndarray, steps: int = 500, lr: float = 0.5) -> Probe: 1175 """Fit a probe by plain gradient descent on binary cross-entropy. 1176 1177 The model under study is frozen: only the probe's d + 1 numbers change. 1178 The gradient of the mean cross-entropy with respect to the logit is 1179 simply (p − y), which is why the update is one line. 1180 """ 1181 w, b = np.zeros(H.shape[1]), 0.0 1182 for _ in range(steps): 1183 error = sigmoid(H @ w + b) - y # (n,): how far each prediction is from its label 1184 w -= lr * H.T @ error / len(y) 1185 b -= lr * float(error.mean()) 1186 return Probe(w, b)
Fit a probe by plain gradient descent on binary cross-entropy.
The model under study is frozen: only the probe's d + 1 numbers change. The gradient of the mean cross-entropy with respect to the logit is simply (p − y), which is why the update is one line.
1189class PlantedModel: 1190 """A toy model whose hidden layer stores two known features as directions. 1191 1192 Each example has a sentiment (0 = negative, 1 = positive) and a tense 1193 (0 = present, 1 = past). The hidden state is 1194 1195 h = (±1) · sentiment_direction + (±1) · tense_direction + noise 1196 1197 in `d` neurons, with both directions random, so no single neuron is 1198 either feature. The model's output reads sentiment and, by construction, 1199 gives tense a weight of exactly zero: the tense information is in the 1200 hidden state, but nothing downstream uses it. 1201 """ 1202 1203 FEATURES = ("sentiment", "tense") 1204 1205 def __init__(self, d: int = 32, noise: float = 0.4, seed: int = 0): 1206 rng = np.random.default_rng(seed) 1207 A = rng.standard_normal((d, 2)) 1208 A /= np.linalg.norm(A, axis=0) # unit directions, not perpendicular to each other 1209 self.directions = dict(zip(self.FEATURES, A.T)) 1210 # The pseudo-inverse's first row is the readout with weight 1 on sentiment and 0 on tense. 1211 self.readout = np.linalg.pinv(A)[0] 1212 self.d, self.noise = d, noise 1213 1214 def hidden(self, labels: dict[str, np.ndarray], noise: np.ndarray) -> np.ndarray: 1215 """(n, d) hidden states for 0/1 labels, with the noise passed in so a flip changes only the label.""" 1216 signs = {name: 2 * np.asarray(labels[name]) - 1 for name in self.FEATURES} # 0/1 -> −1/+1 1217 return sum(np.outer(signs[name], self.directions[name]) for name in self.FEATURES) + noise 1218 1219 def output(self, H: np.ndarray) -> np.ndarray: 1220 """The model's own score for each example: positive means "positive review".""" 1221 return H @ self.readout 1222 1223 def sample(self, n: int, seed: int = 0) -> tuple[np.ndarray, dict[str, np.ndarray], np.ndarray]: 1224 rng = np.random.default_rng(seed) 1225 labels = {name: rng.integers(0, 2, n) for name in self.FEATURES} 1226 noise = rng.normal(0, self.noise, (n, self.d)) 1227 return self.hidden(labels, noise), labels, noise
A toy model whose hidden layer stores two known features as directions.
Each example has a sentiment (0 = negative, 1 = positive) and a tense (0 = present, 1 = past). The hidden state is
h = (±1) · sentiment_direction + (±1) · tense_direction + noise
in d neurons, with both directions random, so no single neuron is
either feature. The model's output reads sentiment and, by construction,
gives tense a weight of exactly zero: the tense information is in the
hidden state, but nothing downstream uses it.
1205 def __init__(self, d: int = 32, noise: float = 0.4, seed: int = 0): 1206 rng = np.random.default_rng(seed) 1207 A = rng.standard_normal((d, 2)) 1208 A /= np.linalg.norm(A, axis=0) # unit directions, not perpendicular to each other 1209 self.directions = dict(zip(self.FEATURES, A.T)) 1210 # The pseudo-inverse's first row is the readout with weight 1 on sentiment and 0 on tense. 1211 self.readout = np.linalg.pinv(A)[0] 1212 self.d, self.noise = d, noise
1219 def output(self, H: np.ndarray) -> np.ndarray: 1220 """The model's own score for each example: positive means "positive review".""" 1221 return H @ self.readout
The model's own score for each example: positive means "positive review".
1223 def sample(self, n: int, seed: int = 0) -> tuple[np.ndarray, dict[str, np.ndarray], np.ndarray]: 1224 rng = np.random.default_rng(seed) 1225 labels = {name: rng.integers(0, 2, n) for name in self.FEATURES} 1226 noise = rng.normal(0, self.noise, (n, self.d)) 1227 return self.hidden(labels, noise), labels, noise
1230def flip_effect(model: PlantedModel, feature: str, n: int = 200, seed: int = 0) -> float: 1231 """Average change in the model's output when only `feature` is flipped, noise held fixed. 1232 1233 This is an intervention, not a reading: it asks whether the model *uses* 1234 the feature, which no probe can answer. 1235 """ 1236 H, labels, noise = model.sample(n, seed) 1237 flipped = dict(labels, **{feature: 1 - labels[feature]}) 1238 return float(np.mean(np.abs(model.output(model.hidden(flipped, noise)) - model.output(H))))
Average change in the model's output when only feature is flipped, noise held fixed.
This is an intervention, not a reading: it asks whether the model uses the feature, which no probe can answer.
1241def probe_report(n_train: int = 100, n_test: int = 1000, seed: int = 0) -> dict[str, dict[str, float]]: 1242 """Train three probes on the same frozen hidden states and score each on its own training set and on unseen examples. 1243 1244 * sentiment: stored and used by the model 1245 * tense: stored but never used by the model's output 1246 * random labels: the control task, with nothing real to find 1247 """ 1248 model = PlantedModel() 1249 H_train, y_train, _ = model.sample(n_train, seed) 1250 H_test, y_test, _ = model.sample(n_test, seed + 1) 1251 rng = np.random.default_rng(seed + 2) 1252 y_train["random labels"] = rng.integers(0, 2, n_train) 1253 y_test["random labels"] = rng.integers(0, 2, n_test) 1254 report = {} 1255 for name in ("sentiment", "tense", "random labels"): 1256 probe = train_probe(H_train, y_train[name]) 1257 report[name] = {"train": probe.accuracy(H_train, y_train[name]), "test": probe.accuracy(H_test, y_test[name])} 1258 return report
Train three probes on the same frozen hidden states and score each on its own training set and on unseen examples.
- sentiment: stored and used by the model
- tense: stored but never used by the model's output
- random labels: the control task, with nothing real to find
1272def logit_lens(residuals: np.ndarray, W_U: np.ndarray, final_norm=None) -> np.ndarray: 1273 """Probabilities the model would give if it stopped here: softmax(W_U · norm(h)) for every state in `residuals`.""" 1274 x = final_norm(residuals) if final_norm is not None else residuals 1275 return softmax(x @ W_U.T, axis=-1)
Probabilities the model would give if it stopped here: softmax(W_U · norm(h)) for every state in residuals.
1278def worked_lens() -> list[dict]: 1279 """The worked example: the lens after the embedding, after layer 1 and after layer 2.""" 1280 states = [LENS_START] 1281 for write in LENS_WRITES: 1282 states.append(states[-1] + write) # a residual layer adds; it never overwrites 1283 probs = logit_lens(np.array(states), LENS_UNEMBED) 1284 return [ 1285 dict(layer=i, residual=s, logits=LENS_UNEMBED @ s, probs=p.tolist(), top=LENS_VOCAB[int(np.argmax(p))]) 1286 for i, (s, p) in enumerate(zip(states, probs)) 1287 ]
The worked example: the lens after the embedding, after layer 1 and after layer 2.
1290def residual_stream(model: TinyGPT, ids: np.ndarray) -> np.ndarray: 1291 """(n_layers + 1, seq, d): the residual state after the embedding and after every block.""" 1292 ids = np.asarray(ids) 1293 x = model.wte[ids] + model.wpe[: len(ids)] 1294 states = [x] 1295 for block in model.blocks: 1296 x = block(x) 1297 states.append(x) 1298 return np.stack(states)
(n_layers + 1, seq, d): the residual state after the embedding and after every block.
1301def tinygpt_logit_lens(model: TinyGPT, ids: np.ndarray) -> np.ndarray: 1302 """(n_layers + 1, seq, vocab) logits: every layer's residual state read through the final norm and the output layer. 1303 1304 Returns logits rather than probabilities so the top layer can be 1305 compared number for number with the model's own output. 1306 """ 1307 R = residual_stream(model, ids) 1308 return layer_norm(R, model.lnf_g, model.lnf_b) @ model.wte.T
(n_layers + 1, seq, vocab) logits: every layer's residual state read through the final norm and the output layer.
Returns logits rather than probabilities so the top layer can be compared number for number with the model's own output.
1316class LandmarkModel: 1317 """A two-layer, one-head transformer built by hand to answer "<landmark> is in" -> city. 1318 1319 Residual dimensions (a readable basis, so you can check it by eye; no 1320 tool in this lesson relies on it): 1321 1322 0-2 which landmark (Eiffel, Colosseum, Big Ben) 1323 3 "this token is a landmark" 1324 4-5 which filler word ("is", "in") 1325 6-8 which city (Paris, Rome, London) 1326 1327 Layer 1 is an MLP that looks up a fact at every position: a landmark's 1328 identity in, its city out. Layer 2 is an attention head: the word "in" 1329 queries for the landmark flag, finds the landmark and copies its city 1330 dimensions into the last position. The output layer reads the city 1331 dimensions at the last position. 1332 """ 1333 1334 TOKENS = ["Eiffel", "Colosseum", "Big Ben", "is", "in"] 1335 CITIES = ["Paris", "Rome", "London"] 1336 STAGES = ("resid_embed", "resid_mlp", "resid_attn") 1337 D = 9 1338 1339 def __init__(self, gain: float = 3.0, sharpness: float = 4.0): 1340 D = self.D 1341 self.E = np.zeros((len(self.TOKENS), D)) 1342 for i in range(3): 1343 self.E[i, i] = 1.0 # landmark identity 1344 self.E[i, 3] = 1.0 # landmark flag 1345 self.E[3, 4] = self.E[4, 5] = 1.0 # "is", "in" 1346 # MLP: three hidden units, one per landmark; each writes `gain` onto its city. 1347 self.W_in = np.zeros((D, 3)) 1348 self.W_in[:3, :3] = np.eye(3) 1349 self.W_out = np.zeros((3, D)) 1350 self.W_out[:3, 6:9] = gain * np.eye(3) 1351 # Attention, one head of width 1: "in" asks (query), the landmark flag answers (key). 1352 self.W_q = np.zeros((D, 1)) 1353 self.W_q[5, 0] = sharpness 1354 self.W_k = np.zeros((D, 1)) 1355 self.W_k[3, 0] = sharpness 1356 self.W_v = np.zeros((D, D)) 1357 self.W_v[6:9, 6:9] = np.eye(3) # the value is the city part of the state, copied as is 1358 self.W_U = np.zeros((3, D)) 1359 self.W_U[:, 6:9] = np.eye(3) 1360 1361 def run(self, tokens: list[str], patch: dict | None = None, answer: tuple[str, str] = ("Paris", "Rome")) -> dict: 1362 """Run the model, optionally overwriting activations. 1363 1364 `patch` maps (site, position) to a vector that replaces that activation 1365 during this run. Sites: "resid_embed", "mlp_out", "resid_mlp", 1366 "attn_out", "resid_attn". Returns every stage's residual state 1367 (3, seq, D), the component outputs, the last position's logits, and 1368 the logit difference between the two `answer` cities. 1369 """ 1370 patch = patch or {} 1371 1372 def apply(site: str, acts: np.ndarray) -> np.ndarray: 1373 acts = acts.copy() 1374 for (s, pos), vector in patch.items(): 1375 if s == site: 1376 acts[pos] = vector 1377 return acts 1378 1379 ids = [self.TOKENS.index(t) for t in tokens] 1380 x = apply("resid_embed", self.E[ids]) 1381 stages = [x] 1382 mlp_out = apply("mlp_out", np.maximum(x @ self.W_in, 0) @ self.W_out) 1383 x = apply("resid_mlp", x + mlp_out) 1384 stages.append(x) 1385 mixed, weights = scaled_dot_product_attention(x @ self.W_q, x @ self.W_k, x @ self.W_v, mask=causal_mask(len(ids))) 1386 attn_out = apply("attn_out", mixed) 1387 x = apply("resid_attn", x + attn_out) 1388 stages.append(x) 1389 logits = x[-1] @ self.W_U.T # only the last position predicts the next word 1390 good, bad = (self.CITIES.index(c) for c in answer) 1391 return dict( 1392 resid=np.stack(stages), mlp_out=mlp_out, attn_out=attn_out, attn_weights=weights, 1393 logits=logits, logit_diff=float(logits[good] - logits[bad]), 1394 )
A two-layer, one-head transformer built by hand to answer "
Residual dimensions (a readable basis, so you can check it by eye; no tool in this lesson relies on it):
0-2 which landmark (Eiffel, Colosseum, Big Ben)
3 "this token is a landmark"
4-5 which filler word ("is", "in")
6-8 which city (Paris, Rome, London)
Layer 1 is an MLP that looks up a fact at every position: a landmark's identity in, its city out. Layer 2 is an attention head: the word "in" queries for the landmark flag, finds the landmark and copies its city dimensions into the last position. The output layer reads the city dimensions at the last position.
1339 def __init__(self, gain: float = 3.0, sharpness: float = 4.0): 1340 D = self.D 1341 self.E = np.zeros((len(self.TOKENS), D)) 1342 for i in range(3): 1343 self.E[i, i] = 1.0 # landmark identity 1344 self.E[i, 3] = 1.0 # landmark flag 1345 self.E[3, 4] = self.E[4, 5] = 1.0 # "is", "in" 1346 # MLP: three hidden units, one per landmark; each writes `gain` onto its city. 1347 self.W_in = np.zeros((D, 3)) 1348 self.W_in[:3, :3] = np.eye(3) 1349 self.W_out = np.zeros((3, D)) 1350 self.W_out[:3, 6:9] = gain * np.eye(3) 1351 # Attention, one head of width 1: "in" asks (query), the landmark flag answers (key). 1352 self.W_q = np.zeros((D, 1)) 1353 self.W_q[5, 0] = sharpness 1354 self.W_k = np.zeros((D, 1)) 1355 self.W_k[3, 0] = sharpness 1356 self.W_v = np.zeros((D, D)) 1357 self.W_v[6:9, 6:9] = np.eye(3) # the value is the city part of the state, copied as is 1358 self.W_U = np.zeros((3, D)) 1359 self.W_U[:, 6:9] = np.eye(3)
1361 def run(self, tokens: list[str], patch: dict | None = None, answer: tuple[str, str] = ("Paris", "Rome")) -> dict: 1362 """Run the model, optionally overwriting activations. 1363 1364 `patch` maps (site, position) to a vector that replaces that activation 1365 during this run. Sites: "resid_embed", "mlp_out", "resid_mlp", 1366 "attn_out", "resid_attn". Returns every stage's residual state 1367 (3, seq, D), the component outputs, the last position's logits, and 1368 the logit difference between the two `answer` cities. 1369 """ 1370 patch = patch or {} 1371 1372 def apply(site: str, acts: np.ndarray) -> np.ndarray: 1373 acts = acts.copy() 1374 for (s, pos), vector in patch.items(): 1375 if s == site: 1376 acts[pos] = vector 1377 return acts 1378 1379 ids = [self.TOKENS.index(t) for t in tokens] 1380 x = apply("resid_embed", self.E[ids]) 1381 stages = [x] 1382 mlp_out = apply("mlp_out", np.maximum(x @ self.W_in, 0) @ self.W_out) 1383 x = apply("resid_mlp", x + mlp_out) 1384 stages.append(x) 1385 mixed, weights = scaled_dot_product_attention(x @ self.W_q, x @ self.W_k, x @ self.W_v, mask=causal_mask(len(ids))) 1386 attn_out = apply("attn_out", mixed) 1387 x = apply("resid_attn", x + attn_out) 1388 stages.append(x) 1389 logits = x[-1] @ self.W_U.T # only the last position predicts the next word 1390 good, bad = (self.CITIES.index(c) for c in answer) 1391 return dict( 1392 resid=np.stack(stages), mlp_out=mlp_out, attn_out=attn_out, attn_weights=weights, 1393 logits=logits, logit_diff=float(logits[good] - logits[bad]), 1394 )
Run the model, optionally overwriting activations.
patch maps (site, position) to a vector that replaces that activation
during this run. Sites: "resid_embed", "mlp_out", "resid_mlp",
"attn_out", "resid_attn". Returns every stage's residual state
(3, seq, D), the component outputs, the last position's logits, and
the logit difference between the two answer cities.
1397def fraction_restored(patched: float, clean: float, corrupt: float) -> float: 1398 """How much of the clean-vs-corrupt gap one patch closes: 0 = nothing, 1 = everything.""" 1399 return (patched - corrupt) / (clean - corrupt)
How much of the clean-vs-corrupt gap one patch closes: 0 = nothing, 1 = everything.
1402def patching_map(model: LandmarkModel, clean: list[str], corrupt: list[str]) -> np.ndarray: 1403 """(stage, position) grid: patch each residual state from the clean run into the corrupted run, one at a time.""" 1404 clean_run, corrupt_run = model.run(clean), model.run(corrupt) 1405 grid = np.zeros((len(model.STAGES), len(clean))) 1406 for s, stage in enumerate(model.STAGES): 1407 for pos in range(len(clean)): 1408 patched = model.run(corrupt, patch={(stage, pos): clean_run["resid"][s, pos]}) 1409 grid[s, pos] = fraction_restored(patched["logit_diff"], clean_run["logit_diff"], corrupt_run["logit_diff"]) 1410 return grid
(stage, position) grid: patch each residual state from the clean run into the corrupted run, one at a time.
1413def lens_map(model: LandmarkModel, tokens: list[str], answer: str = "Paris") -> np.ndarray: 1414 """(stage, position) grid of the logit lens's probability for `answer`.""" 1415 resid = model.run(tokens)["resid"] 1416 return logit_lens(resid, model.W_U)[..., model.CITIES.index(answer)]
(stage, position) grid of the logit lens's probability for answer.
1424def pentagon() -> np.ndarray: 1425 """(2, 5): five unit feature directions spread 72° apart in a 2-neuron space.""" 1426 angles = np.radians(72 * np.arange(5)) 1427 return np.stack([np.cos(angles), np.sin(angles)])
(2, 5): five unit feature directions spread 72° apart in a 2-neuron space.
1430def superposition_readout(W: np.ndarray, b: np.ndarray, x: np.ndarray, relu: bool = True) -> np.ndarray: 1431 """Squeeze features x into h = W x, then read them back out: ReLU(Wᵀ h + b).""" 1432 z = W.T @ (W @ x) + b 1433 return np.maximum(z, 0) if relu else z
Squeeze features x into h = W x, then read them back out: ReLU(Wᵀ h + b).
1436def sample_sparse_features(n: int, n_features: int, p_active: float, rng: np.random.Generator) -> np.ndarray: 1437 """(n, n_features): each feature is on with probability p_active, at a strength uniform in [0, 1).""" 1438 on = rng.random((n, n_features)) < p_active 1439 return on * rng.random((n, n_features))
(n, n_features): each feature is on with probability p_active, at a strength uniform in [0, 1).
1442def superposition_loss_and_grads(W: np.ndarray, b: np.ndarray, X: np.ndarray, importance: np.ndarray | None = None): 1443 """Importance-weighted reconstruction loss of the toy model, and its gradients, derived by hand. 1444 1445 Forward: H = X Wᵀ (batch, d), Z = H W + b (batch, n), Y = ReLU(Z) 1446 Loss: mean over the batch of Σᵢ Iᵢ (Yᵢ − Xᵢ)² 1447 W appears twice (squeeze and unsqueeze), so its gradient has two terms. 1448 """ 1449 importance = np.ones(X.shape[1]) if importance is None else importance 1450 H = X @ W.T 1451 Z = H @ W + b 1452 Y = np.maximum(Z, 0) 1453 loss = float(np.mean(np.sum(importance * (Y - X) ** 2, axis=1))) 1454 dZ = 2 * importance * (Y - X) / len(X) * (Z > 0) # ReLU passes gradient only where it was on 1455 dH = dZ @ W.T 1456 grads = {"W": H.T @ dZ + dH.T @ X, "b": dZ.sum(axis=0)} 1457 return loss, grads
Importance-weighted reconstruction loss of the toy model, and its gradients, derived by hand.
Forward: H = X Wᵀ (batch, d), Z = H W + b (batch, n), Y = ReLU(Z) Loss: mean over the batch of Σᵢ Iᵢ (Yᵢ − Xᵢ)² W appears twice (squeeze and unsqueeze), so its gradient has two terms.
1460def train_superposition( 1461 p_active: float, n_features: int = 5, d: int = 2, importance_decay: float = 0.8, 1462 steps: int = 2000, batch: int = 1024, lr: float = 0.01, seed: int = 0, 1463) -> dict: 1464 """Train the toy model of Elhage et al. (2022) on features that are on with probability `p_active`. 1465 1466 Feature i matters `importance_decay ** i` as much as feature 0, so when 1467 the model cannot keep everything it has a reason to choose. 1468 """ 1469 rng = np.random.default_rng(seed) 1470 importance = importance_decay ** np.arange(n_features) 1471 params = {"W": rng.normal(0, 0.5, (d, n_features)), "b": np.zeros(n_features)} 1472 opt = {k: Adam(lr=lr) for k in params} 1473 losses = [] 1474 for _ in range(steps): 1475 X = sample_sparse_features(batch, n_features, p_active, rng) 1476 loss, grads = superposition_loss_and_grads(params["W"], params["b"], X, importance) 1477 for k in params: 1478 params[k] = opt[k].step(params[k], grads[k]) 1479 losses.append(loss) 1480 return dict(params, losses=np.array(losses))
Train the toy model of Elhage et al. (2022) on features that are on with probability p_active.
Feature i matters importance_decay ** i as much as feature 0, so when
the model cannot keep everything it has a reason to choose.
1483def features_a_neuron_responds_to(W: np.ndarray, neuron: int, threshold: float = 0.25) -> list[int]: 1484 """Indices of the features that push this neuron up by more than `threshold` per unit of feature.""" 1485 return [k for k in range(W.shape[1]) if W[neuron, k] > threshold]
Indices of the features that push this neuron up by more than threshold per unit of feature.
1488def finite_difference_gradient(f, w: np.ndarray, eps: float = 1e-5) -> np.ndarray: 1489 """Slow, obviously correct gradient: nudge each entry up and down and watch the loss.""" 1490 grad = np.zeros_like(w) 1491 for idx in np.ndindex(w.shape): 1492 old = w[idx] 1493 w[idx] = old + eps 1494 up = f(w) 1495 w[idx] = old - eps 1496 down = f(w) 1497 w[idx] = old 1498 grad[idx] = (up - down) / (2 * eps) 1499 return grad
Slow, obviously correct gradient: nudge each entry up and down and watch the loss.
1507def sae_objective(h: np.ndarray, f: np.ndarray, decoder: np.ndarray, lam: float, b_d: np.ndarray | float = 0.0) -> float: 1508 """The sparse autoencoder's loss for one example and one proposed code f: ‖h − (D f + b_d)‖² + λ Σ|fᵢ|.""" 1509 residual = h - (decoder @ f + b_d) 1510 return float(residual @ residual + lam * np.abs(f).sum())
The sparse autoencoder's loss for one example and one proposed code f: ‖h − (D f + b_d)‖² + λ Σ|fᵢ|.
1513class SparseAutoencoder: 1514 """encode: f = ReLU(W_e (h − b_d) + b_e); decode: ĥ = W_d f + b_d. 1515 1516 `n_latents` is usually much bigger than `d` (overcomplete). Each decoder 1517 column is kept at length 1, so the L1 penalty cannot be dodged by making 1518 codes tiny and decoder columns huge. 1519 """ 1520 1521 def __init__(self, d: int, n_latents: int, seed: int = 0): 1522 rng = np.random.default_rng(seed) 1523 W_d = rng.standard_normal((d, n_latents)) 1524 self.W_d = W_d / np.linalg.norm(W_d, axis=0) # (d, m): one unit direction per latent 1525 self.W_e = self.W_d.T.copy() # (m, d): start each latent listening for its own direction 1526 self.b_e = np.zeros(n_latents) 1527 self.b_d = np.zeros(d) 1528 1529 def encode(self, H: np.ndarray) -> np.ndarray: 1530 return np.maximum((H - self.b_d) @ self.W_e.T + self.b_e, 0) 1531 1532 def decode(self, F: np.ndarray) -> np.ndarray: 1533 return F @ self.W_d.T + self.b_d 1534 1535 def reconstruction_error(self, H: np.ndarray) -> float: 1536 """Fraction of the data's variance the autoencoder fails to rebuild (0 = perfect).""" 1537 residual = self.decode(self.encode(H)) - H 1538 return float(np.sum(residual**2) / np.sum((H - H.mean(axis=0)) ** 2)) 1539 1540 def loss_and_grads(self, H: np.ndarray, lam: float) -> tuple[float, dict[str, np.ndarray]]: 1541 """Mean over the batch of ‖h − ĥ‖² + λ Σ fᵢ, and the gradient for every parameter, by hand.""" 1542 n = len(H) 1543 centred = H - self.b_d 1544 pre = centred @ self.W_e.T + self.b_e # (n, m) 1545 F = np.maximum(pre, 0) 1546 residual = F @ self.W_d.T + self.b_d - H # (n, d) 1547 loss = float(np.sum(residual**2) / n + lam * F.sum() / n) # F ≥ 0, so |F| = F 1548 d_hat = 2 * residual / n 1549 dF = d_hat @ self.W_d + lam / n 1550 d_pre = dF * (pre > 0) 1551 grads = { 1552 "W_d": d_hat.T @ F, 1553 "W_e": d_pre.T @ centred, 1554 "b_e": d_pre.sum(axis=0), 1555 # b_d is added back in the decoder and subtracted in the encoder. 1556 "b_d": d_hat.sum(axis=0) - (d_pre @ self.W_e).sum(axis=0), 1557 } 1558 return loss, grads
encode: f = ReLU(W_e (h − b_d) + b_e); decode: ĥ = W_d f + b_d.
n_latents is usually much bigger than d (overcomplete). Each decoder
column is kept at length 1, so the L1 penalty cannot be dodged by making
codes tiny and decoder columns huge.
1521 def __init__(self, d: int, n_latents: int, seed: int = 0): 1522 rng = np.random.default_rng(seed) 1523 W_d = rng.standard_normal((d, n_latents)) 1524 self.W_d = W_d / np.linalg.norm(W_d, axis=0) # (d, m): one unit direction per latent 1525 self.W_e = self.W_d.T.copy() # (m, d): start each latent listening for its own direction 1526 self.b_e = np.zeros(n_latents) 1527 self.b_d = np.zeros(d)
1535 def reconstruction_error(self, H: np.ndarray) -> float: 1536 """Fraction of the data's variance the autoencoder fails to rebuild (0 = perfect).""" 1537 residual = self.decode(self.encode(H)) - H 1538 return float(np.sum(residual**2) / np.sum((H - H.mean(axis=0)) ** 2))
Fraction of the data's variance the autoencoder fails to rebuild (0 = perfect).
1540 def loss_and_grads(self, H: np.ndarray, lam: float) -> tuple[float, dict[str, np.ndarray]]: 1541 """Mean over the batch of ‖h − ĥ‖² + λ Σ fᵢ, and the gradient for every parameter, by hand.""" 1542 n = len(H) 1543 centred = H - self.b_d 1544 pre = centred @ self.W_e.T + self.b_e # (n, m) 1545 F = np.maximum(pre, 0) 1546 residual = F @ self.W_d.T + self.b_d - H # (n, d) 1547 loss = float(np.sum(residual**2) / n + lam * F.sum() / n) # F ≥ 0, so |F| = F 1548 d_hat = 2 * residual / n 1549 dF = d_hat @ self.W_d + lam / n 1550 d_pre = dF * (pre > 0) 1551 grads = { 1552 "W_d": d_hat.T @ F, 1553 "W_e": d_pre.T @ centred, 1554 "b_e": d_pre.sum(axis=0), 1555 # b_d is added back in the decoder and subtracted in the encoder. 1556 "b_d": d_hat.sum(axis=0) - (d_pre @ self.W_e).sum(axis=0), 1557 } 1558 return loss, grads
Mean over the batch of ‖h − ĥ‖² + λ Σ fᵢ, and the gradient for every parameter, by hand.
1561def train_sae( 1562 H: np.ndarray, n_latents: int = 5, lam: float = 0.3, steps: int = 2000, batch: int = 512, lr: float = 0.01, seed: int = 0 1563) -> SparseAutoencoder: 1564 """Fit a sparse autoencoder to activations H (n, d) with Adam, renormalizing decoder columns after every step.""" 1565 rng = np.random.default_rng(seed) 1566 sae = SparseAutoencoder(H.shape[1], n_latents, seed) 1567 names = ("W_e", "b_e", "W_d", "b_d") 1568 opt = {k: Adam(lr=lr) for k in names} 1569 for _ in range(steps): 1570 _, grads = sae.loss_and_grads(H[rng.integers(0, len(H), batch)], lam) 1571 for k in names: 1572 setattr(sae, k, opt[k].step(getattr(sae, k), grads[k])) 1573 sae.W_d /= np.linalg.norm(sae.W_d, axis=0) 1574 return sae
Fit a sparse autoencoder to activations H (n, d) with Adam, renormalizing decoder columns after every step.
1577def match_features(true_directions: np.ndarray, learned_directions: np.ndarray) -> np.ndarray: 1578 """For each true feature (a column), the cosine similarity of the learned column that points most nearly the same way.""" 1579 t = true_directions / np.linalg.norm(true_directions, axis=0) 1580 ell = learned_directions / np.linalg.norm(learned_directions, axis=0) 1581 return (t.T @ ell).max(axis=1)
For each true feature (a column), the cosine similarity of the learned column that points most nearly the same way.
1593def figures() -> dict: 1594 """Plot this lesson's data. matplotlib is imported here, and only here, 1595 so the lesson itself needs nothing beyond NumPy.""" 1596 import matplotlib 1597 1598 matplotlib.use("Agg") 1599 import matplotlib.pyplot as plt 1600 1601 BLUE, RED, GREEN, MUTED = "#2563eb", "#dc2626", "#059669", "#9ca3af" 1602 figs = {} 1603 1604 # --- 1. Features as directions ----------------------------------------- 1605 fig, ax = plt.subplots(figsize=(4.6, 4.2)) 1606 amounts = {"positive": 0.9, "past tense": 0.4} 1607 h = compose_features(amounts, WORKED_DIRECTIONS) 1608 for (name, d), color in zip(WORKED_DIRECTIONS.items(), (BLUE, GREEN)): 1609 ax.annotate("", d, (0, 0), arrowprops=dict(arrowstyle="->", color=color, lw=2)) 1610 # Labels sit beyond the arrow tip, above it for the upper arrow and below it for the lower one. 1611 ax.text(d[0] * 1.1, d[1] * 1.1, f'"{name}"\ndirection', color=color, ha="center", va="bottom" if d[1] > 0 else "top", fontsize=9) 1612 foot = amounts[name] * d 1613 ax.plot([h[0], foot[0]], [h[1], foot[1]], ls=":", color=color) 1614 ax.plot(*foot, "o", color=color, ms=4) 1615 ax.annotate("", h, (0, 0), arrowprops=dict(arrowstyle="->", color=RED, lw=2.5)) 1616 ax.text(h[0] + 0.04, h[1] + 0.03, "h = (0.86, 0.48)", color=RED) 1617 ax.set_xlim(-0.2, 1.15) 1618 ax.set_ylim(-0.9, 1.1) 1619 ax.set_aspect("equal") 1620 ax.axhline(0, color=MUTED, lw=0.8) 1621 ax.axvline(0, color=MUTED, lw=0.8) 1622 ax.set_xlabel("neuron 1") 1623 ax.set_ylabel("neuron 2") 1624 ax.set_title("Two features, two directions, one hidden state") 1625 figs["directions"] = fig 1626 1627 # --- 2. Probes: what they can read, and what the model uses ------------ 1628 report = probe_report() 1629 model = PlantedModel() 1630 names = ["sentiment", "tense", "random labels"] 1631 fig, (a1, a2) = plt.subplots(1, 2, figsize=(9, 3.4)) 1632 x = np.arange(len(names)) 1633 a1.bar(x - 0.2, [report[n]["train"] for n in names], 0.4, color=MUTED, label="training examples") 1634 a1.bar(x + 0.2, [report[n]["test"] for n in names], 0.4, color=BLUE, label="unseen examples") 1635 for xi, n in zip(x, names): 1636 a1.text(xi + 0.2, report[n]["test"] + 0.02, f"{report[n]['test']:.2f}", ha="center") 1637 a1.axhline(0.5, color=RED, ls="--", lw=1, label="chance") 1638 a1.set_xticks(x, names) 1639 a1.set_ylim(0, 1.3) 1640 a1.set_ylabel("probe accuracy") 1641 a1.set_title("What a probe can read") 1642 a1.legend(frameon=False, loc="upper right", fontsize=8) 1643 effects = [flip_effect(model, f) for f in ("sentiment", "tense")] 1644 a2.bar([0, 1], effects, 0.5, color=[BLUE, MUTED]) 1645 for xi, e in enumerate(effects): 1646 a2.text(xi, e + 0.05, f"{e:.2f}", ha="center") 1647 a2.set_xticks([0, 1], ["flip sentiment", "flip tense"]) 1648 a2.set_ylim(0, 2.5) 1649 a2.set_ylabel("change in the model's output") 1650 a2.set_title("What the model actually uses") 1651 fig.tight_layout() 1652 figs["probe_vs_use"] = fig 1653 1654 # --- 3. Logit lens: worked example and the landmark model ------------- 1655 fig, (a1, a2) = plt.subplots(1, 2, figsize=(9.5, 3.6), gridspec_kw={"width_ratios": [1, 1.2]}) 1656 lens = worked_lens() 1657 layers = np.arange(len(lens)) 1658 for k, (word, color) in enumerate(zip(LENS_VOCAB, (BLUE, GREEN, MUTED))): 1659 a1.bar(layers + (k - 1) * 0.27, [row["probs"][k] for row in lens], 0.27, color=color, label=word) 1660 a1.set_xticks(layers, ["embedding", "layer 1", "layer 2"]) 1661 a1.set_ylabel("lens probability") 1662 a1.set_ylim(0, 1) 1663 a1.set_title("Worked example: the guess firms up") 1664 a1.legend(frameon=False, fontsize=8) 1665 grid = lens_map(LandmarkModel(), CLEAN) 1666 im = a2.imshow(grid, cmap="Blues", vmin=0, vmax=1) 1667 for (r, c), v in np.ndenumerate(grid): 1668 a2.text(c, r, f"{v:.2f}", ha="center", va="center", color="white" if v > 0.6 else "black") 1669 a2.set_xticks(range(3), CLEAN) 1670 a2.set_yticks(range(3), STAGE_LABELS) 1671 a2.set_title('Landmark model: P("Paris") by layer and word') 1672 a2.grid(False) 1673 fig.colorbar(im, ax=a2, fraction=0.046) 1674 fig.tight_layout() 1675 figs["logit_lens"] = fig 1676 1677 # --- 4. Activation patching map --------------------------------------- 1678 grid = patching_map(LandmarkModel(), CLEAN, CORRUPT) 1679 fig, ax = plt.subplots(figsize=(5.6, 3.4)) 1680 im = ax.imshow(grid, cmap="Greens", vmin=0, vmax=1) 1681 for (r, c), v in np.ndenumerate(grid): 1682 ax.text(c, r, f"{v:.2f}", ha="center", va="center", color="white" if v > 0.6 else "black") 1683 ax.set_xticks(range(3), [f"{a} / {b}" if a != b else a for a, b in zip(CLEAN, CORRUPT)]) 1684 ax.set_yticks(range(3), STAGE_LABELS) 1685 ax.set_xlabel("position patched (clean / corrupted word)") 1686 ax.set_title("Fraction of the answer restored by one patch") 1687 ax.grid(False) 1688 fig.colorbar(im, ax=ax, fraction=0.046) 1689 figs["patching"] = fig 1690 1691 # --- 5. Superposition: dense keeps two, sparse keeps five ------------- 1692 fig, axes = plt.subplots(1, 3, figsize=(11, 3.7), gridspec_kw={"width_ratios": [1, 1, 1.2]}) 1693 colors = [BLUE, GREEN, "#d97706", "#7c3aed", RED] 1694 for ax, p, title in ((axes[0], 1.0, "Dense features: keeps 2"), (axes[1], 0.05, "Sparse features: keeps all 5")): 1695 W = train_superposition(p_active=p)["W"] 1696 for k in range(W.shape[1]): 1697 if np.linalg.norm(W[:, k]) > 0.1: 1698 ax.annotate("", W[:, k], (0, 0), arrowprops=dict(arrowstyle="->", color=colors[k], lw=2)) 1699 ax.text(*(W[:, k] * 1.18), f"f{k}", color=colors[k], ha="center", va="center") 1700 else: # a dropped feature: an arrow this short would be all arrowhead 1701 ax.plot(*W[:, k], "o", color=colors[k], ms=4) 1702 ax.set_xlim(-1.45, 1.45) 1703 ax.set_ylim(-1.45, 1.45) 1704 ax.set_aspect("equal") 1705 ax.set_title(f"{title}\n(each feature on {p:.0%} of the time)") 1706 ax.set_xlabel("neuron 1") 1707 ax.set_ylabel("neuron 2") 1708 W = pentagon() 1709 x = np.arange(5) 1710 axes[2].bar(x - 0.2, W[0], 0.4, color=BLUE, label="neuron 1") 1711 axes[2].bar(x + 0.2, W[1], 0.4, color=MUTED, label="neuron 2") 1712 axes[2].axhline(0, color="#4b5563", lw=0.8) 1713 axes[2].set_xticks(x, [f"f{k}" for k in x]) 1714 axes[2].set_ylabel("how much the feature moves the neuron") 1715 axes[2].set_title("Each neuron answers to several features") 1716 axes[2].legend(frameon=False, fontsize=8) 1717 fig.tight_layout() 1718 figs["superposition"] = fig 1719 1720 # --- 6. Sparse autoencoder: recovering the planted features ------------ 1721 rng = np.random.default_rng(1) 1722 H = sample_sparse_features(20_000, 5, 0.05, rng) @ pentagon().T 1723 fig, (a1, a2) = plt.subplots(1, 2, figsize=(9.5, 3.9)) 1724 shown = H[np.abs(H).sum(axis=1) > 0][:600] 1725 a1.scatter(shown[:, 0], shown[:, 1], s=3, color=MUTED, alpha=0.5, label="hidden states") 1726 sae = train_sae(H, n_latents=5, lam=0.3) 1727 for k in range(5): 1728 t = pentagon()[:, k] 1729 a1.plot([0, 1.25 * t[0]], [0, 1.25 * t[1]], color=MUTED, lw=6, alpha=0.4, label="planted features" if k == 0 else None) 1730 d = sae.W_d[:, k] 1731 a1.annotate("", 1.1 * d, (0, 0), arrowprops=dict(arrowstyle="->", color=RED, lw=1.8)) 1732 a1.plot([], [], color=RED, label="learned latents") 1733 a1.set_aspect("equal") 1734 a1.set_xlim(-1.4, 1.4) 1735 a1.set_ylim(-1.4, 1.4) 1736 a1.set_title("λ = 0.3: each latent lands on a feature") 1737 a1.legend(frameon=False, fontsize=8, loc="lower left") 1738 lams = [0.003, 0.01, 0.03, 0.1, 0.2, 0.3, 0.6, 1.0] 1739 worst, error = [], [] 1740 for lam in lams: 1741 s = train_sae(H, n_latents=5, lam=lam) 1742 worst.append(match_features(pentagon(), s.W_d).min()) 1743 error.append(s.reconstruction_error(H)) 1744 a2.plot(lams, worst, "o-", color=BLUE, label="worst feature match (cosine)") 1745 a2.plot(lams, error, "s-", color=RED, label="variance left unexplained") 1746 a2.set_xscale("log") 1747 a2.set_xlabel("sparsity penalty λ") 1748 a2.set_ylim(0, 1.05) 1749 a2.set_title("The penalty trades rebuild for meaning") 1750 a2.legend(frameon=False, fontsize=8) 1751 fig.tight_layout() 1752 figs["sae"] = fig 1753 1754 return figs
Plot this lesson's data. matplotlib is imported here, and only here, so the lesson itself needs nothing beyond NumPy.
1762def demo() -> None: 1763 banner("1. A feature is a direction") 1764 h = compose_features({"positive": 0.9, "past tense": 0.4}, WORKED_DIRECTIONS) 1765 say( 1766 f""" 1767 Mix 0.9 of the "positive" direction (0.6, 0.8) with 0.4 of the "past 1768 tense" direction (0.8, -0.6) and the hidden state is h = ({h[0]:.2f}, 1769 {h[1]:.2f}). Neither neuron equals either feature. A dot product with 1770 each direction reads the amounts back: 1771 """ 1772 ) 1773 table(["feature", "direction · h"], [(k, v) for k, v in read_features(h, WORKED_DIRECTIONS).items()], floatfmt=".2f") 1774 takeaway("Features live in directions, not in individual neurons.") 1775 1776 banner("2. Probes: reading a property out of the hidden state") 1777 say( 1778 f""" 1779 A probe is a tiny classifier trained on frozen hidden states. The 1780 worked probe w = (1, 0.5), b = 0 reads h = (2, -1) as 1781 σ(1.5) = {probe_probability([2.0, -1.0], [1.0, 0.5], 0.0):.3f}. Now three 1782 probes on a toy model with 32 neurons, trained on 100 examples and 1783 scored on 1,000 unseen ones: 1784 """ 1785 ) 1786 report = probe_report() 1787 table(["probe target", "train accuracy", "unseen accuracy"], [(k, v["train"], v["test"]) for k, v in report.items()], floatfmt=".3f") 1788 model = PlantedModel() 1789 say( 1790 f""" 1791 Random labels fit the training set to {report['random labels']['train']:.0%} and then 1792 score a coin flip on unseen examples: accuracy only counts on held-out 1793 data. Tense is read at {report['tense']['test']:.0%}, yet flipping it moves the 1794 model's output by {flip_effect(model, 'tense'):.2f}, while flipping sentiment moves 1795 it by {flip_effect(model, 'sentiment'):.2f}. 1796 """ 1797 ) 1798 takeaway("A probe shows information is present, not that the model uses it.") 1799 1800 banner("3. The logit lens: what would the model say if it stopped here?") 1801 table( 1802 ["after", "residual", "P(Paris)", "P(London)", "P(Rome)", "top"], 1803 [(["embedding", "layer 1", "layer 2"][r["layer"]], str(np.round(r["residual"], 2)), *r["probs"], r["top"]) for r in worked_lens()], 1804 floatfmt=".2f", 1805 ) 1806 tiny = TinyGPT(vocab_size=11, d_model=16, n_layers=3, n_heads=2, max_len=8, seed=4) 1807 ids = np.array([3, 1, 4, 1, 5]) 1808 agree = np.allclose(tinygpt_logit_lens(tiny, ids)[-1], tiny(ids)) 1809 say(f"On primer.ml.transformer.TinyGPT, the lens at the top layer equals the model's own logits: {agree}.") 1810 matrix('landmark model, lens P("Paris"): rows = layers, cols = Eiffel, is, in', lens_map(LandmarkModel(), CLEAN)) 1811 takeaway("The Eiffel position knows Paris after layer 1; the last word only learns it in layer 2.") 1812 1813 banner("4. Activation patching: which activations carry the answer?") 1814 lm = LandmarkModel() 1815 clean, corrupt = lm.run(CLEAN)["logit_diff"], lm.run(CORRUPT)["logit_diff"] 1816 say( 1817 f""" 1818 Logit difference Paris minus Rome: clean "Eiffel is in" gives {clean:+.2f}, 1819 corrupted "Colosseum is in" gives {corrupt:+.2f}. Copy one clean 1820 residual state into the corrupted run at a time and measure the 1821 fraction of the gap that comes back: 1822 """ 1823 ) 1824 matrix("fraction restored: rows = layers, cols = positions", patching_map(lm, CLEAN, CORRUPT), precision=2) 1825 takeaway("The answer sits at the subject early and at the last word late: a two-step circuit, found by intervention.") 1826 1827 banner("5. Superposition: five features in two neurons") 1828 table( 1829 ["feature", "readout when only f0 is on (no ReLU)", "with bias -0.31 and ReLU"], 1830 [ 1831 (f"f{k}", a, c) 1832 for k, (a, c) in enumerate( 1833 zip( 1834 superposition_readout(pentagon(), np.zeros(5), np.eye(5)[0], relu=False), 1835 superposition_readout(pentagon(), np.full(5, -0.31), np.eye(5)[0]), 1836 ) 1837 ) 1838 ], 1839 floatfmt=".3f", 1840 ) 1841 for p in (1.0, 0.05): 1842 W = train_superposition(p_active=p)["W"] 1843 print(f"trained with each feature on {p:>4.0%} of the time: feature lengths {np.round(np.linalg.norm(W, axis=0), 2)}") 1844 print() 1845 say( 1846 """ 1847 Dense features: the model keeps the two most important and drops the 1848 rest. Sparse features: it keeps all five, 72 degrees apart, and a 1849 negative bias filters the small leaks. In the ideal pentagon, neuron 1 answers to 1850 features 0, 1 and 4: it is polysemantic. 1851 """ 1852 ) 1853 1854 banner("6. Sparse autoencoders: getting the features back") 1855 rng = np.random.default_rng(1) 1856 H = sample_sparse_features(20_000, 5, 0.05, rng) @ pentagon().T 1857 rows = [] 1858 for lam in (0.01, 0.3): 1859 sae = train_sae(H, n_latents=5, lam=lam) 1860 rows.append((lam, sae.reconstruction_error(H), match_features(pentagon(), sae.W_d).min())) 1861 table(["λ", "variance unexplained", "worst feature match (cosine)"], rows, floatfmt=".3f") 1862 takeaway( 1863 "Any two directions can rebuild a 2-D space; only the sparsity penalty makes the " 1864 "autoencoder find the five directions the model really used." 1865 )