AtP*: attribution patching at industrial scale
DeepMind found two ways the cheap gradient approximation to activation patching misses important components, and patched both without giving up the speed. A method deep dive on why the linear approximation breaks at saturated attention and how the fix works.
The cost problem with activation patching
Activation patching is the honest way to ask which parts of a model cause a behaviour. You run a clean prompt and a corrupted prompt, swap one component's activation from the corrupted run into the clean run, and measure how much the output moves. The trouble is the sweep. Every component needs its own forward pass, and once you count attention heads per position and individual MLP neurons per position, the number of components in a 12B model runs into the millions.
Attribution patching, AtP for short, is the shortcut. Take a first order Taylor expansion of the loss around the clean activation, and the effect of patching a node becomes the difference between its corrupted and clean activation dotted with the gradient of the loss at that node. Two forward passes and one backward pass give you an estimate for every node at once. János Kramár, Tom Lieberum, Rohin Shah and Neel Nanda at Google DeepMind posted a paper on March 1 that asks how often that shortcut is wrong, and what to do about it.
Where the linear approximation lies
The first failure is attention saturation. Queries and keys feed a softmax, and when the attention pattern is already near zero or near one for a given pair of positions, the softmax sits in a flat region. A linear approximation of a flat function says nothing happens when you move the input, but a large enough patch can push the attention weight across the whole curve. The gradient is small and the true effect is large, so AtP reports a false negative on exactly the heads that matter most for a sharp behaviour.
The second failure is cancellation. A node can have a direct effect on the logits and an indirect effect routed through later components, and the two can have opposite signs of similar size. The paper gives examples of nodes whose direct effect ratio is 5.35, 12.2, or zero, meaning the direct path is several times bigger than the net effect. When you approximate two large terms and subtract them, small errors in each swamp the small true difference. That is a general property of the estimator, and no amount of extra precision in the gradient fixes it.
The two patches that make AtP*
For saturation the answer is to stop linearising the part that is nonlinear. The QK fix recomputes the actual change in attention weights when the query or key is patched, and applies the gradient only downstream of the softmax, where the approximation is reasonable. For queries this costs roughly one extra attention computation. For keys a naive version would cost O(T cubed) in sequence length, and the paper gives a variant that gets it down to O(T squared). In total the QK fix adds less than one forward pass worth of compute.
For cancellation the answer is GradDrop. For each of the L layers, compute an AtP estimate with that layer's gradient contribution zeroed out, so that any cancellation routed through that layer is broken. Then average the absolute values across the L estimates. That is L backward passes instead of one, but each reuses the same stored activations, so it remains far cheaper than a per node sweep. AtP with both fixes is what the authors call AtP*.
What the comparison showed
The test bed is the Pythia suite from 410M up to 12B parameters, with two kinds of node. Attention nodes split each head at each position into query, key, value and output. Neuron nodes are individual MLP neurons per position. Prompt pairs come from indirect object identification, a factual recall task about cities, and a random pair drawn from the Pile. The metric is the negative log probability of the clean target token.
The question the paper actually asks is how many forward passes a method needs before its ranking includes the true top nodes, where truth comes from an exhaustive sweep. On Pythia-12B, AtP* recovers the top 100 MLP neurons on the city task within about 500 forward passes, and the top 100 attention nodes on the IOI task within about 700. The baselines are honest alternatives, including random subsampling, patching blocks of nodes and recursively subdividing the hot ones, and an iterative brute force in a smart order. AtP already beats all of them, and AtP* is a clear further step, with the QK fix doing most of the work on attention nodes.
One detail we found useful. When the authors averaged over a distribution of prompts rather than a single pair, AtP with only the QK fix did about as well as full AtP*. Cancellation seems to be a property of specific prompt pairs, and it washes out when you average. That tells you which fix to reach for depending on how you set up the experiment.
Knowing what you missed
The part of the paper we expect to reuse is the diagnostic. A fast approximate method is only useful if you can say how bad its false negatives are. The authors take the top K nodes found by AtP*, then patch random subsets of the remaining nodes and measure the effects. Using Welch's t-test they invert the hypothesis test to find the smallest effect size that could still be hiding among the unranked nodes at 90, 99 or 99.9 percent confidence. That turns a heuristic into a bound.
What we would like to see next is the same treatment applied to sparse autoencoder features rather than raw neurons and heads, since that is where much of the field is moving and the node count is even larger. We would also want to know how the saturation failure behaves in models trained with newer attention variants, where the softmax geometry is different. For now, if you are localising a behaviour in a model with more than a few billion parameters, this paper gives you a defensible reason to skip the full sweep.
Sources
From the foundation