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.

Read these first

On this page 5
  1. Why it exists
  2. How it works
  3. Where you have already seen it
  4. Remember this
  5. What to learn next

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

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

bash
pip install transformers torch

The first run downloads gpt2, roughly 500 MB.

Patching "France" into a "Germany" prompt, one layer at a time

patching_demo.py
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}")
Output
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

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):

text
patched_output = run(x_corrupt, with a set to a_clean)
effect(a) = metric(patched_output) - metric(run(x_corrupt))
  • a_clean is the value of activation site a recorded from the clean run.
  • metric is 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