Skip to content

feat(attribution_patching): faithfulness (edge ablation) on TransformerBridge - #1830

Open
janmenjayap wants to merge 6 commits into
TransformerLensOrg:devfrom
janmenjayap:feat/attribution-patching-faithfulness
Open

janmenjayap wants to merge 6 commits into
TransformerLensOrg:devfrom
janmenjayap:feat/attribution-patching-faithfulness

Conversation

@janmenjayap

@janmenjayap janmenjayap commented Sep 27, 2026 •

Copy link
Copy Markdown
Contributor

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.

from transformer_lens.tools.analysis import (
    EdgeAttributionConfig,
    attribution_patch,
    faithfulness,
)

ranked = attribution_patch(
    model, clean, corrupt, metric_fn, config=EdgeAttributionConfig(granularity="edge")
)
report = faithfulness(model, clean, corrupt, metric_fn, ranked.top_edges(k=200))
print(report.recovered, report.circuit_size, report.total_edges)
  • Forward-only writer-output capture, names-filtered to the writer families the ablation needs (hook_embed, attn.hook_result, hook_mlp_out). Kept separate from cache_activation_and_gradient, whose contract is grad-retaining and which refuses to run without autograd.
  • Reader-input rewrite at the pre-LN fork hooks (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=...) returning recovered, 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.
  • Edge-class breakout (into-Q/K vs into-V vs into-MLP vs into-logits) so the Risk-3 attention-nonlinearity limitation is measured rather than hidden.

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-189 fires hook(forked) then _apply_ln1_per_head; generalized_components/block.py:173 captures "ln2's input on pre-norm blocks"). The correction is consequently a plain add and subtract in d_model space 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

  • Empty circuit reproduces the corrupt metric; full circuit reproduces the clean metric, each with a discriminating assertion so neither can pass on a no-op correction.
  • The mutation-checked reconstruction test from PR2 is reused verbatim for the ablation correction terms. The ablation's correction terms inherit the same indexing failure modes as the scorer, so the guard is shared rather than duplicated.
  • A random edge set of the same size recovers markedly less than the attribution circuit, with both numbers reported.
  • An edge outside the graph raises rather than silently ablating nothing; a zero metric gap raises rather than returning an undefined fraction.

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

Budget Ranked circuit Random edge set
50 +0.293 ~0.000
200 +0.624 ~0.000
1000 +0.997 ~0.000

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.

Class Recovered Edges
into_qk +1.087 125,280
into_v −0.182 62,640
into_mlp +0.086 6,084
into_logits +0.000 942

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

  1. 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.

  2. Batched input is rejected. faithfulness raises 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.

  3. 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 with test_faithfulness_recovery_is_graded_not_binary, which restricts to readout-position edges and asserts 0 < recovered < 1.

  4. 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

  • PR4 - EAP-IG (integrated gradients). ig_steps > 1 continues to raise NotImplementedError.
  • PR5 - oracle parity against the pinned hannamw/eap-ig reference, IOI circuit recovery, and the demo notebook.

The proposal's "vertical slice only" checklist item is satisfied when PR5 merges.


Type of change

  • New feature (non-breaking change which adds functionality)

Checklist

  • I have commented my code, particularly in hard-to-understand areas
  • I have made corresponding changes to the documentation (docs/source/content/analysis_tools.md)
  • My changes generate no new warnings
  • I have added tests that prove my fix is effective or that my feature works
  • New and existing unit tests pass locally with my changes
  • I have not rewritten tests relating to key interfaces which would affect backward compatibility

Testing

Check Result
pytest tests/unit/tools/test_attribution_patching.py 82 passed
pytest tests/integration/test_attribution_patching.py 3 passed (real gpt2-small)
pytest tests/unit -m "not slow" 6723 passed, 0 failed
mypy . clean, 385 source files
black --check / isort --check-only / pycln --check clean

No HookedTransformer reference added; all model interaction goes through
TransformerBridge. No # type: ignore added.


Commits

# Commit
1 feat(attribution_patching): writer-output capture in one forward
2 feat(attribution_patching): reader-input rewrite for edge ablation
3 feat(attribution_patching): faithfulness() entry point
4 feat(attribution_patching): edge-class breakout in the faithfulness report
5 test(attribution_patching): random-edge-set baseline honesty
6 docs(attribution_patching): record the faithfulness contract in the module docstring

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 jlarson4 left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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]

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

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?

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

None yet

Projects

None yet

Development

Successfully merging this pull request may close these issues.

[Proposal] Attribution Patching + EAP/EAP-IG: linearized activation patching and faithful edge-level circuit discovery

2 participants