feat(attribution_patching): faithfulness (edge ablation) on TransformerBridge - #1830
janmenjayap wants to merge 6 commits into
Conversation
Add faithfulness(), which ablates every edge outside a candidate circuit and reports how much of the clean-to-corrupt metric gap the circuit recovers. A ranked edge list is a hypothesis; this measures whether it actually explains the behavior. The measurement is forward-only, so it takes no gradient and is independent of the linearization that ranked the edges. FaithfulnessConfig(ablation=...) selects the replacement for an ablated writer. The default, corrupt, substitutes the corrupt run's own contribution, matching the evaluation the external EAP-IG reference reports so the two compare like quantities; mean substitutes the dataset mean instead. The circuit accepts either bare (writer, reader) pairs or the (writer, reader, score) triples top_edges returns, so a discovered edge set feeds straight in. An edge outside the graph raises rather than silently ablating nothing, and a zero metric gap raises rather than returning an undefined fraction. Batched input is rejected: the metric reads a single example, so a batch would need per-pair aggregation the caller should own, and silently ablating only the first row would return a plausible but wrong number. Export faithfulness, FaithfulnessConfig, and FaithfulnessResult, and document the API and the ablation default.
…eport Add EdgeClass and _edge_class, and populate FaithfulnessResult.edge_class_recovered and edge_class_counts. Edges into Q and K pass through the softmax, so the linearized ranking is least trustworthy there; edges into V, the MLP, and the terminal readout are linear. The breakout makes that asymmetry measurable instead of hiding it in one aggregate. Each class is measured leave-one-out: keep every edge outside it, so the entry says how much of the gap survives without that class. Keeping only the class instead reports roughly zero for every class, since no single class alone reconstructs the behavior, so the leave-one-out form is the one that carries signal. Move the zero-gap guard ahead of the breakout, which divides by the gap. Thresholds in the integration test are set from measured values on gpt2-small: keeping everything except into-Q/K recovers about 1.09 of the gap, while removing into-V drops it to about -0.18 and removing into-logits to about 0.00.
Add a nonlinear edge toy and a test asserting the ranked circuit beats same-size random edge sets by a wide margin. A ranking that does no better than chance still produces a plausible-looking recovered number, so the comparison against random sets is what makes the number mean anything. The toy is nonlinear on purpose. On a linear model any circuit holding every edge into the logits reader recovers the metric outright, so the random baseline could tie by accident; a GELU on the attention output path removes that degeneracy. The block gains an overridable attention-output transform so the nonlinear variant reuses the parent's hook graph. The budget is fixed at 8 edges, where the ranked circuit recovers about 1.89 of the gap and random sets recover about 0.00. Larger budgets are not usable as a baseline: at 16 edges a random set can overshoot to 8.29, since ablating a large arbitrary set can move the metric arbitrarily far. The test also asserts the baseline is not degenerate, and a companion test pins the draw to a fixed seed.
…odule docstring The module docstring still described ablate-outside faithfulness as unimplemented, which stopped being true when the entry point landed. It now states what faithfulness does and the two facts the exactness of the correction rests on: the residual stream is a running sum, so a reader's input is exactly the sum of its writers' contributions, and the reader fork hooks fire before the layer norm, so the correction needs no norm scale. The pre-LN placement is the assumption the whole ablation rests on and the first thing a reviewer will question, so it belongs in the module-level contract rather than only in the helper that implements it.
jlarson4
left a comment
There was a problem hiding this comment.
Here is the first round of review for this one!
|
|
||
| class_recovered: dict[EdgeClass, float] = {} | ||
| for edge_class in class_edges: | ||
| outside = [edge for edge in edges if _edge_class(edge) != edge_class] |
There was a problem hiding this comment.
This drops each class from the full graph, not from circuit, so edge_class_recovered is identical for every circuit passed in, from empty to full. It describes the prompt rather than the circuit, so it can't show whether the ranking is weaker on Q/K edges. Could you rework the breakout so it's measured against the circuit being evaluated?
| """ | ||
| totals: dict[str, torch.Tensor] = {} | ||
| count = 0 | ||
| for tokens in (clean, corrupt): |
There was a problem hiding this comment.
With one pair enforced, this mean is the midpoint of the clean and corrupt runs, so every ablated writer keeps half its clean contribution. An empty circuit then recovers about half the gap, because the scored example sets its own baseline. Could you source the mean from something other than the pair being evaluated?
| metric=clean_cache.metric, | ||
| ) | ||
|
|
||
| mutated_scores = _edge_effects(perturbed_clean_cache, corrupt_cache, edges) |
There was a problem hiding this comment.
This scores edges with _edge_effects and never calls _ablate_edges, so it repeats the scorer guard instead of testing the ablation. The empty and full boundary tests can't fill that gap: an empty circuit rebuilds the logits reader from replacements alone, which erases any upstream error, and a full circuit installs no correction hooks. Could you pin a partial circuit's ablation against an independently computed reference, on the gpt2 bridge as well as the toy?
| replacements = _ablation_replacement(model, corrupt) | ||
|
|
||
| with _edge_hook_flags(model): | ||
| ablated_logits = _ablate_edges(model, corrupt, edges, [], replacements) |
There was a problem hiding this comment.
This ablates the corrupt tokens with corrupt replacements, so every correction is zero and the test passes even with the ablation hooks removed. test_faithfulness_recovers_nothing_for_the_empty_circuit already covers this boundary on the clean tokens. Could this test go?
Description
Third PR of the five-PR vertical slice tracked by #1742. PR1 (#1750) shipped the node substrate; PR2 (#1781) shipped edge scoring. This PR adds the Risk-2 subsystem: edge ablation and the
faithfulness()report.Part of #1742
What this adds
A ranked edge list is a hypothesis, not a result.
faithfulness()tests it by ablating every edge outside a candidate circuit and reporting how much of the clean-to-corrupt metric gap the circuit recovers.hook_embed,attn.hook_result,hook_mlp_out). Kept separate fromcache_activation_and_gradient, whose contract is grad-retaining and which refuses to run without autograd.hook_q/k/v_input,hook_mlp_in), subtracting excluded writers' live contributions and adding their replacements.faithfulness(model, clean, corrupt, metric_fn, circuit, config=...)returningrecovered,full_metric,corrupt_metric,circuit_size,total_edges, plus a per-edge-class breakout.FaithfulnessConfig(ablation=...), defaulting to"corrupt"(matching the pinned oracle's headline evaluation), with"mean"as an option.Why the correction is exact
The residual stream is a running sum, so a reader's input is exactly the sum of its writers' contributions. Rebuilding that input from the included writers plus replacements for the excluded ones is therefore an exact correction, not an approximation.
The reader fork hooks fire before the layer norm (verified on
dev:generalized_components/attention.py:186-189fireshook(forked)then_apply_ln1_per_head;generalized_components/block.py:173captures "ln2's input on pre-norm blocks"). The correction is consequently a plain add and subtract ind_modelspace with no norm scale to divide out. This invariant is recorded in the module docstring, since it is the assumption the whole PR rests on.Several reader nodes share one hook point (one per head, one per position), so the rewrite is grouped by reader node and each hook touches only its own head slot. One hook per excluded edge would apply the correction repeatedly and overshoot.
Guards
Measured results
gpt2-small, fp32, CPU,"The capital of France is"/"The capital of Russia is", Paris-minus-Moscow logit difference. Graph size: 194,946 edges.Recovery vs circuit budget
Edge-class breakout (leave-one-out, budget 200)
Each entry is the recovery with that class's edges removed, so a class whose removal collapses recovery is load-bearing.
Reading: removing the softmax-fed class leaves the circuit intact, while removing the linear classes collapses it. That is the Risk-3 asymmetry made visible.
Deviations from the plan, with reasons
The breakout is leave-one-out, not keep-only-class. The plan specified restricting the circuit to each class's edges. Measured on
gpt2-small, that framing is degenerate: every class reports ~0.0000, because no single class alone reconstructs the behavior. Leave-one-out is informative (table above). The field docstring records why.Batched input is rejected.
faithfulnessraises on more than one clean/corrupt pair. The metric reads a single example, so a batch would need per-pair aggregation the caller should own; silently ablating only the first row would return a plausible but wrong number.The plan's monotonicity test was replaced. The plan listed
test_faithfulness_recovers_more_for_a_larger_circuit. On this position-wise toy only edges into the readout position can move the metric, so a prefix of the edge list is largely invisible to it and the monotonicity claim was not testable as written. Replaced withtest_faithfulness_recovery_is_graded_not_binary, which restricts to readout-position edges and asserts0 < recovered < 1.The random-baseline budget is fixed at 8 edges on the nonlinear toy. At 16 edges a random set can overshoot to +8.29, since ablating a large arbitrary set can move the metric arbitrarily far; 8 is where the margin is unambiguous (ranked +1.89 vs random ~0.00).
Still to land
ig_steps > 1continues to raiseNotImplementedError.hannamw/eap-igreference, IOI circuit recovery, and the demo notebook.The proposal's "vertical slice only" checklist item is satisfied when PR5 merges.
Type of change
Checklist
docs/source/content/analysis_tools.md)Testing
pytest tests/unit/tools/test_attribution_patching.pypytest tests/integration/test_attribution_patching.pygpt2-small)pytest tests/unit -m "not slow"mypy .black --check/isort --check-only/pycln --checkNo
HookedTransformerreference added; all model interaction goes throughTransformerBridge. No# type: ignoreadded.Commits
feat(attribution_patching): writer-output capture in one forwardfeat(attribution_patching): reader-input rewrite for edge ablationfeat(attribution_patching): faithfulness() entry pointfeat(attribution_patching): edge-class breakout in the faithfulness reporttest(attribution_patching): random-edge-set baseline honestydocs(attribution_patching): record the faithfulness contract in the module docstring