Looking Inside a Trained Model
Activation patching
Activation patching swaps one internal number from a correct run into a broken run, to test which layer is actually causing a prediction rather than only correlated with it.
- 9 min read
- 3 reading levels
- Published
Read these first
On this page 5
One lesson, three depths. Pick the one that fits you today — you can switch any time.
Beginner — No maths. Plain English.
Activation patching swaps one piece of a model's internal working from a correct run into a broken run. This finds which piece actually mattered.
Imagine a dish that tastes wrong. You suspect one ingredient. So you swap only that ingredient back to the original, keep everything else the same, and taste it again.
If the dish tastes right again, you found the ingredient that mattered. Activation patching does the same test on a model, swapping one internal number at a time instead of an ingredient.
Why it exists
The logit lens shows what a model guesses at each layer. It does not prove that layer caused the guess. Two things can look connected without one actually causing the other.
To prove a layer matters, you need to change only that layer and watch what happens to the output. If nothing changes, that layer was not responsible, no matter how promising it looked.
This is the same logic doctors use to test a drug. Change one thing, hold everything else fixed, and measure the difference.
How it works
Correct prompt: "The capital of France is" -> "Paris"
Broken prompt: "The capital of Germany is" -> "Berlin"
Step 1: run the CORRECT prompt, save its internal numbers.
Step 2: run the BROKEN prompt, but swap in the correct
prompt's saved numbers at one specific layer.
Step 3: check if the output flips back toward "Paris".
If it does, that layer was carrying the "France" information.
If it does not, that layer was not the one responsible.Repeat this one layer at a time. It builds a map of where a specific piece of information lives in the model.
Where you have already seen it
- AI safety audits. Testing which part of a model is responsible for a harmful or biased output, before trying to fix it.
- Model debugging. Tracing exactly where a model's "reasoning" went off track, for one wrong answer.
- Research into how facts are stored. Locating which layer holds a specific fact, like a country's capital, inside a trained model.
Remember this
- Activation patching swaps one internal value between two runs, to test if it caused the output.
- It answers a causal question the logit lens cannot: not "what does this layer guess", but "does this layer matter".
- Changing one thing at a time, and measuring the effect, is the same logic behind a controlled experiment.
What to learn next
- The logit lens — the observational technique this method turns into a causal test.
- Where a model stores a fact — a direct application of this exact method.
- Induction heads — a circuit originally confirmed using this same patching approach.
Developer — Code and libraries.
This patches one prompt's fact into another, layer by layer, and watches how much of the correct answer's probability comes back.
Setup
pip install transformers torchThe first run downloads gpt2, roughly 500 MB.
Patching "France" into a "Germany" prompt, one layer at a time
import torch
from transformers import GPT2LMHeadModel, GPT2Tokenizer
tok = GPT2Tokenizer.from_pretrained("gpt2")
model = GPT2LMHeadModel.from_pretrained("gpt2")
model.eval()
clean = tok("The capital of France is", return_tensors="pt")
corrupted = tok("The capital of Germany is", return_tensors="pt")
paris_id = tok(" Paris")["input_ids"][0]
COUNTRY_POS = 3 # the token position holding "France" / "Germany"
# Step 1: run the clean prompt, save every block's output.
clean_cache = {}
def save_hook(layer_idx):
def hook(module, inputs, output):
clean_cache[layer_idx] = output.detach().clone()
return hook
handles = [model.transformer.h[i].register_forward_hook(save_hook(i)) for i in range(12)]
with torch.no_grad():
model(**clean)
for h in handles: h.remove()
with torch.no_grad():
baseline = torch.softmax(model(**corrupted).logits[0, -1], dim=-1)[paris_id].item()
print(f"corrupted prompt alone: P(Paris) = {baseline:.4f}")
# Step 2: patch one block at a time, only at the country-name position.
def patch_hook(layer_idx):
def hook(module, inputs, output):
patched = output.clone()
patched[:, COUNTRY_POS, :] = clean_cache[layer_idx][:, COUNTRY_POS, :]
return patched
return hook
for layer_idx in range(12):
handle = model.transformer.h[layer_idx].register_forward_hook(patch_hook(layer_idx))
with torch.no_grad():
prob = torch.softmax(model(**corrupted).logits[0, -1], dim=-1)[paris_id].item()
handle.remove()
print(f"patch block {layer_idx:2}: P(Paris) = {prob:.4f}")corrupted prompt alone: P(Paris) = 0.0013 patch block 0: P(Paris) = 0.0318 patch block 1: P(Paris) = 0.0324 patch block 2: P(Paris) = 0.0308 patch block 3: P(Paris) = 0.0322 patch block 4: P(Paris) = 0.0344 patch block 5: P(Paris) = 0.0299 patch block 6: P(Paris) = 0.0348 patch block 7: P(Paris) = 0.0341 patch block 8: P(Paris) = 0.0337 patch block 9: P(Paris) = 0.0067 patch block 10: P(Paris) = 0.0027 patch block 11: P(Paris) = 0.0013
Patching at the country-name position works for blocks 0 through 8: it consistently lifts P(Paris) about 25 times above baseline. From block 9 onward, the effect fades to nothing.
Line by line
COUNTRY_POS = 3 targets only the token that actually differs between the two prompts. Patching a position that is identical in both prompts anyway would tell you nothing.
The patch replaces one position, not the whole sequence. patched[:, COUNTRY_POS, :] = ... leaves every other position untouched, which is what makes this a precise test rather than a blunt one.
The fade after block 8 is not a bug. By around block 9, the model has already moved the country information from position 3 forward to the final position, where the prediction is actually read out. Patching position 3 that late no longer matters, because the model already used that information earlier.
Common mistakes
Patching the wrong token position. Patching the final position instead of the country-name position tells a different, complementary story. Both are worth trying, since they reveal different steps of the same computation.
Concluding one layer "is" the fact, from a single prompt pair. This experiment uses exactly one France/Germany pair. Real findings need averaging over many different fact pairs, to rule out one prompt being an outlier.
Forgetting to remove the forward hook. A hook left registered keeps patching every future forward pass through that model, silently corrupting anything run afterward. Always pair register_forward_hook with .remove().
Try it yourself
Change COUNTRY_POS to -1, the final token position, and rerun. Compare the pattern of results against the one above.
Where the country-name-position patching worked early and faded late, the final-position patching should show close to the opposite shape: little effect early, then a rise in later layers. That is the same information, caught at a different point in its journey through the model.
What to learn next
- Where a model stores a fact — separating this layer-level result into attention's contribution versus the MLP's.
- Steering a model with activation vectors — using a similar intervention to change behaviour on purpose, not only to test it.
- The logit lens — the technique usually run first, to decide which layers are worth patching.
Researcher — Mathematics and papers.
Formal setup
Given a clean input x_clean (produces the desired output) and a corrupted input x_corrupt (produces an undesired output), and a specific activation site a (a layer, position and, optionally, attention head):
patched_output = run(x_corrupt, with a set to a_clean)
effect(a) = metric(patched_output) - metric(run(x_corrupt))a_cleanis the value of activation sitearecorded from the clean run.metricis typically a logit difference or probability for the answer token, as used in the developer block.effect(a)isolates the causal contribution of exactly that site, holding every other computation on the corrupted input's own path fixed.
This is denoising patching, clean-into-corrupted, used above. The complementary direction, noising (corrupted-into-clean), asks the opposite question: which sites, if broken, are sufficient to break an otherwise-correct output. Both directions appear throughout the literature, and answer subtly different questions.
Causal tracing and the ROME methodology
Meng et al. (2022), Locating and Editing Factual Associations in GPT, formalised a specific protocol called causal tracing. Corrupt the subject tokens with noise, then patch clean activations back one at a time, at every layer and token position. This produces a full grid of restoration effects.
Averaged over many factual-recall prompts (the CounterFact dataset), this produces a consistent, replicated finding. The strongest restoration effect concentrates at mid-layer MLP sublayers, at the last token of the subject, rather than at attention or the final position. This finding directly motivated ROME, a method for editing a stored fact by modifying those MLP weights directly.
Path patching and edge-level attribution
Wang et al. (2022), Interpretability in the Wild, extended activation patching from patching whole nodes to patching specific edges between two components. This isolates not only which layer matters, but which specific upstream head feeds which downstream head. It is how the IOI (indirect object identification) circuit in GPT-2 small was mapped, head by head.
Complexity
Exhaustive activation patching over every layer and token position costs O(L * n) forward passes for a sequence of length n and L layers, each pass otherwise identical in cost to normal inference. Head-level and edge-level patching multiply this further by the number of heads or edges under study, which is why automated circuit-discovery tools generally prune the search space heuristically rather than patching exhaustively.
Key references
- Vig, J. et al. (2020). Investigating Gender Bias in Language Models Using Causal Mediation Analysis. NeurIPS — an early application of this causal framework to bias.
- Meng, K. et al. (2022). Locating and Editing Factual Associations in GPT. arXiv:2202.05262
- Wang, K. et al. (2022). Interpretability in the Wild: A Circuit for Indirect Object Identification in GPT-2 small. arXiv:2211.00593
Current state and open problems
Activation patching, and its refinements such as path patching, are currently the standard toolkit for making a causal claim, not only a correlational one, about what a specific part of a model is doing.
The open problem is scale. Fact-level causal tracing on a small model, as in the developer block, is cheap. Doing the same exhaustive search across every layer, head and position of a frontier-scale model is not. Most current circuit-discovery work either studies small models directly, or uses automated, approximate attribution methods to narrow the search before running any exact patching.
What to learn next
- Where a model stores a fact — the specific MLP-localisation finding this technique produced.
- Sparse autoencoders for feature discovery — a complementary technique for finding interpretable directions, rather than testing a chosen one.
- Induction heads — a circuit whose two-head structure was confirmed using exactly this method.