From 824430847170bb0233224fe741be9bc2f65fc281 Mon Sep 17 00:00:00 2001 From: janmenjayap Date: Sun, 27 Sep 2026 13:21:26 +0530 Subject: [PATCH 1/6] feat(attribution_patching): writer-output capture in one forward --- tests/unit/tools/test_attribution_patching.py | 105 ++++++++++++++++++ .../tools/analysis/attribution_patching.py | 104 +++++++++++++++++ 2 files changed, 209 insertions(+) diff --git a/tests/unit/tools/test_attribution_patching.py b/tests/unit/tools/test_attribution_patching.py index c2e809a2d6..402f0f814b 100644 --- a/tests/unit/tools/test_attribution_patching.py +++ b/tests/unit/tools/test_attribution_patching.py @@ -31,10 +31,13 @@ _edge_hook_names, _ensure_edge_hook_flags, _node_effects, + _reader_hook_names, _required_hook_names, _writer_hook_name, + _writer_hook_names, attribution_patch, cache_activation_and_gradient, + capture_writer_outputs, enumerate_edges, enumerate_nodes, ) @@ -1955,3 +1958,105 @@ def test_top_edges_can_surface_a_head_to_logits_edge() -> None: # At least one carries a real signed effect, so the ranking is not over zeros. assert any(abs(score) > 1e-6 for _writer, _reader, score in head_to_logits) + + +# --------------------------------------------------------------------------- +# Writer-output capture (forward-only) +# --------------------------------------------------------------------------- +# +# Rewriting a reader's input during a live forward pass needs the writers' +# current contributions rather than a gradient estimate. The capture is +# therefore forward-only and names-filtered to the writer families, the +# opposite contract from cache_activation_and_gradient. + + +def test_writer_hook_names_cover_the_three_writer_families() -> None: + names = _writer_hook_names(N_LAYERS) + + assert names[0] == "hook_embed" + for layer in range(N_LAYERS): + assert f"blocks.{layer}.attn.hook_result" in names + assert f"blocks.{layer}.hook_mlp_out" in names + # Writer families only: no reader inputs, and no pre-hook_result attn-z. + assert not any("hook_q_input" in name for name in names) + assert not any("hook_mlp_in" in name for name in names) + assert not any(name.endswith("attn.hook_z") for name in names) + + +def test_reader_hook_names_cover_the_reader_input_families() -> None: + names = _reader_hook_names(N_LAYERS) + + for layer in range(N_LAYERS): + for input_hook in ("hook_q_input", "hook_k_input", "hook_v_input"): + assert f"blocks.{layer}.attn.{input_hook}" in names + assert f"blocks.{layer}.hook_mlp_in" in names + # Reader inputs only: no writer outputs. + assert not any("hook_result" in name for name in names) + assert not any("hook_mlp_out" in name for name in names) + + +def test_capture_writer_outputs_populates_only_the_writer_family() -> None: + model = _EdgeScoringToyBridge() + tokens = model.to_tokens("prompt") + + with _edge_hook_flags(model): + captured = capture_writer_outputs(model, tokens) + + assert set(captured) == set(_writer_hook_names(N_LAYERS)) + # The reader-input family and the pre-hook_result attn-z are not captured. + assert "blocks.0.attn.hook_q_input" not in captured + assert "blocks.0.hook_mlp_in" not in captured + assert "blocks.0.attn.hook_z" not in captured + + assert captured["hook_embed"].shape == (1, SEQ_LEN, D_MODEL) + for layer in range(N_LAYERS): + assert captured[f"blocks.{layer}.attn.hook_result"].shape == ( + 1, + SEQ_LEN, + N_HEADS, + D_MODEL, + ) + assert captured[f"blocks.{layer}.hook_mlp_out"].shape == ( + 1, + SEQ_LEN, + D_MODEL, + ) + + +def test_capture_writer_outputs_returns_detached_tensors() -> None: + model = _EdgeScoringToyBridge() + tokens = model.to_tokens("prompt") + + with _edge_hook_flags(model): + captured = capture_writer_outputs(model, tokens) + + for tensor in captured.values(): + assert not tensor.requires_grad + assert tensor.grad_fn is None + + +def test_capture_writer_outputs_raises_when_the_filter_matches_nothing() -> None: + model = _EdgeScoringToyBridge() + tokens = model.to_tokens("prompt") + + with pytest.raises(ValueError, match="matched no hook points"): + capture_writer_outputs(model, tokens, names_filter=["blocks.0.not_a_hook"]) + + +def test_capture_writer_outputs_needs_no_grad() -> None: + """The capture runs with autograd off, where the gradient helper refuses. + + The two helpers have opposite contracts: this one is forward-only and must + work under ``torch.no_grad()``, while ``cache_activation_and_gradient`` + exists to retain gradients and raises without autograd. + """ + model = _EdgeScoringToyBridge() + tokens = model.to_tokens("prompt") + metric = _metric_fn(answer=1, wrong=2) + + with _edge_hook_flags(model), torch.no_grad(): + captured = capture_writer_outputs(model, tokens) + assert set(captured) == set(_writer_hook_names(N_LAYERS)) + + with pytest.raises(ValueError, match="autograd"): + cache_activation_and_gradient(model, tokens, metric) diff --git a/transformer_lens/tools/analysis/attribution_patching.py b/transformer_lens/tools/analysis/attribution_patching.py index 6cf392d34c..aac47120b3 100644 --- a/transformer_lens/tools/analysis/attribution_patching.py +++ b/transformer_lens/tools/analysis/attribution_patching.py @@ -453,6 +453,40 @@ def _edge_hook_names(n_layers: int) -> list[str]: ) +def _writer_hook_names(n_layers: int) -> list[str]: + """The hook points holding a writer's own residual-stream contribution. + + The three writer families an edge sweep scores: the token embedding write, + each attention head's per-head output (``attn.hook_result``, the + decomposition of the head's contribution after it is projected into the + residual stream), and each layer's MLP output. These are the hook points + :func:`_writer_hook_name` resolves a single writer node to, listed as a + family so a capture can filter to exactly them. + """ + names = ["hook_embed"] + for layer in range(n_layers): + names.append(f"blocks.{layer}.attn.hook_result") + names.append(f"blocks.{layer}.hook_mlp_out") + return names + + +def _reader_hook_names(n_layers: int) -> list[str]: + """The hook points holding a reader's input. + + The residual each reader consumes: the split ``attn.hook_q_input`` / + ``hook_k_input`` / ``hook_v_input`` per head, and each layer's MLP entry + ``hook_mlp_in``. The terminal logits reader reads the final + ``hook_resid_post``, which :func:`_required_edge_reader_hook_names` lists. + """ + names: list[str] = [] + for layer in range(n_layers): + names.append(f"blocks.{layer}.attn.hook_q_input") + names.append(f"blocks.{layer}.attn.hook_k_input") + names.append(f"blocks.{layer}.attn.hook_v_input") + names.append(f"blocks.{layer}.hook_mlp_in") + return names + + def enumerate_edges(model: Any, cache: GradientCache) -> list[tuple[Node, Node]]: """Enumerate every writer -> reader edge in the residual-stream graph. @@ -707,6 +741,76 @@ def hook(grad: torch.Tensor, *, hook: Any) -> None: return GradientCache(activations=activations, gradients=gradients, metric=metric.detach()) +def capture_writer_outputs( + model: Any, + tokens: torch.Tensor, + names_filter: NamesFilter = None, +) -> dict[str, torch.Tensor]: + """Capture writer contributions in one forward-only pass. + + Rewriting a reader's input during a live forward pass needs the writers' + *current* contributions rather than a gradient estimate. This runs a single + forward under :func:`torch.no_grad` with forward hooks only: no backward is + driven and no gradient is retained. That is the opposite contract from + :func:`cache_activation_and_gradient`, which exists to retain gradients and + refuses to run without autograd, so the two are kept as separate helpers + rather than one helper with a mode switch. + + The default filter is the writer-output family only -- ``hook_embed``, + ``blocks.*.attn.hook_result``, and ``blocks.*.hook_mlp_out``. Caching the + reader-input family as well would roughly double the captured memory for no + benefit here, since a caller rewriting a reader's input reads that input + from the live forward rather than from a cache. + + Memory caveat: ``attn.hook_result`` is a per-head + ``[batch, seq, n_heads, d_model]`` tensor, so the capture is fine on a model + the size of gpt2-small and does not scale to models with many heads or + layers. The hook point only fires while ``cfg.use_attn_result`` is on; see + :func:`_edge_hook_flags`. + + Args: + model: A ``TransformerBridge`` (or compatible) exposing ``cfg.n_layers``, + ``hook_dict``, and the ``hooks()`` context manager. + tokens: Input token ids for a single forward pass. + names_filter: Restricts which hook points are captured. ``None`` (the + default) captures the writer-output family. On a real Bridge the + gated points (``attn.hook_result``, the split-QKV inputs, + ``hook_mlp_in``) raise in ``add_hook`` unless their ``set_use_*`` + flag is on, so a filter reaching them must be paired with + :func:`_edge_hook_flags`. + + Returns: + Detached activation tensors keyed by hook name, for every hook point + matching ``names_filter``. + + Raises: + ValueError: if ``names_filter`` matches no hook point. + """ + if names_filter is None: + names_filter = _writer_hook_names(int(model.cfg.n_layers)) + predicate = _as_predicate(names_filter) + names = [name for name in model.hook_dict if predicate(name)] + if not names: + raise ValueError("names_filter matched no hook points") + + captured: dict[str, torch.Tensor] = {} + + def make_fwd_hook(name: str) -> Callable[..., None]: + def hook(tensor: torch.Tensor, *, hook: Any) -> None: + del hook + if isinstance(tensor, torch.Tensor): + captured[name] = tensor.detach().clone() + return None + + return hook + + fwd_hooks = [(name, make_fwd_hook(name)) for name in names] + with torch.no_grad(), model.hooks(fwd_hooks=fwd_hooks): + model(tokens) + + return captured + + def _node_effects( clean_cache: GradientCache, corrupt_cache: GradientCache, From 6bc510f43939714187cc1b4055e71bb8fa67d5ed Mon Sep 17 00:00:00 2001 From: janmenjayap Date: Sun, 27 Sep 2026 18:03:14 +0530 Subject: [PATCH 2/6] feat(attribution_patching): reader-input rewrite for edge ablation --- tests/unit/tools/test_attribution_patching.py | 307 ++++++++++++++++-- .../tools/analysis/attribution_patching.py | 176 +++++++++- 2 files changed, 459 insertions(+), 24 deletions(-) diff --git a/tests/unit/tools/test_attribution_patching.py b/tests/unit/tools/test_attribution_patching.py index 402f0f814b..e5a96fce04 100644 --- a/tests/unit/tools/test_attribution_patching.py +++ b/tests/unit/tools/test_attribution_patching.py @@ -24,12 +24,14 @@ EdgeAttributionConfig, GradientCache, Node, + _ablate_edges, _assert_edges_unique, _check_required_hooks, _edge_effects, _edge_hook_flags, _edge_hook_names, _ensure_edge_hook_flags, + _excluded_writers_by_reader, _node_effects, _reader_hook_names, _required_hook_names, @@ -1722,6 +1724,53 @@ def test_edge_scores_sum_to_the_readers_direct_input_change() -> None: assert checked > 0 +def _assert_mutation_only_changes_the_perturbed_writers_edges( + edges: list[tuple[Node, Node]], + writer: Node, + perturbation: torch.Tensor, + baseline_scores: dict[tuple[Node, Node], float], + mutated_scores: dict[tuple[Node, Node], float], + corrupt_cache: GradientCache, +) -> None: + """Assert a single-writer perturbation moved exactly that writer's edges. + + Both the edge scorer and the ablation correction read a writer's delta and a + reader's gradient from independently indexed tensors (writer position/head, + reader position/head). A slicing bug that mixed up either index could leak a + perturbation into an edge whose writer was never touched, or fail to move an + edge whose writer was. Perturbing a single head's slice of a shared per-head + tensor also exercises the narrower case: sibling heads and positions inside + the *same* tensor must stay untouched. + + Shared by the edge-scoring and edge-ablation tests so the two cannot drift + apart: the ablation's correction terms inherit the same indexing failure + modes as the scorer. + """ + changed = 0 + moved = 0 + for edge in edges: + edge_writer, reader = edge + if edge_writer == writer: + grad = corrupt_cache.gradients[reader.hook_name] + assert grad is not None + if reader.kind in ("q_input", "k_input", "v_input"): + grad_vec = grad[0, reader.position, reader.head] + else: + grad_vec = grad[0, reader.position] + shift = float((perturbation * grad_vec).sum()) + expected = baseline_scores[edge] + shift + assert mutated_scores[edge] == pytest.approx(expected) + changed += 1 + if abs(shift) > 1e-6: + assert mutated_scores[edge] != pytest.approx(baseline_scores[edge]) + moved += 1 + else: + assert mutated_scores[edge] == baseline_scores[edge] + + assert changed > 0 + assert moved > 0 # the perturbation genuinely shifts scores at the readout position + + def test_edge_effects_mutation_only_changes_the_perturbed_writers_edges() -> None: """Perturbing one writer's captured contribution changes only that writer's edges. @@ -1771,29 +1820,9 @@ def test_edge_effects_mutation_only_changes_the_perturbed_writers_edges() -> Non mutated_scores = _edge_effects(perturbed_clean_cache, corrupt_cache, edges) - changed = 0 - moved = 0 - for edge in edges: - edge_writer, reader = edge - if edge_writer == writer: - grad = corrupt_cache.gradients[reader.hook_name] - assert grad is not None - if reader.kind in ("q_input", "k_input", "v_input"): - grad_vec = grad[0, reader.position, reader.head] - else: - grad_vec = grad[0, reader.position] - shift = float((perturbation * grad_vec).sum()) - expected = baseline_scores[edge] + shift - assert mutated_scores[edge] == pytest.approx(expected) - changed += 1 - if abs(shift) > 1e-6: - assert mutated_scores[edge] != pytest.approx(baseline_scores[edge]) - moved += 1 - else: - assert mutated_scores[edge] == baseline_scores[edge] - - assert changed > 0 - assert moved > 0 # the perturbation genuinely shifts scores at the readout position + _assert_mutation_only_changes_the_perturbed_writers_edges( + edges, writer, perturbation, baseline_scores, mutated_scores, corrupt_cache + ) # --------------------------------------------------------------------------- @@ -2060,3 +2089,235 @@ def test_capture_writer_outputs_needs_no_grad() -> None: with pytest.raises(ValueError, match="autograd"): cache_activation_and_gradient(model, tokens, metric) + + +# --------------------------------------------------------------------------- +# Reader-input rewrite (edge ablation core) +# --------------------------------------------------------------------------- +# +# The residual stream is a running sum, so a reader's input is exactly the sum +# of its writers' contributions. Ablating an edge means subtracting that +# writer's live contribution and adding its replacement at the reader's fork +# hook, which fires before the layer norm -- an exact correction, not an +# approximation. The two boundary cases pin the rewrite: excluding every edge +# must reproduce the corrupt run, excluding none must reproduce the clean run. + + +def _ablation_replacement( + model: _EdgeScoringToyBridge, tokens: torch.Tensor +) -> dict[str, torch.Tensor]: + """The value standing in for an ablated writer: the corrupt run's own.""" + with _edge_hook_flags(model): + return capture_writer_outputs(model, tokens) + + +def test_excluded_writers_by_reader_groups_only_out_of_circuit_edges() -> None: + model = _EdgeScoringToyBridge() + _ensure_edge_hook_flags(model) + tokens = model.to_tokens("prompt") + metric = _metric_fn(answer=1, wrong=2) + edge_hook_names = _edge_hook_names(N_LAYERS) + + corrupt_cache = cache_activation_and_gradient( + model, tokens, metric, names_filter=edge_hook_names + ) + edges = enumerate_edges(model, corrupt_cache) + + # Keeping the whole graph excludes nothing. + assert _excluded_writers_by_reader(edges, edges) == {} + + # Keeping nothing excludes every writer, grouped once per reader. + excluded = _excluded_writers_by_reader(edges, []) + readers = {reader for _writer, reader in edges} + assert set(excluded) == readers + for reader, writers in excluded.items(): + expected = [writer for writer, edge_reader in edges if edge_reader == reader] + assert writers == expected + + # A single kept edge excludes exactly that edge's writer from its reader. + # Pick a reader with more than one incoming edge, so the exclusion is + # observable rather than emptying the reader's writer list. + kept_edge = next( + edge for edge in edges if len([writer for writer, reader in edges if reader == edge[1]]) > 1 + ) + excluded = _excluded_writers_by_reader(edges, [kept_edge]) + kept_writer, kept_reader = kept_edge + incoming = [writer for writer, reader in edges if reader == kept_reader] + assert kept_writer not in excluded[kept_reader] + assert len(excluded[kept_reader]) == len(incoming) - 1 + + +def test_ablating_the_empty_circuit_reproduces_the_corrupt_metric() -> None: + """Excluding every edge rebuilds each reader's input as the corrupt run's.""" + model = _EdgeScoringToyBridge() + _ensure_edge_hook_flags(model) + clean = torch.tensor([[1, 2, 3]]) + corrupt = torch.tensor([[3, 2, 1]]) + metric = _metric_fn(answer=1, wrong=2) + edge_hook_names = _edge_hook_names(N_LAYERS) + + corrupt_cache = cache_activation_and_gradient( + model, corrupt, metric, names_filter=edge_hook_names + ) + clean_cache = cache_activation_and_gradient( + model, clean, metric, names_filter=edge_hook_names, compute_gradient=False + ) + edges = enumerate_edges(model, corrupt_cache) + replacements = _ablation_replacement(model, corrupt) + + with _edge_hook_flags(model): + ablated_logits = _ablate_edges(model, corrupt, edges, [], replacements) + + ablated = float(metric(ablated_logits)) + assert ablated == pytest.approx(float(corrupt_cache.metric), abs=1e-6) + # Discriminating: the correction genuinely moved the run, so this cannot + # pass on an ablation that silently did nothing. + assert ablated != pytest.approx(float(clean_cache.metric), abs=1e-6) + + +def test_ablating_the_full_circuit_reproduces_the_clean_metric() -> None: + """Keeping every edge leaves each reader's input untouched.""" + model = _EdgeScoringToyBridge() + _ensure_edge_hook_flags(model) + clean = torch.tensor([[1, 2, 3]]) + corrupt = torch.tensor([[3, 2, 1]]) + metric = _metric_fn(answer=1, wrong=2) + edge_hook_names = _edge_hook_names(N_LAYERS) + + corrupt_cache = cache_activation_and_gradient( + model, corrupt, metric, names_filter=edge_hook_names + ) + clean_cache = cache_activation_and_gradient( + model, clean, metric, names_filter=edge_hook_names, compute_gradient=False + ) + edges = enumerate_edges(model, corrupt_cache) + replacements = _ablation_replacement(model, corrupt) + + with _edge_hook_flags(model): + ablated_logits = _ablate_edges(model, clean, edges, edges, replacements) + + ablated = float(metric(ablated_logits)) + assert ablated == pytest.approx(float(clean_cache.metric), abs=1e-6) + # Discriminating: the clean and corrupt runs genuinely differ, so this + # cannot pass on a degenerate pair. + assert ablated != pytest.approx(float(corrupt_cache.metric), abs=1e-6) + + +def test_ablation_installs_one_hook_per_reader_node() -> None: + """Each reader node is corrected once, not once per incoming edge. + + Several reader nodes share one hook point (one per head, one per position), + and each rewrites a disjoint slice of it. Installing one hook per excluded + edge would apply the correction repeatedly and overshoot, so the count of + hooks on each fork point must equal the number of reader nodes reading it -- + not the number of edges into those nodes. + """ + model = _EdgeScoringToyBridge() + _ensure_edge_hook_flags(model) + tokens = model.to_tokens("prompt") + metric = _metric_fn(answer=1, wrong=2) + edge_hook_names = _edge_hook_names(N_LAYERS) + + corrupt_cache = cache_activation_and_gradient( + model, tokens, metric, names_filter=edge_hook_names + ) + edges = enumerate_edges(model, corrupt_cache) + replacements = _ablation_replacement(model, tokens) + + readers = {reader for _writer, reader in edges} + assert len(readers) > 0 + # Some readers have several incoming edges, so a per-edge hook would + # register several hooks on the same point. + crowded = { + reader + for reader in readers + if len([writer for writer, edge_reader in edges if edge_reader == reader]) > 1 + } + assert crowded + + # The ablation's hooks live only for the duration of its forward, so count + # them from a pre-hook on the model root, which fires while they are + # registered. + counts: dict[str, int] = {} + + def snapshot(_module: nn.Module, _args: tuple) -> None: + for reader in readers: + counts[reader.hook_name] = len(model.hook_dict[reader.hook_name].fwd_hooks) + + handle = model.register_forward_pre_hook(snapshot) + try: + with _edge_hook_flags(model): + _ablate_edges(model, tokens, edges, [], replacements) + finally: + handle.remove() + + expected = { + name: len([reader for reader in readers if reader.hook_name == name]) + for name in {reader.hook_name for reader in readers} + } + assert counts == expected + # Fewer hooks than edges: the grouping collapsed the per-edge corrections. + assert sum(counts.values()) < len(edges) + + +def test_ablation_correction_terms_isolate_the_intended_edges() -> None: + """Perturbing one writer's contribution moves only that writer's edges. + + The ablation's correction terms read a writer's contribution and a reader's + position from independently indexed tensors, the same failure mode the edge + scorer has. This reuses the scorer's mutation guard rather than restating it, + so the two cannot drift apart. + """ + model = _EdgeScoringToyBridge() + _ensure_edge_hook_flags(model) + clean = torch.tensor([[1, 2, 3]]) + corrupt = torch.tensor([[3, 2, 1]]) + metric = _metric_fn(answer=1, wrong=2) + edge_hook_names = _edge_hook_names(N_LAYERS) + + clean_cache = cache_activation_and_gradient( + model, clean, metric, names_filter=edge_hook_names, compute_gradient=False + ) + corrupt_cache = cache_activation_and_gradient( + model, corrupt, metric, names_filter=edge_hook_names + ) + edges = enumerate_edges(model, corrupt_cache) + _assert_edges_unique(edges) + baseline_scores = _edge_effects(clean_cache, corrupt_cache, edges) + + writer = Node(kind="attn_head_out", layer=0, head=0, position=SEQ_LEN - 1) + assert any(edge_writer == writer for edge_writer, _reader in edges) + + perturbation = torch.full((D_MODEL,), 0.37, dtype=clean_cache.activations["hook_embed"].dtype) + writer_name = _writer_hook_name(writer) + perturbed_activations = dict(clean_cache.activations) + perturbed_activations[writer_name] = perturbed_activations[writer_name].clone() + perturbed_activations[writer_name][0, writer.position, writer.head] += perturbation + perturbed_clean_cache = GradientCache( + activations=perturbed_activations, + gradients=clean_cache.gradients, + metric=clean_cache.metric, + ) + + mutated_scores = _edge_effects(perturbed_clean_cache, corrupt_cache, edges) + + _assert_mutation_only_changes_the_perturbed_writers_edges( + edges, writer, perturbation, baseline_scores, mutated_scores, corrupt_cache + ) + + +def test_ablation_raises_when_a_writer_contribution_was_not_captured() -> None: + """An uncaptured writer raises rather than silently skipping the correction.""" + model = _EdgeScoringToyBridge() + _ensure_edge_hook_flags(model) + tokens = model.to_tokens("prompt") + metric = _metric_fn(answer=1, wrong=2) + edge_hook_names = _edge_hook_names(N_LAYERS) + + corrupt_cache = cache_activation_and_gradient( + model, tokens, metric, names_filter=edge_hook_names + ) + edges = enumerate_edges(model, corrupt_cache) + + with _edge_hook_flags(model), pytest.raises(ValueError, match="no contribution captured"): + _ablate_edges(model, tokens, edges, [], {}) diff --git a/transformer_lens/tools/analysis/attribution_patching.py b/transformer_lens/tools/analysis/attribution_patching.py index aac47120b3..200f40cd27 100644 --- a/transformer_lens/tools/analysis/attribution_patching.py +++ b/transformer_lens/tools/analysis/attribution_patching.py @@ -33,7 +33,17 @@ from contextlib import contextmanager from dataclasses import dataclass, field -from typing import Any, Callable, Iterator, Literal, Optional, Sequence, Union +from typing import ( + Any, + Callable, + Collection, + Iterator, + Literal, + Optional, + Sequence, + Union, + cast, +) import torch @@ -950,6 +960,170 @@ def _aggregate_edge_scores_to_writer_nodes( return totals +def _writer_contribution_slice(writer: Node, tensor: torch.Tensor) -> torch.Tensor: + """The writer's own contribution at its position, and head if it has one. + + A captured writer tensor holds every position (and, for an attention head, + every head); an edge carries only the writer's contribution at the edge's + position, so this selects that slice. The result is a ``d_model`` vector, + which broadcasts into a reader's input -- a per-head reader receives the + writer once per head. + """ + if writer.kind == "attn_head_out": + return tensor[0, writer.position, writer.head] + return tensor[0, writer.position] + + +def _excluded_writers_by_reader( + edges: Sequence[tuple[Node, Node]], + circuit: Collection[tuple[Node, Node]], +) -> dict[Node, list[Node]]: + """Group the out-of-circuit writers feeding each reader node. + + A reader node's input is rebuilt once, from the writers whose edges into it + lie outside ``circuit``. Grouping by reader node keeps that rebuild to a + single hook per node: one hook per excluded edge would apply the correction + once per edge instead of once per node. + """ + circuit_set = set(circuit) + excluded: dict[Node, list[Node]] = {} + for writer, reader in edges: + if (writer, reader) not in circuit_set: + excluded.setdefault(reader, []).append(writer) + return excluded + + +def _make_writer_capture_hook(name: str, sink: dict[str, torch.Tensor]) -> Callable[..., None]: + """Record a writer's live contribution for the reader hooks that follow it. + + A writer always precedes the readers it feeds, so by the time a reader's + fork hook fires the contributions it needs are already in ``sink``. The + tensor is detached rather than cloned: it is only read, and the reader hook + clones its own input before rewriting it. + """ + + def hook(tensor: torch.Tensor, *, hook: Any) -> None: + del hook + if isinstance(tensor, torch.Tensor): + sink[name] = tensor.detach() + return None + + return hook + + +def _make_edge_ablation_hook( + reader: Node, + excluded_writers: Sequence[Node], + live_writers: dict[str, torch.Tensor], + replacement_writers: dict[str, torch.Tensor], +) -> Callable[..., torch.Tensor]: + """Rebuild one reader's input with its out-of-circuit writers replaced. + + The residual stream is a running sum, so a reader's input is exactly the sum + of its writers' contributions. Removing an edge therefore means subtracting + that writer's live contribution and adding its replacement -- an exact + correction rather than an approximation. The reader's fork hook fires before + the layer norm, so the correction is a plain add and subtract in ``d_model`` + space with no norm scale to divide out. + + A per-head reader's input is the residual replicated across heads, so the + correction is applied to that reader's own head slot. Several reader nodes + share one hook point (one per head, one per position), and each rewrites a + disjoint slice, so no writer's contribution is subtracted twice. + + Args: + reader: The reader whose input is rebuilt. + excluded_writers: Writers whose edges into ``reader`` lie outside the + circuit. + live_writers: Contributions captured during the current forward, keyed + by writer hook name. + replacement_writers: The value standing in for each excluded writer, + keyed by writer hook name. + + Returns: + A forward hook returning the rebuilt input. + + Raises: + ValueError: if an excluded writer's contribution was not captured, which + means its hook family was not enabled for this forward. + """ + if reader.kind in ("q_input", "k_input", "v_input"): + # Node.__post_init__ guarantees a head for the per-head reader kinds. + index: tuple[int, ...] = (0, reader.position, cast(int, reader.head)) + else: + index = (0, reader.position) + + def hook(tensor: torch.Tensor, *, hook: Any) -> torch.Tensor: + del hook + rebuilt = tensor.clone() + for writer in excluded_writers: + name = _writer_hook_name(writer) + live = live_writers.get(name) + replacement = replacement_writers.get(name) + if live is None or replacement is None: + raise ValueError( + f"cannot ablate the edge from {writer} into {reader}: no " + f"contribution captured at {name!r}. Enable the writer's hook " + "family before running the ablation." + ) + rebuilt[index] = ( + rebuilt[index] + - _writer_contribution_slice(writer, live) + + _writer_contribution_slice(writer, replacement) + ) + return rebuilt + + return hook + + +def _ablate_edges( + model: Any, + tokens: torch.Tensor, + edges: Sequence[tuple[Node, Node]], + circuit: Collection[tuple[Node, Node]], + replacement_writers: dict[str, torch.Tensor], +) -> torch.Tensor: + """Run one forward with every edge outside ``circuit`` ablated. + + Captures each writer's live contribution and rewrites each reader's input at + its fork hook, so the forward sees a model in which only the circuit's edges + carry the run's own values. Forward-only: no gradient is involved, and the + caller must have the writer and reader hook families enabled (see + :func:`_edge_hook_flags`). + + Args: + model: A ``TransformerBridge`` (or compatible). + tokens: Input token ids for the ablated forward. + edges: The full edge list, as :func:`enumerate_edges` returns it. + circuit: The edges to keep. Every other edge is ablated. + replacement_writers: The value standing in for each ablated writer, + keyed by writer hook name. + + Returns: + The model's logits under the ablation, detached. The rewrite is a + forward-only intervention, so the pass runs under :func:`torch.no_grad` + and no gradient is retained. + """ + excluded_by_reader = _excluded_writers_by_reader(edges, circuit) + + live_writers: dict[str, torch.Tensor] = {} + fwd_hooks: list[tuple[str, Callable[..., Any]]] = [ + (name, _make_writer_capture_hook(name, live_writers)) + for name in _writer_hook_names(int(model.cfg.n_layers)) + ] + fwd_hooks.extend( + ( + reader.hook_name, + _make_edge_ablation_hook(reader, writers, live_writers, replacement_writers), + ) + for reader, writers in excluded_by_reader.items() + ) + + with torch.no_grad(), model.hooks(fwd_hooks=fwd_hooks): + logits = model(tokens) + return logits + + def attribution_patch( model: Any, clean: torch.Tensor, From b49bc65b36dcc678f21915420ac93b2f5ba0ef37 Mon Sep 17 00:00:00 2001 From: janmenjayap Date: Sun, 27 Sep 2026 18:24:33 +0530 Subject: [PATCH 3/6] feat(attribution_patching): faithfulness() entry point 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. --- docs/source/content/analysis_tools.md | 37 +++ .../integration/test_attribution_patching.py | 72 +++++- tests/unit/tools/test_attribution_patching.py | 189 ++++++++++++++++ transformer_lens/tools/analysis/__init__.py | 6 + .../tools/analysis/attribution_patching.py | 211 ++++++++++++++++++ 5 files changed, 511 insertions(+), 4 deletions(-) diff --git a/docs/source/content/analysis_tools.md b/docs/source/content/analysis_tools.md index fcb334a2fb..7d3381245c 100644 --- a/docs/source/content/analysis_tools.md +++ b/docs/source/content/analysis_tools.md @@ -86,6 +86,43 @@ 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. + +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, diff --git a/tests/integration/test_attribution_patching.py b/tests/integration/test_attribution_patching.py index ca3a8a2dc8..416f9033a2 100644 --- a/tests/integration/test_attribution_patching.py +++ b/tests/integration/test_attribution_patching.py @@ -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 @@ -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 @@ -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" @@ -71,3 +77,61 @@ 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 diff --git a/tests/unit/tools/test_attribution_patching.py b/tests/unit/tools/test_attribution_patching.py index e5a96fce04..387e556be7 100644 --- a/tests/unit/tools/test_attribution_patching.py +++ b/tests/unit/tools/test_attribution_patching.py @@ -22,6 +22,8 @@ from transformer_lens.tools.analysis.attribution_patching import ( AttributionResult, EdgeAttributionConfig, + FaithfulnessConfig, + FaithfulnessResult, GradientCache, Node, _ablate_edges, @@ -42,6 +44,7 @@ capture_writer_outputs, enumerate_edges, enumerate_nodes, + faithfulness, ) D_MODEL = 4 @@ -2321,3 +2324,189 @@ def test_ablation_raises_when_a_writer_contribution_was_not_captured() -> None: with _edge_hook_flags(model), pytest.raises(ValueError, match="no contribution captured"): _ablate_edges(model, tokens, edges, [], {}) + + +# --------------------------------------------------------------------------- +# Faithfulness (public API) +# --------------------------------------------------------------------------- +# +# faithfulness() ablates every edge outside a candidate circuit and reports how +# much of the clean-to-corrupt metric gap the circuit recovers. The two boundary +# circuits pin it: keeping nothing must reproduce the corrupt run, keeping +# everything must reproduce the clean run. + + +def _faithfulness_toy() -> tuple[_EdgeScoringToyBridge, torch.Tensor, torch.Tensor, Callable]: + model = _EdgeScoringToyBridge() + clean = torch.tensor([[1, 2, 3]]) + corrupt = torch.tensor([[3, 2, 1]]) + return model, clean, corrupt, _metric_fn(answer=1, wrong=2) + + +def _graph_edges(model: _EdgeScoringToyBridge, corrupt: torch.Tensor, metric: Callable) -> list: + _ensure_edge_hook_flags(model) + with _edge_hook_flags(model): + corrupt_cache = cache_activation_and_gradient( + model, corrupt, metric, names_filter=_edge_hook_names(N_LAYERS) + ) + return enumerate_edges(model, corrupt_cache) + + +def test_faithfulness_recovers_the_full_metric_for_the_full_circuit() -> None: + model, clean, corrupt, metric = _faithfulness_toy() + edges = _graph_edges(model, corrupt, metric) + + report = faithfulness(model, clean, corrupt, metric, edges) + + assert isinstance(report, FaithfulnessResult) + assert report.recovered == pytest.approx(1.0, abs=1e-6) + assert report.circuit_size == len(edges) + assert report.total_edges == len(edges) + + +def test_faithfulness_recovers_nothing_for_the_empty_circuit() -> None: + model, clean, corrupt, metric = _faithfulness_toy() + + report = faithfulness(model, clean, corrupt, metric, []) + + assert report.recovered == pytest.approx(0.0, abs=1e-6) + assert report.circuit_size == 0 + assert report.total_edges > 0 + + +def test_faithfulness_full_and_corrupt_metrics_match_direct_evaluation() -> None: + model, clean, corrupt, metric = _faithfulness_toy() + + report = faithfulness(model, clean, corrupt, metric, []) + + with torch.no_grad(): + assert report.full_metric == pytest.approx(float(metric(model(clean))), abs=1e-6) + assert report.corrupt_metric == pytest.approx(float(metric(model(corrupt))), abs=1e-6) + # The pair genuinely separates, so the recovered fraction is well defined. + assert report.full_metric != pytest.approx(report.corrupt_metric) + + +def test_faithfulness_defaults_to_corrupt_ablation() -> None: + assert FaithfulnessConfig().ablation == "corrupt" + + +def test_faithfulness_mean_ablation_differs_from_corrupt_ablation() -> None: + """The two ablation modes are genuinely different measurements.""" + model, clean, corrupt, metric = _faithfulness_toy() + edges = _graph_edges(model, corrupt, metric) + # Keep a strict subset, so some edges are actually ablated. + circuit = edges[: len(edges) // 2] + + corrupt_report = faithfulness(model, clean, corrupt, metric, circuit) + mean_report = faithfulness( + model, clean, corrupt, metric, circuit, config=FaithfulnessConfig(ablation="mean") + ) + + assert corrupt_report.recovered != pytest.approx(mean_report.recovered) + # Both bound the same gap, so the clean/corrupt endpoints agree. + assert corrupt_report.full_metric == pytest.approx(mean_report.full_metric) + assert corrupt_report.corrupt_metric == pytest.approx(mean_report.corrupt_metric) + + +def test_faithfulness_accepts_top_edges_output() -> None: + """A ranked edge list feeds straight in, scores and all.""" + model, clean, corrupt, metric = _faithfulness_toy() + ranked = attribution_patch( + model, clean, corrupt, metric, config=EdgeAttributionConfig(granularity="edge") + ).top_edges(k=5) + + report = faithfulness(model, clean, corrupt, metric, ranked) + + assert report.circuit_size == len(ranked) + assert math.isfinite(report.recovered) + + +def test_faithfulness_raises_on_an_edge_outside_the_graph() -> None: + model, clean, corrupt, metric = _faithfulness_toy() + # The graph never mixes positions, so a cross-position pair is not an edge. + bogus = (Node(kind="embed", position=1), Node(kind="mlp_in", layer=0, position=0)) + + with pytest.raises(ValueError, match="not in the graph"): + faithfulness(model, clean, corrupt, metric, [bogus]) + + +def test_faithfulness_raises_when_the_metric_gap_is_zero() -> None: + model, clean, _, metric = _faithfulness_toy() + + with pytest.raises(ValueError, match="same metric"): + faithfulness(model, clean, clean, metric, []) + + +def test_faithfulness_raises_on_token_length_mismatch() -> None: + model, _, corrupt, metric = _faithfulness_toy() + shorter = torch.tensor([[1, 2]]) + + with pytest.raises(ValueError, match="same length"): + faithfulness(model, shorter, corrupt, metric, []) + + +def test_faithfulness_raises_on_more_than_one_pair() -> None: + """Batched input is rejected rather than silently half-ablated. + + 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. + """ + model, clean, corrupt, metric = _faithfulness_toy() + batched_clean = torch.cat([clean, clean], dim=0) + batched_corrupt = torch.cat([corrupt, corrupt], dim=0) + + with pytest.raises(ValueError, match="one clean/corrupt pair at a time"): + faithfulness(model, batched_clean, batched_corrupt, metric, []) + + +def test_faithfulness_rejects_model_in_training_mode() -> None: + model, clean, corrupt, metric = _faithfulness_toy() + model.train() + + with pytest.raises(ValueError, match="evaluation mode"): + faithfulness(model, clean, corrupt, metric, []) + + +def test_faithfulness_restores_caller_hook_flags() -> None: + model, clean, corrupt, metric = _faithfulness_toy() + assert model.cfg.use_attn_result is False + assert model.cfg.use_split_qkv_input is False + assert model.cfg.use_hook_mlp_in is False + + faithfulness(model, clean, corrupt, metric, []) + + assert model.cfg.use_attn_result is False + assert model.cfg.use_split_qkv_input is False + assert model.cfg.use_hook_mlp_in is False + + +def test_faithfulness_restores_hook_flags_when_it_raises() -> None: + model, clean, corrupt, metric = _faithfulness_toy() + bogus = (Node(kind="embed", position=1), Node(kind="mlp_in", layer=0, position=0)) + + with pytest.raises(ValueError, match="not in the graph"): + faithfulness(model, clean, corrupt, metric, [bogus]) + + assert model.cfg.use_attn_result is False + assert model.cfg.use_split_qkv_input is False + assert model.cfg.use_hook_mlp_in is False + + +def test_faithfulness_recovery_is_graded_not_binary() -> None: + """A partial circuit lands strictly between the corrupt and clean metrics. + + Only edges into the readout position can move this toy's metric, since it is + position-wise and the metric reads the last position. Keeping half of those + must therefore recover part of the gap rather than all or none of it. + """ + model, clean, corrupt, metric = _faithfulness_toy() + edges = _graph_edges(model, corrupt, metric) + readout = [edge for edge in edges if edge[1].position == SEQ_LEN - 1] + assert readout + partial = readout[: len(readout) // 2] + + report = faithfulness(model, clean, corrupt, metric, partial) + + assert 0.0 < report.recovered < 1.0 + assert report.circuit_size == len(partial) diff --git a/transformer_lens/tools/analysis/__init__.py b/transformer_lens/tools/analysis/__init__.py index f2df062659..6740a4932f 100644 --- a/transformer_lens/tools/analysis/__init__.py +++ b/transformer_lens/tools/analysis/__init__.py @@ -39,8 +39,11 @@ from transformer_lens.tools.analysis.attribution_patching import ( AttributionResult, EdgeAttributionConfig, + FaithfulnessConfig, + FaithfulnessResult, Node, attribution_patch, + faithfulness, ) from transformer_lens.tools.analysis.backward_lens import ( BackwardLens, @@ -133,6 +136,8 @@ "DegenerateDirectionError", "DirectLogitAttribution", "EdgeAttributionConfig", + "FaithfulnessConfig", + "FaithfulnessResult", "FunctionSpec", "HeadAffinityPair", "HeadAffinityResult", @@ -163,6 +168,7 @@ "decompose_head", "direct_logit_attribution", "estimate_occupancy", + "faithfulness", "fit_sparse_probe", "get_act_patch_direct_path", "get_act_patch_direct_path_all_sources", diff --git a/transformer_lens/tools/analysis/attribution_patching.py b/transformer_lens/tools/analysis/attribution_patching.py index 200f40cd27..2e85ebd240 100644 --- a/transformer_lens/tools/analysis/attribution_patching.py +++ b/transformer_lens/tools/analysis/attribution_patching.py @@ -56,6 +56,7 @@ "embed", "attn_head_out", "mlp_out", "q_input", "k_input", "v_input", "mlp_in", "logits" ] Granularity = Literal["node", "edge"] +AblationMode = Literal["corrupt", "mean"] @dataclass @@ -248,6 +249,50 @@ def top_edges(self, k: int = 10) -> list[tuple[Node, Node, float]]: return [(writer, reader, score) for (writer, reader), score in ranked[:k]] +@dataclass(frozen=True) +class FaithfulnessConfig: + """Configuration for an ablate-outside faithfulness measurement. + + Attributes: + ablation: How an out-of-circuit edge's writer is replaced. + ``"corrupt"`` substitutes the corrupt run's own contribution, which + is the evaluation the pinned external reference reports and the + default here so the two compare like quantities. ``"mean"`` + substitutes the dataset mean of that writer's contribution over the + clean/corrupt batch, which is the usual choice when no single + corrupt run is meaningful. + """ + + ablation: AblationMode = "corrupt" + + +@dataclass +class FaithfulnessResult: + """How much of the clean-to-corrupt metric gap a candidate circuit recovers. + + Attributes: + recovered: Fraction of the clean-to-corrupt metric gap the circuit + recovers, ``(m_ablated - m_corrupt) / (m_clean - m_corrupt)``. A + fraction, not a percentage: ``1.0`` means ablating every edge + outside the circuit reproduces the clean metric, ``0.0`` means it + reproduces the corrupt metric. Values outside ``[0, 1]`` are + possible and meaningful -- a circuit can overshoot the clean run. + full_metric: The clean run's metric, the top of the gap. + corrupt_metric: The corrupt run's metric, the bottom of the gap. + circuit_size: Number of edges kept. + total_edges: Number of edges in the graph, so the budget is visible. + edge_class_recovered: Per-edge-class ``recovered`` values, keyed by edge + class. Empty unless the caller asks for the breakout. + """ + + recovered: float + full_metric: float + corrupt_metric: float + circuit_size: int + total_edges: int + edge_class_recovered: dict[str, float] = field(default_factory=dict) + + def _required_hook_names(n_layers: int, granularity: Granularity = "node") -> list[str]: """Hook points a sweep at ``granularity`` reads. @@ -1253,3 +1298,169 @@ def attribution_patch( node_scores = {node: total / batch for node, total in totals.items()} return AttributionResult(node_scores=node_scores) + + +def _normalize_circuit( + circuit: Sequence[tuple[Node, Node]] | Sequence[tuple[Node, Node, float]], +) -> list[tuple[Node, Node]]: + """Reduce a circuit to ``(writer, reader)`` pairs, dropping any scores. + + Accepts either the bare edge list :func:`enumerate_edges` returns or the + ranked ``(writer, reader, score)`` triples ``AttributionResult.top_edges`` + returns, so a discovered edge set feeds straight into + :func:`faithfulness` without reshaping. + """ + normalized: list[tuple[Node, Node]] = [] + for entry in circuit: + if len(entry) == 2: + writer, reader = entry + elif len(entry) == 3: + writer, reader, _score = entry + else: + raise ValueError( + f"circuit entries must be (writer, reader) or (writer, reader, score), " + f"got a {len(entry)}-tuple" + ) + normalized.append((writer, reader)) + return normalized + + +def _mean_writer_contributions( + model: Any, + clean: torch.Tensor, + corrupt: torch.Tensor, +) -> dict[str, torch.Tensor]: + """The dataset mean of each writer's contribution over the prompt pairs. + + Averages the clean and corrupt runs together, since a replacement stands in + for an out-of-circuit edge regardless of which direction the run moves. + """ + totals: dict[str, torch.Tensor] = {} + count = 0 + for tokens in (clean, corrupt): + for index in range(int(tokens.shape[0])): + captured = capture_writer_outputs(model, tokens[index : index + 1]) + for name, tensor in captured.items(): + running = totals.get(name) + totals[name] = tensor if running is None else running + tensor + count += 1 + if count == 0: + raise ValueError("faithfulness needs at least one clean/corrupt pair") + return {name: total / count for name, total in totals.items()} + + +def faithfulness( + model: Any, + clean: torch.Tensor, + corrupt: torch.Tensor, + metric_fn: MetricFn, + circuit: Sequence[tuple[Node, Node]] | Sequence[tuple[Node, Node, float]], + config: FaithfulnessConfig = FaithfulnessConfig(), +) -> FaithfulnessResult: + """Measure how much of the clean-to-corrupt metric gap a circuit recovers. + + Ablates every edge *outside* ``circuit`` and reports the metric the model + then produces. The residual stream is a running sum, so each reader's input + is rebuilt exactly -- subtracting the excluded writers' live contributions + and adding their replacements -- rather than approximated. A circuit that + recovers most of the gap is a faithful account of the behavior; one that + recovers little is not, however well it ranks. + + This is a forward-only measurement: no gradient is taken, so it is + independent of the linearization :func:`attribution_patch` uses to rank + edges. Feed a ranked edge set straight in from + ``attribution_patch(..., config=EdgeAttributionConfig(granularity="edge")).top_edges(k=...)``. + + Args: + model: A ``TransformerBridge`` (or compatible) exposing ``cfg.n_layers``, + ``hook_dict``, and ``hooks()``. + clean: Clean token ids, shape ``[batch, seq]``. + corrupt: Corrupt token ids, shape ``[batch, seq]``, paired row-by-row + with ``clean``. + metric_fn: Maps single-example logits to a scalar. + circuit: The edges to keep, as ``(writer, reader)`` pairs or the + ``(writer, reader, score)`` triples ``top_edges`` returns. Every + other edge in the graph is ablated. + config: Ablation configuration. Defaults to replacing an ablated writer + with the corrupt run's own contribution. + + Returns: + A :class:`FaithfulnessResult` with the recovered fraction, the clean and + corrupt metrics bounding it, and the circuit's size against the graph's. + + Raises: + ValueError: if ``clean``/``corrupt`` are not 2D, hold a different number + of pairs, hold more than one pair, a pair tokenizes to different + lengths, the model or a submodule is in training mode, a circuit edge + is not in the graph, or the clean and corrupt metrics are equal (so + the recovered fraction is undefined). + """ + if clean.ndim != 2 or corrupt.ndim != 2: + raise ValueError( + "faithfulness expects 2D [batch, seq] token tensors, got clean " + f"{tuple(clean.shape)} and corrupt {tuple(corrupt.shape)}" + ) + if clean.shape[0] != corrupt.shape[0]: + raise ValueError( + "clean and corrupt must hold the same number of prompt pairs, got " + f"{clean.shape[0]} and {corrupt.shape[0]}" + ) + if clean.shape[0] != 1: + raise ValueError( + "faithfulness measures one clean/corrupt pair at a time, since the " + f"metric reads a single example; got {clean.shape[0]} pairs. Loop over " + "pairs and aggregate the recovered fractions yourself." + ) + if clean.shape[1] != corrupt.shape[1]: + raise ValueError( + "each clean/corrupt pair must tokenize to the same length; got clean " + f"length {clean.shape[1]} and corrupt length {corrupt.shape[1]}. " + "Faithfulness aligns activations position-by-position." + ) + + require_eval_mode(model, operation="faithfulness()") + + kept = _normalize_circuit(circuit) + n_layers = int(model.cfg.n_layers) + hook_names = _edge_hook_names(n_layers) + + with _edge_hook_flags(model): + corrupt_cache = cache_activation_and_gradient( + model, corrupt, metric_fn, names_filter=hook_names + ) + edges = enumerate_edges(model, corrupt_cache) + graph = set(edges) + unknown = [edge for edge in kept if edge not in graph] + if unknown: + raise ValueError( + f"circuit holds {len(unknown)} edge(s) that are not in the graph, " + f"starting with {unknown[0]}; a stale or mistyped circuit would " + "otherwise ablate nothing." + ) + + if config.ablation == "mean": + replacements = _mean_writer_contributions(model, clean, corrupt) + else: + replacements = capture_writer_outputs(model, corrupt) + + # Forward-only throughout: no gradient is taken, so the clean endpoint + # is evaluated without building a graph. + with torch.no_grad(): + clean_metric = float(metric_fn(model(clean))) + corrupt_metric = float(corrupt_cache.metric) + ablated_metric = float(metric_fn(_ablate_edges(model, clean, edges, kept, replacements))) + + gap = clean_metric - corrupt_metric + if gap == 0.0: + raise ValueError( + "the clean and corrupt runs produce the same metric, so the recovered " + "fraction is undefined; choose a pair the metric separates." + ) + + return FaithfulnessResult( + recovered=(ablated_metric - corrupt_metric) / gap, + full_metric=clean_metric, + corrupt_metric=corrupt_metric, + circuit_size=len(kept), + total_edges=len(edges), + ) From e19ececf8e94d862ed035a6635f07d8fa232cfe9 Mon Sep 17 00:00:00 2001 From: janmenjayap Date: Sun, 27 Sep 2026 19:05:41 +0530 Subject: [PATCH 4/6] feat(attribution_patching): edge-class breakout in the faithfulness report 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. --- docs/source/content/analysis_tools.md | 7 ++ .../integration/test_attribution_patching.py | 41 ++++++++++ tests/unit/tools/test_attribution_patching.py | 81 +++++++++++++++++++ .../tools/analysis/attribution_patching.py | 76 ++++++++++++++--- 4 files changed, 195 insertions(+), 10 deletions(-) diff --git a/docs/source/content/analysis_tools.md b/docs/source/content/analysis_tools.md index 7d3381245c..4d5cee853f 100644 --- a/docs/source/content/analysis_tools.md +++ b/docs/source/content/analysis_tools.md @@ -121,6 +121,13 @@ 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 diff --git a/tests/integration/test_attribution_patching.py b/tests/integration/test_attribution_patching.py index 416f9033a2..80c8b9eb21 100644 --- a/tests/integration/test_attribution_patching.py +++ b/tests/integration/test_attribution_patching.py @@ -135,3 +135,44 @@ def test_faithfulness_recovers_most_of_the_metric_on_gpt2_small(gpt2_bridge) -> 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 diff --git a/tests/unit/tools/test_attribution_patching.py b/tests/unit/tools/test_attribution_patching.py index 387e556be7..1a6a4023b7 100644 --- a/tests/unit/tools/test_attribution_patching.py +++ b/tests/unit/tools/test_attribution_patching.py @@ -29,6 +29,7 @@ _ablate_edges, _assert_edges_unique, _check_required_hooks, + _edge_class, _edge_effects, _edge_hook_flags, _edge_hook_names, @@ -2510,3 +2511,83 @@ def test_faithfulness_recovery_is_graded_not_binary() -> None: assert 0.0 < report.recovered < 1.0 assert report.circuit_size == len(partial) + + +# --------------------------------------------------------------------------- +# Edge-class breakout +# --------------------------------------------------------------------------- +# +# Edges into Q and K pass through the softmax, so they are expected to be less +# faithful than edges into V, the MLP, or the terminal readout. The breakout +# measures that per class instead of hiding it in one aggregate. + + +def test_edge_class_partitions_edges_exhaustively_and_without_overlap() -> None: + model, clean, corrupt, metric = _faithfulness_toy() + edges = _graph_edges(model, corrupt, metric) + + classes = [_edge_class(edge) for edge in edges] + + # Every edge maps to exactly one class, and the classes partition the graph. + assert set(classes) == {"into_qk", "into_v", "into_mlp", "into_logits"} + assert len(classes) == len(edges) + + report = faithfulness(model, clean, corrupt, metric, []) + assert sum(report.edge_class_counts.values()) == report.total_edges + assert set(report.edge_class_counts) == set(classes) + + +def test_edge_class_maps_each_reader_kind_to_its_class() -> None: + writer = Node(kind="embed", position=0) + + assert _edge_class((writer, Node(kind="q_input", layer=0, head=0, position=0))) == "into_qk" + assert _edge_class((writer, Node(kind="k_input", layer=0, head=0, position=0))) == "into_qk" + assert _edge_class((writer, Node(kind="v_input", layer=0, head=0, position=0))) == "into_v" + assert _edge_class((writer, Node(kind="mlp_in", layer=0, position=0))) == "into_mlp" + assert _edge_class((writer, Node(kind="logits", layer=0, position=0))) == "into_logits" + + # A writer kind is not a reader, so it has no class. + with pytest.raises(ValueError, match="no edge class"): + _edge_class((writer, Node(kind="mlp_out", layer=0, position=0))) + + +def test_edge_class_recovered_reconciles_with_the_aggregate() -> None: + model, clean, corrupt, metric = _faithfulness_toy() + edges = _graph_edges(model, corrupt, metric) + + report = faithfulness(model, clean, corrupt, metric, edges) + + assert set(report.edge_class_recovered) == { + "into_qk", + "into_v", + "into_mlp", + "into_logits", + } + for edge_class, recovered in report.edge_class_recovered.items(): + assert math.isfinite(recovered), f"non-finite recovery for {edge_class}" + assert report.edge_class_counts[edge_class] > 0 + # Keeping every edge is the aggregate, so the full-circuit report recovers + # the whole gap. + assert report.recovered == pytest.approx(1.0, abs=1e-6) + + +def test_edge_class_breakout_is_leave_one_out() -> None: + """Each class's entry is the recovery with that class's edges removed. + + Keeping only a class would report 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. On this toy the classes are separable, so + dropping one must move the number away from the full-circuit ``1.0``. + """ + model, clean, corrupt, metric = _faithfulness_toy() + edges = _graph_edges(model, corrupt, metric) + + report = faithfulness(model, clean, corrupt, metric, edges) + + assert report.edge_class_recovered + assert all(math.isfinite(value) for value in report.edge_class_recovered.values()) + assert sum(report.edge_class_counts.values()) == report.total_edges + # Leave-one-out is not the aggregate: removing a class changes the number. + assert any( + value != pytest.approx(report.recovered) for value in report.edge_class_recovered.values() + ) diff --git a/transformer_lens/tools/analysis/attribution_patching.py b/transformer_lens/tools/analysis/attribution_patching.py index 2e85ebd240..2a694483b1 100644 --- a/transformer_lens/tools/analysis/attribution_patching.py +++ b/transformer_lens/tools/analysis/attribution_patching.py @@ -57,6 +57,7 @@ ] Granularity = Literal["node", "edge"] AblationMode = Literal["corrupt", "mean"] +EdgeClass = Literal["into_qk", "into_v", "into_mlp", "into_logits"] @dataclass @@ -281,8 +282,17 @@ class FaithfulnessResult: corrupt_metric: The corrupt run's metric, the bottom of the gap. circuit_size: Number of edges kept. total_edges: Number of edges in the graph, so the budget is visible. - edge_class_recovered: Per-edge-class ``recovered`` values, keyed by edge - class. Empty unless the caller asks for the breakout. + edge_class_recovered: Recovered fraction when every edge *outside* this + class is kept, keyed by class. A class whose removal collapses + recovery is load-bearing; one whose removal leaves recovery near + ``1.0`` is not. Keeping only the class instead would report roughly + ``0.0`` for every class, since no single class alone reconstructs + the behavior, so the leave-one-out form is the informative one. + Edges into Q and K pass through the softmax, so ``"into_qk"`` is + expected to be the least faithful class; edges into V, the MLP, and + the terminal logits readout are linear. + edge_class_counts: Number of edges in each class, so a class with few + edges is not over-read. """ recovered: float @@ -290,7 +300,8 @@ class FaithfulnessResult: corrupt_metric: float circuit_size: int total_edges: int - edge_class_recovered: dict[str, float] = field(default_factory=dict) + edge_class_recovered: dict[EdgeClass, float] = field(default_factory=dict) + edge_class_counts: dict[EdgeClass, int] = field(default_factory=dict) def _required_hook_names(n_layers: int, granularity: Granularity = "node") -> list[str]: @@ -1325,6 +1336,26 @@ def _normalize_circuit( return normalized +def _edge_class(edge: tuple[Node, Node]) -> EdgeClass: + """The class an edge belongs to, keyed by what its reader consumes. + + Q and K share a class because both feed the attention score, so both pass + through the same softmax nonlinearity; V, the MLP entry, and the terminal + logits readout are each linear in the residual they read. The split exists + to make that asymmetry measurable rather than hidden in one aggregate. + """ + reader = edge[1] + if reader.kind in ("q_input", "k_input"): + return "into_qk" + if reader.kind == "v_input": + return "into_v" + if reader.kind == "mlp_in": + return "into_mlp" + if reader.kind == "logits": + return "into_logits" + raise ValueError(f"{reader.kind} is a writer kind and has no edge class") + + def _mean_writer_contributions( model: Any, clean: torch.Tensor, @@ -1386,7 +1417,11 @@ def faithfulness( Returns: A :class:`FaithfulnessResult` with the recovered fraction, the clean and - corrupt metrics bounding it, and the circuit's size against the graph's. + corrupt metrics bounding it, the circuit's size against the graph's, and + a per-edge-class breakout. The breakout measures each class + leave-one-out, so it costs one extra ablation per class; edges into Q and + K are expected to be the least faithful, since they pass through the + softmax. Raises: ValueError: if ``clean``/``corrupt`` are not 2D, hold a different number @@ -1448,14 +1483,33 @@ def faithfulness( with torch.no_grad(): clean_metric = float(metric_fn(model(clean))) corrupt_metric = float(corrupt_cache.metric) + gap = clean_metric - corrupt_metric + if gap == 0.0: + raise ValueError( + "the clean and corrupt runs produce the same metric, so the recovered " + "fraction is undefined; choose a pair the metric separates." + ) ablated_metric = float(metric_fn(_ablate_edges(model, clean, edges, kept, replacements))) - gap = clean_metric - corrupt_metric - if gap == 0.0: - raise ValueError( - "the clean and corrupt runs produce the same metric, so the recovered " - "fraction is undefined; choose a pair the metric separates." - ) + # Break the report out by edge class, so the attention nonlinearity's + # cost is measured rather than hidden in the aggregate. Each class is + # measured leave-one-out: keep every edge outside it, so the value says + # how much of the gap survives without that class. One extra ablation + # per class. + class_counts: dict[EdgeClass, int] = {} + class_edges: dict[EdgeClass, list[tuple[Node, Node]]] = {} + for edge in edges: + edge_class = _edge_class(edge) + class_counts[edge_class] = class_counts.get(edge_class, 0) + 1 + class_edges.setdefault(edge_class, []).append(edge) + + class_recovered: dict[EdgeClass, float] = {} + for edge_class in class_edges: + outside = [edge for edge in edges if _edge_class(edge) != edge_class] + class_metric = float( + metric_fn(_ablate_edges(model, clean, edges, outside, replacements)) + ) + class_recovered[edge_class] = (class_metric - corrupt_metric) / gap return FaithfulnessResult( recovered=(ablated_metric - corrupt_metric) / gap, @@ -1463,4 +1517,6 @@ def faithfulness( corrupt_metric=corrupt_metric, circuit_size=len(kept), total_edges=len(edges), + edge_class_recovered=class_recovered, + edge_class_counts=class_counts, ) From a77d3a627f14eda480121e7a19f1903573791901 Mon Sep 17 00:00:00 2001 From: janmenjayap Date: Sun, 27 Sep 2026 19:16:17 +0530 Subject: [PATCH 5/6] test(attribution_patching): random-edge-set baseline honesty 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. --- tests/unit/tools/test_attribution_patching.py | 125 +++++++++++++++++- 1 file changed, 124 insertions(+), 1 deletion(-) diff --git a/tests/unit/tools/test_attribution_patching.py b/tests/unit/tools/test_attribution_patching.py index 1a6a4023b7..f244503d2d 100644 --- a/tests/unit/tools/test_attribution_patching.py +++ b/tests/unit/tools/test_attribution_patching.py @@ -1385,7 +1385,7 @@ def forward(self, residual: torch.Tensor) -> torch.Tensor: ) w_o_per_head = self.w_o.weight.reshape(d_model, self.n_heads, self.d_head).permute(1, 2, 0) - per_head_out_raw = torch.einsum("bshd,hdm->bshm", z, w_o_per_head) + per_head_out_raw = self._attention_output(torch.einsum("bshd,hdm->bshm", z, w_o_per_head)) if self.cfg.use_attn_result: per_head_out = self.hook_result(per_head_out_raw) else: @@ -1396,6 +1396,28 @@ def forward(self, residual: torch.Tensor) -> torch.Tensor: mlp_out = self.hook_mlp_out(self.w_mlp(mlp_in)) return self.hook_resid_post(residual + mlp_out) + def _attention_output(self, per_head_out: torch.Tensor) -> torch.Tensor: + """Transform the per-head attention output before it joins the residual. + + Identity here. A subclass can make the block nonlinear, which is what + stops a circuit from reproducing the metric merely by containing every + edge into the logits reader. + """ + return per_head_out + + +class _NonlinearEdgeScoringBlock(_EdgeScoringBlock): + """``_EdgeScoringBlock`` with a GELU on the attention output path. + + The nonlinearity breaks the linear relation between a writer's contribution + and the metric, so a circuit cannot recover the metric just by containing + every edge into the logits reader. That is what makes a random edge set of + the same size a meaningful baseline rather than a coin flip. + """ + + def _attention_output(self, per_head_out: torch.Tensor) -> torch.Tensor: + return torch.nn.functional.gelu(per_head_out) + class _EdgeScoringToyBridge(_LinearToyBridge): """A tiny ``TransformerBridge`` whose edge-granularity hooks all sit on the data path. @@ -1475,6 +1497,30 @@ def set_use_hook_mlp_in(self, use_hook_mlp_in: bool) -> None: self.cfg.use_hook_mlp_in = use_hook_mlp_in +class _NonlinearEdgeScoringToyBridge(_EdgeScoringToyBridge): + """``_EdgeScoringToyBridge`` whose attention output passes through a GELU. + + Reuses the parent's hook graph and ``hooks()`` plumbing but swaps in + :class:`_NonlinearEdgeScoringBlock`, so a writer's contribution reaches the + metric nonlinearly. On a linear toy any circuit holding every edge into the + logits reader recovers the metric outright, which makes a random baseline + meaningless; the nonlinearity is what gives the comparison teeth. + """ + + def __init__(self, *, dtype: torch.dtype = torch.float32) -> None: + super().__init__(dtype=dtype) + torch.manual_seed(1) + self.blocks = nn.ModuleList( + [ + _NonlinearEdgeScoringBlock(D_MODEL, N_HEADS, D_HEAD, layer, dtype, self.cfg) + for layer in range(N_LAYERS) + ] + ) + # The parent's eval() ran before these blocks existed, so they would + # otherwise start in training mode and trip require_eval_mode. + self.eval() + + def test_attribution_patch_edge_granularity_scores_every_edge_with_finite_values() -> None: model = _EdgeScoringToyBridge() clean = torch.tensor([[1, 2, 3]]) @@ -2591,3 +2637,80 @@ def test_edge_class_breakout_is_leave_one_out() -> None: assert any( value != pytest.approx(report.recovered) for value in report.edge_class_recovered.values() ) + + +# --------------------------------------------------------------------------- +# Random-edge-set baseline +# --------------------------------------------------------------------------- +# +# A ranked circuit is only meaningful if it beats an arbitrary edge set of the +# same size. The toy is nonlinear here on purpose: on a linear model any circuit +# holding every edge into the logits reader recovers the metric outright, so the +# baseline could tie by accident. + + +def _random_edge_sets( + edges: list[tuple[Node, Node]], size: int, draws: int, seed: int +) -> list[list[tuple[Node, Node]]]: + """``draws`` distinct random edge sets of ``size`` edges, from a fixed seed.""" + generator = torch.Generator().manual_seed(seed) + return [ + [edges[index] for index in torch.randperm(len(edges), generator=generator)[:size].tolist()] + for _ in range(draws) + ] + + +def test_random_edge_set_recovers_markedly_less_than_the_attribution_circuit() -> None: + """The ranked circuit beats random edge sets of the same size by a wide margin. + + This is the honesty guard: a ranking that does no better than chance would + still produce a plausible-looking ``recovered`` number, so the comparison + against same-size random sets is what makes the number mean anything. + """ + model = _NonlinearEdgeScoringToyBridge() + clean = torch.tensor([[1, 2, 3]]) + corrupt = torch.tensor([[3, 2, 1]]) + metric = _metric_fn(answer=1, wrong=2) + + result = attribution_patch( + model, clean, corrupt, metric, config=EdgeAttributionConfig(granularity="edge") + ) + edges = list(result.edge_scores) + budget = 8 + ranked = result.top_edges(k=budget) + ranked_report = faithfulness(model, clean, corrupt, metric, ranked) + + random_reports = [ + faithfulness(model, clean, corrupt, metric, circuit) + for circuit in _random_edge_sets(edges, budget, draws=3, seed=0) + ] + random_recovered = [report.recovered for report in random_reports] + best_random = max(random_recovered) + + # The baseline is not degenerate: if every random set recovered the whole + # gap, the toy would be too easy to discriminate and the comparison would be + # vacuous. + assert all(value < 1.0 for value in random_recovered), random_recovered + assert ranked_report.recovered > best_random + 0.2, ( + f"ranked recovered {ranked_report.recovered:.4f} vs random " + f"{[round(value, 4) for value in random_recovered]}" + ) + + +def test_random_edge_set_baseline_is_reproducible_under_a_fixed_seed() -> None: + model = _NonlinearEdgeScoringToyBridge() + clean = torch.tensor([[1, 2, 3]]) + corrupt = torch.tensor([[3, 2, 1]]) + metric = _metric_fn(answer=1, wrong=2) + + result = attribution_patch( + model, clean, corrupt, metric, config=EdgeAttributionConfig(granularity="edge") + ) + edges = list(result.edge_scores) + + first = _random_edge_sets(edges, 8, draws=3, seed=0) + second = _random_edge_sets(edges, 8, draws=3, seed=0) + + assert first == second + # Different seeds give different draws, so the seed is actually used. + assert first != _random_edge_sets(edges, 8, draws=3, seed=1) From b1532ca1724df8fde93c17ebb6437ab22462252a Mon Sep 17 00:00:00 2001 From: janmenjayap Date: Sun, 27 Sep 2026 19:24:24 +0530 Subject: [PATCH 6/6] docs(attribution_patching): record the faithfulness contract in the module 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. --- .../tools/analysis/attribution_patching.py | 16 +++++++++++++--- 1 file changed, 13 insertions(+), 3 deletions(-) diff --git a/transformer_lens/tools/analysis/attribution_patching.py b/transformer_lens/tools/analysis/attribution_patching.py index 2a694483b1..f5a9aad08d 100644 --- a/transformer_lens/tools/analysis/attribution_patching.py +++ b/transformer_lens/tools/analysis/attribution_patching.py @@ -23,10 +23,20 @@ should filter to the hook families their analysis actually reads. Scope: this build ships node and edge granularity with plain attribution -(``ig_steps=1``). The integrated-gradient path (EAP-IG, ``ig_steps>1``) and -ablate-outside faithfulness are not implemented yet; their API is declared here, -and ``ig_steps>1`` raises :class:`NotImplementedError`, so downstream code can pin +(``ig_steps=1``), plus ablate-outside faithfulness. The integrated-gradient path +(EAP-IG, ``ig_steps>1``) is not implemented yet; its API is declared here, and +``ig_steps>1`` raises :class:`NotImplementedError`, so downstream code can pin against a stable surface now. + +Faithfulness intervenes rather than estimating: it ablates every edge outside a +candidate circuit and reports how much of the clean-to-corrupt metric gap the +circuit recovers. The residual stream is a running sum, so a reader's input is +exactly the sum of its writers' contributions, and rebuilding that input from +the included writers plus replacements for the excluded ones is an exact +correction rather than an approximation. The reader's fork hooks +(``attn.hook_q_input`` / ``hook_k_input`` / ``hook_v_input``, ``hook_mlp_in``) +fire before the layer norm, so the correction is a plain add and subtract in +``d_model`` space with no norm scale to divide out. """ from __future__ import annotations