Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
44 changes: 44 additions & 0 deletions docs/source/content/analysis_tools.md
Original file line number Diff line number Diff line change
Expand Up @@ -86,6 +86,50 @@ Both node and edge granularity use plain attribution with `ig_steps=1`.

API: {func}`~transformer_lens.tools.analysis.attribution_patching.attribution_patch`.

#### Faithfulness: check that a circuit actually explains the behavior

A ranked edge list is a hypothesis, not a result. `faithfulness()` tests it by
ablating every edge *outside* the candidate circuit and reporting how much of the
clean-to-corrupt metric gap the circuit recovers:

```python
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=50))
print(report.recovered, report.circuit_size, report.total_edges)
```

`recovered` is a fraction, not a percentage: `1.0` means the circuit reproduces
the clean metric, `0.0` means it reproduces the corrupt metric. Values outside
`[0, 1]` are possible and meaningful, since a circuit can overshoot.

The residual stream is a running sum, so each reader's input is rebuilt exactly
by subtracting the excluded writers' live contributions and adding their
replacements. This is a forward-only measurement: it takes no gradient, so it is
independent of the linearization that ranked the edges in the first place.

`FaithfulnessConfig(ablation=...)` selects the replacement. The default,
`"corrupt"`, substitutes the corrupt run's own contribution, which is the
evaluation the external EAP-IG reference reports; `"mean"` substitutes the
dataset mean of that writer's contribution instead. Compare against a random
edge set of the same size before claiming a circuit is meaningful.

`report.edge_class_recovered` breaks the result out by edge class, measured
leave-one-out: each entry is the recovery with that class's edges removed, so a
class whose removal collapses recovery is load-bearing. Edges into Q and K pass
through the softmax, so they are the least faithfully ranked; edges into V, the
MLP, and the terminal readout are linear. `report.edge_class_counts` reports how
many edges each class holds, so a small class is not over-read.

API: {func}`~transformer_lens.tools.analysis.attribution_patching.faithfulness`.

### Direct Path Patching: state the path and its approximation

`get_act_patch_direct_path` fixes a source head and sweeps later destination heads,
Expand Down
113 changes: 109 additions & 4 deletions tests/integration/test_attribution_patching.py
Original file line number Diff line number Diff line change
@@ -1,11 +1,11 @@
"""Integration guard: ``attribution_patch`` yields finite node scores on a real Bridge.
"""Integration guard: attribution patching and faithfulness on a real Bridge. Bridge.

The model-free unit suite runs against ``_LinearToyBridge``, which overrides
``hook_dict``, ``hooks()``, and ``check_hooks_to_add`` — the three behaviours the
two real-Bridge failures depend on. Its hook points have no conversion and its
gate check is a no-op, so a green unit suite says nothing about a real Bridge.

This test boots a real GPT-2 Bridge and exercises the two paths the toy bridge
This test boots a real GPT-2 Bridge and exercises the paths the toy bridge
hides:

- gradients captured *through* hook conversions — ``blocks.*.attn.hook_z`` hands a
Expand All @@ -14,7 +14,9 @@
in that converted shape; and
- the default ``names_filter`` (``None``) falling back to the node hook set on a
real ``hook_dict``, which also exposes gated points (``hook_mlp_in``,
``attn.hook_result``, split-QKV inputs) that ``add_hook`` would reject.
``attn.hook_result``, split-QKV inputs) that ``add_hook`` would reject; and
- the edge-ablation rewrite, whose per-head fork hooks and pre-LN placement only
exist on a real attention bridge.
"""

from __future__ import annotations
Expand All @@ -25,7 +27,11 @@
import torch

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

CLEAN_PROMPT = "The capital of France is"
CORRUPT_PROMPT = "The capital of Russia is"
Expand Down Expand Up @@ -71,3 +77,102 @@ def test_attribution_patch_scores_every_node_finite_on_real_bridge(gpt2_bridge)
# conversion bug broke, mlp_out and embed round out the node graph.
families = {node.kind for node in result.node_scores}
assert families == {"embed", "attn_head_out", "mlp_out"}


def test_faithfulness_recovers_most_of_the_metric_on_gpt2_small(gpt2_bridge) -> None:
"""A ranked edge circuit recovers far more of the gap than a random one.

The ablation rewrites each reader's input at its pre-LN fork hook, so this
only works if the real Bridge's split-QKV and MLP-entry hooks sit where the
rewrite assumes. A circuit built from the top-ranked edges should recover a
large share of the clean-to-corrupt gap; a random edge set of the same size
should recover essentially none, which is what makes the number meaningful
rather than an artifact of ablating almost nothing.

Thresholds are set from measured values on this exact model and prompt pair
(gpt2-small, fp32, CPU): the top 200 edges recover about 0.62 of the gap,
while a random 200 of the 194,946 edges recover about 0.00. Recovery grows
with the budget -- roughly 0.29 at 50 edges and 0.997 at 1000 -- so the
budget is fixed here rather than left to drift.
"""
clean = gpt2_bridge.to_tokens(CLEAN_PROMPT)
corrupt = gpt2_bridge.to_tokens(CORRUPT_PROMPT)
assert clean.shape == corrupt.shape, "prompts must tokenize to the same length"

answer_id = int(gpt2_bridge.to_tokens(" Paris")[0, -1].item())
wrong_id = int(gpt2_bridge.to_tokens(" Moscow")[0, -1].item())
metric_fn = _logit_diff_metric(answer_id, wrong_id)

result = attribution_patch(
gpt2_bridge,
clean,
corrupt,
metric_fn,
config=EdgeAttributionConfig(granularity="edge"),
)
budget = 200
ranked = result.top_edges(k=budget)
assert len(ranked) == budget

report = faithfulness(gpt2_bridge, clean, corrupt, metric_fn, ranked)

assert math.isfinite(report.recovered)
assert report.circuit_size == budget
assert report.total_edges > report.circuit_size
assert report.full_metric != pytest.approx(report.corrupt_metric)
assert report.recovered > 0.5

# A random edge set of the same size recovers essentially nothing, so the
# ranking is doing real work rather than the number being an artifact of the
# circuit's size.
all_edges = list(result.edge_scores)
generator = torch.Generator().manual_seed(0)
sampled = torch.randperm(len(all_edges), generator=generator)[:budget].tolist()
random_report = faithfulness(
gpt2_bridge, clean, corrupt, metric_fn, [all_edges[index] for index in sampled]
)

assert random_report.circuit_size == report.circuit_size
assert random_report.recovered < 0.1
assert report.recovered > random_report.recovered + 0.4


def test_edge_class_breakout_marks_into_qk_as_the_least_faithful_class(gpt2_bridge) -> None:
"""Removing the into-Q/K edges is what breaks the circuit's faithfulness.

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 measures that by keeping every edge outside one class,
so a class whose removal collapses recovery is the one carrying the
nonlinearity.

Thresholds are set from measured values on this exact model and prompt pair
(gpt2-small, fp32, CPU): 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.
"""
clean = gpt2_bridge.to_tokens(CLEAN_PROMPT)
corrupt = gpt2_bridge.to_tokens(CORRUPT_PROMPT)
answer_id = int(gpt2_bridge.to_tokens(" Paris")[0, -1].item())
wrong_id = int(gpt2_bridge.to_tokens(" Moscow")[0, -1].item())
metric_fn = _logit_diff_metric(answer_id, wrong_id)

result = attribution_patch(
gpt2_bridge,
clean,
corrupt,
metric_fn,
config=EdgeAttributionConfig(granularity="edge"),
)
report = faithfulness(gpt2_bridge, clean, corrupt, metric_fn, result.top_edges(k=200))

breakout = report.edge_class_recovered
assert set(breakout) == {"into_qk", "into_v", "into_mlp", "into_logits"}
assert all(math.isfinite(value) for value in breakout.values())
assert sum(report.edge_class_counts.values()) == report.total_edges

# Dropping the softmax-fed class leaves the circuit intact; dropping the
# linear classes does not.
assert breakout["into_qk"] > 0.9
assert breakout["into_v"] < 0.5
assert breakout["into_logits"] < 0.5
Loading
Loading