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
1 change: 1 addition & 0 deletions .github/workflows/checks.yml
Original file line number Diff line number Diff line change
Expand Up @@ -443,6 +443,7 @@ jobs:
- "Patchscopes_Generation_Demo"
- "Santa_Coder"
# - "stable_lm"
- "SVD_Circuits_Demo"
- "T5"
requires_hf_token: [false]
include:
Expand Down
1,009 changes: 1,009 additions & 0 deletions demos/SVD_Circuits_Demo.ipynb

Large diffs are not rendered by default.

1 change: 1 addition & 0 deletions docs/make_docs.py
Original file line number Diff line number Diff line change
Expand Up @@ -43,6 +43,7 @@ def copy_demos(_app: Optional[Any] = None):
"Jacobian_Lens_Coordinate_Patch_Benchmark_Demo.ipynb",
"Jacobian_Lens_Decomposition_Demo.ipynb",
"Main_Demo.ipynb",
"SVD_Circuits_Demo.ipynb",
]

if copy_to_dir.exists():
Expand Down
6 changes: 5 additions & 1 deletion docs/source/content/analysis_tools.md
Original file line number Diff line number Diff line change
Expand Up @@ -139,8 +139,12 @@ For SVD head decomposition, inspect the rank report before assigning meaning to
individual direction. Near-equal singular values define a subspace whose basis can
rotate; numerically null directions are also unsuitable for individual attribution.
The weight decomposition alone is not a causal validation of a proposed subfunction.
See [SVD Circuits](svd_circuits.md) for the degeneracy guard, the causal gate, and a
worked example, and the [SVD Circuits demo](../generated/demos/SVD_Circuits_Demo.html)
for a runnable walkthrough.

API: {func}`~transformer_lens.tools.analysis.svd_circuits.decompose_head`.
API: {func}`~transformer_lens.tools.analysis.svd_circuits.decompose_head`,
{func}`~transformer_lens.tools.analysis.svd_circuits.patch_along_directions`.

## Try a geometry question without downloading a model

Expand Down
147 changes: 147 additions & 0 deletions docs/source/content/svd_circuits.md
Original file line number Diff line number Diff line change
@@ -0,0 +1,147 @@
# SVD Circuits

SVD Circuits decomposes a single attention head's QK and OV weight maps into orthogonal
singular directions. Vocab readouts and activation projections describe those directions;
patching measures prompt-specific changes in a caller-selected output metric.

Mathematically distinct directions need not be semantically or causally distinct
subfunctions. A weight-space decomposition or a plausible token projection does not
establish a direction's role. Interventions provide additional measurements, but the
reported gate comparison is not by itself a validation of a named mechanism.

## Definition

For a head at layer $\ell$ with head index $h$, the two maps are

$$
\mathrm{QK} = W_Q W_K^\top, \qquad \mathrm{OV} = W_V W_O,
$$

each of shape $d_{\text{model}} \times d_{\text{model}}$ but of rank at most
$d_{\text{head}}$. Their SVDs are

$$
W_V W_O = U \Sigma V^\top, \qquad \operatorname{rank} \le d_{\text{head}}.
$$

The convention matters, because both factors are $d_{\text{model}}$-wide and a swap
raises no shape error:

- For OV, the columns of $V$ are the residual-stream **output/write** directions, the
ones projected through $W_U$ for a vocab readout. The columns of $U$ span the
value-computation **input** space.
- For QK, the columns of $V$ are the **source/key-read** directions and the columns of
$U$ the **destination/query-read** directions. QK produces no write direction, so it
has no vocab readout.

`FactoredMatrix` computes both SVDs without materialising the $d_{\text{model}}^2$
product.

## The degeneracy guard

Singular directions are unique only when the singular values are distinct. Equal or
near-equal consecutive singular values leave the corresponding subspace defined only up
to an arbitrary rotation, so any statement of the form "direction 3 is the surname
subfunction" is a statement about a basis choice rather than about the model.

`decompose_head` returns a rank report marking degenerate and numerically null
directions. Null directions are arbitrary vectors from the map's null space and are
also unsuitable for individual attribution. The consumers have different guards:

- `logit_signature` requires an isolated, non-null direction and refuses individual
directions inside a degenerate block.
- `patch_along_directions` refuses selections that split a degenerate block. Complete
blocks can be retained or removed as subspaces; empty or full-span retained sets
require an explicit `threshold` because their random controls coincide with them.
- `vocab_readout` and `project_activations` return raw numerical projections, including
columns inside degenerate blocks. They neither refuse these columns nor replace
them with block summaries. Consult the rank report and exclude both `is_degenerate`
and `is_null` before interpreting individual directions.

## The intervention comparison

`patch_along_directions` reconstructs the head's output onto a chosen singular subspace
and reports the resulting change in a caller-supplied metric:

- `delta_metric`: patched metric minus original metric.
- `baseline_delta_metric`: the mean magnitude over random same-width subspaces drawn
inside the head's own OV span. Drawing the control in-span rather than from the full
residual stream makes it the effect of an arbitrary subspace of *this head's* output,
which is the comparison the gate needs.
- `gated`: whether the metric change passed the mode-specific threshold comparison.

The comparison depends on the mode the caller expressed, because `keep=S` and
`ablate=complement(S)` resolve to the same retained set:

| Mode | Meaning | `gated` is |
|---|---|---|
| `keep=S` | retain only $S$ | `abs(delta_metric) < threshold` |
| `ablate=S` | zero $S$, retain the rest | `abs(delta_metric) > threshold` |

With `keep`, a passing comparison means retaining the selected subspace changes the
metric less than the threshold. With `ablate`, it means removing the selected subspace
changes the metric more than the threshold. These describe effects on the chosen metric,
not reconstruction of the head's entire behaviour or identification of its function.

The threshold defaults to `baseline_delta_metric`, a sampled mean magnitude rather than
a confidence bound or p-value. It does not establish statistical separation from
arbitrary directions. The sampled mean can vary with the seed and draw count, so verdicts
near it can change. Pass an explicit `rng` for repeatability and increase `n_baseline` to
sample the mean more thoroughly; neither makes a verdict scientifically conclusive.

## Compatibility mode

The vocab readout projects through the final LayerNorm folded into $W_U$, which requires
`enable_compatibility_mode()` on an adapter that supports folding. A decomposition
records the folded-LayerNorm state it was taken under, and the readout and patch
consumers refuse a decomposition whose state no longer matches the model. See
[Compatibility Mode](compatibility_mode.md).

## Worked example

```python
import torch

from transformer_lens.model_bridge import TransformerBridge
from transformer_lens.tools.analysis.svd_circuits import (
decompose_head,
patch_along_directions,
vocab_readout,
)

model = TransformerBridge.boot_transformers("gpt2", dtype=torch.float32, device="cpu")
model.enable_compatibility_mode()
model.eval()

prompt = "When Mary and John went to the store, John gave a drink to"
ov = decompose_head(model, layer=9, head=9, which=("OV",)).OV

readout = vocab_readout(model, ov, k=10)

mary, john = model.to_single_token(" Mary"), model.to_single_token(" John")
metric = lambda logits: float(logits[0, -1, mary] - logits[0, -1, john])

result = patch_along_directions(
model, ov, prompt, metric, keep=[0], rng=torch.Generator().manual_seed(0)
)
print(result.delta_metric, result.baseline_delta_metric, result.gated)
```

## What this does not establish

- A passing gate reports a metric comparison for one intervention on one prompt. Vocab
readouts and activation coefficients are descriptive projections, not semantic labels.
- Results are single-head and prompt-specific. Nothing here assembles a multi-head circuit.
- The worked example and demo do not establish named subfunctions, statistical separation
from arbitrary directions, or replication of the paper's causal subfunction taxonomy.
- Output-metric changes include downstream responses. A small final-logit effect does not
imply a small direct write contribution, nor identify which downstream components alter
the effect.
- Verdicts near the mean control magnitude can change with the seed or draw count. Fixing
a seed makes the comparison repeatable, not scientifically conclusive.

## Links

- [SVD Circuits demo](../generated/demos/SVD_Circuits_Demo.ipynb)
- Areeb Ahmad, Abhinav Joshi, Ashutosh Modi, "Beyond Components: Singular Vector-Based
Interpretability of Transformer Circuits", [arXiv 2511.20273](https://arxiv.org/abs/2511.20273)
2 changes: 2 additions & 0 deletions docs/source/index.md
Original file line number Diff line number Diff line change
Expand Up @@ -67,6 +67,8 @@ generated/demos/Jacobian_Lens_Coordinate_Patch_Benchmark_Demo
content/backward_lens
content/debugging_numerical_divergence
content/sparse_probing
content/svd_circuits
generated/demos/SVD_Circuits_Demo
generated/demos/Main_Demo
generated/demos/Exploratory_Analysis_Demo
content/special_cases
Expand Down
1 change: 1 addition & 0 deletions makefile
Original file line number Diff line number Diff line change
Expand Up @@ -58,6 +58,7 @@ notebook-test:
$(RUN) pytest --nbval-sanitize-with demos/doc_sanitize.cfg demos/Qwen.ipynb $(RERUN_ARGS)
$(RUN) pytest --nbval-sanitize-with demos/doc_sanitize.cfg demos/Santa_Coder.ipynb $(RERUN_ARGS)
$(RUN) pytest --nbval-sanitize-with demos/doc_sanitize.cfg demos/stable_lm.ipynb $(RERUN_ARGS)
$(RUN) pytest --nbval-sanitize-with demos/doc_sanitize.cfg demos/SVD_Circuits_Demo.ipynb $(RERUN_ARGS)
$(RUN) pytest --nbval-sanitize-with demos/doc_sanitize.cfg demos/SVD_Interpreter_Demo.ipynb $(RERUN_ARGS)
$(RUN) pytest --nbval-sanitize-with demos/doc_sanitize.cfg demos/Tracr_to_Transformer_Lens_Demo.ipynb $(RERUN_ARGS)

Expand Down
190 changes: 190 additions & 0 deletions tests/integration/test_svd_circuits_oracle_parity.py
Original file line number Diff line number Diff line change
@@ -0,0 +1,190 @@
"""Slow qualitative OV sweep on GPT-2-small layer 9 head 9.

Reports prompt-specific gate counts without requiring a minimum or claiming a causal
subfunction split. There is no pinned external numerical oracle: the Beyond Components
paper publishes no numeric table for this head. The repository's ``SVDInterpreter``
cross-check is covered separately in ``test_svd_circuits.py``.

Keep-mode gates compare the metric change with a sampled mean random in-span control
magnitude, not a statistical significance threshold. Counts may vary with seed and draw
count, including zero. The checks cover eligibility, finite results, gate polarity, and
fixed-seed reproducibility. Each sweep costs ``k * (n_baseline + 2)`` forward passes,
so the model tests are marked ``slow`` and excluded from the default tiers.
"""

from __future__ import annotations

import math
from dataclasses import dataclass
from typing import Callable, List, Sequence, Tuple

import pytest
import torch

from transformer_lens.model_bridge import TransformerBridge
from transformer_lens.tools.analysis.svd_circuits import (
HeadSVD,
decompose_head,
patch_along_directions,
)

CLEAN_PROMPT = "When Mary and John went to the store, John gave a drink to"

# Wang et al. (2022), arXiv 2211.00593, identifies L9H9 as a name mover in the IOI circuit.
LAYER, HEAD = 9, 9

TOP_K_DIRECTIONS = 8
BASELINE_DRAWS = 32
BASELINE_SEED = 0


@pytest.fixture(scope="module")
def gpt2_bridge():
model = TransformerBridge.boot_transformers("gpt2", device="cpu", dtype=torch.float32)
model.enable_compatibility_mode()
# Forward-pass tools require evaluation mode.
model.eval()
return model


def _logit_diff_metric(model) -> Callable[[torch.Tensor], float]:
mary_token = model.to_single_token(" Mary")
john_token = model.to_single_token(" John")

def metric(logits: torch.Tensor) -> float:
return float(logits[0, -1, mary_token] - logits[0, -1, john_token])

return metric


@dataclass(frozen=True)
class _GateRow:
"""One swept direction's gate verdict."""

idx: int
sigma: float
delta_metric: float
baseline_delta_metric: float
gated: bool


def _eligible_directions(head_svd: HeadSVD, k: int) -> List[int]:
"""The first ``k`` directions that are attributable on their own.

Fewer than ``k`` may qualify: the degeneracy guard is a feature, so a short sweep
is used as-is rather than treated as an error.
"""
eligible = [
row.idx for row in head_svd.rank_report if not row.is_degenerate and not row.is_null
]
return eligible[:k]


def _gated_directions(
model,
head_svd: HeadSVD,
prompt: torch.Tensor,
metric: Callable[[torch.Tensor], float],
*,
k: int,
n_baseline: int,
seed: int,
) -> List[_GateRow]:
"""Report keep-mode changes against mean same-width random in-span control magnitudes."""
rows: List[_GateRow] = []
for idx in _eligible_directions(head_svd, k):
result = patch_along_directions(
model,
head_svd,
prompt,
metric,
keep=[idx],
rng=torch.Generator().manual_seed(seed),
n_baseline=n_baseline,
)
rows.append(
_GateRow(
idx=idx,
sigma=head_svd.rank_report[idx].sigma,
delta_metric=result.delta_metric,
baseline_delta_metric=result.baseline_delta_metric,
gated=result.gated,
)
)
return rows


def _format_rows(rows: Sequence[_GateRow], *, seed: int, n_baseline: int) -> str:
lines = [
f"seed={seed}, draws={n_baseline}, swept={len(rows)}, "
f"gated={sum(row.gated for row in rows)} (descriptive count)",
f"{'dir':>4} {'sigma':>10} {'delta':>12} {'baseline':>12} {'gated':>6}",
]
for row in rows:
lines.append(
f"{row.idx:>4} {row.sigma:>10.4f} {row.delta_metric:>12.6f} "
f"{row.baseline_delta_metric:>12.6f} {str(row.gated):>6}"
)
return "\n".join(lines)


def _sweep(
gpt2_bridge, *, seed: int = BASELINE_SEED, n_baseline: int = BASELINE_DRAWS
) -> Tuple[_GateRow, ...]:
decomposition = decompose_head(gpt2_bridge, layer=LAYER, head=HEAD, which=("OV",))
ov = decomposition.OV
assert ov is not None
return tuple(
_gated_directions(
gpt2_bridge,
ov,
gpt2_bridge.to_tokens(CLEAN_PROMPT),
_logit_diff_metric(gpt2_bridge),
k=TOP_K_DIRECTIONS,
n_baseline=n_baseline,
seed=seed,
)
)


@pytest.fixture(scope="module")
def default_sweep(gpt2_bridge) -> Tuple[_GateRow, ...]:
"""Share immutable default rows without replacing the fresh reproducibility sweep."""
return _sweep(gpt2_bridge)


@pytest.mark.slow
@pytest.mark.parametrize("seed", [0, 1, 2])
@pytest.mark.parametrize("n_baseline", [16, BASELINE_DRAWS])
def test_qualitative_sweep_has_finite_results_and_consistent_gates(
gpt2_bridge, default_sweep, seed: int, n_baseline: int
) -> None:
rows = (
default_sweep
if (seed, n_baseline) == (BASELINE_SEED, BASELINE_DRAWS)
else _sweep(gpt2_bridge, seed=seed, n_baseline=n_baseline)
)
report = _format_rows(rows, seed=seed, n_baseline=n_baseline)
print(report)

ov = decompose_head(gpt2_bridge, layer=LAYER, head=HEAD, which=("OV",)).OV
assert ov is not None
expected_ids = [row.idx for row in ov.rank_report if not row.is_degenerate and not row.is_null][
:TOP_K_DIRECTIONS
]
actual_ids = [row.idx for row in rows]
assert rows, "no attributable OV directions to sweep"
assert actual_ids == expected_ids, report
assert len(set(actual_ids)) == len(actual_ids), report
for row in rows:
assert math.isfinite(row.delta_metric), report
assert math.isfinite(row.baseline_delta_metric), report
assert row.baseline_delta_metric >= 0, report
assert row.gated == (abs(row.delta_metric) < row.baseline_delta_metric), report


@pytest.mark.slow
def test_sweep_is_reproducible_under_a_fixed_seed(gpt2_bridge, default_sweep) -> None:
second = _sweep(gpt2_bridge)

assert default_sweep == second
1 change: 1 addition & 0 deletions tests/unit/test_make_docs.py
Original file line number Diff line number Diff line change
Expand Up @@ -42,4 +42,5 @@ def test_copy_demos_creates_generated_dir_when_absent(tmp_path, monkeypatch):
"Jacobian_Lens_Coordinate_Patch_Benchmark_Demo.ipynb",
"Jacobian_Lens_Decomposition_Demo.ipynb",
"Main_Demo.ipynb",
"SVD_Circuits_Demo.ipynb",
]
Loading
Loading