transformer_lens.tools.analysis.svd_circuits module

Singular-vector decomposition of a single attention head’s QK and OV maps.

An attention head is characterized by two low-rank linear maps: the query-key map W_Q W_K^T that scores source positions, and the output-value map W_V W_O that writes the attended value back into the residual stream. This tool takes the singular value decomposition of each map for one head and exposes the singular values (how much each direction matters) together with the left and right singular vectors that span the map’s input and output spaces.

The decomposition is weight-space only: it reads W_Q/W_K/W_V/W_O and needs no forward pass, no activation cache, and no compatibility mode. Each map is kept factored through FactoredMatrix, so the d_model x d_model product is never materialized and the returned rank is bounded by d_head.

For a factored map A @ B (A: [ldim, mdim], B: [mdim, rdim]), the SVD’s U columns live in A’s input space (ldim) and V columns live in B’s output space (rdim): feeding x = U[:, i] through the map gives x @ (A @ B) == S[i] * V[:, i], never the reverse. For OV (A = W_V_h, B = W_O_h), U’s columns are therefore the value-computation input directions this head reads from the residual stream, and V’s columns are the output directions it writes back into the residual stream - the ones to project through W_U for a vocab or logit readout. For QK (A = W_Q_h, B = W_K_h.transpose(-1, -2)), both U (destination/query-read) and V (source/key-read) are read directions; QK only ever produces a scalar attention score, so neither is a write direction. The historical .Vh alias returns the same tensor as .V and is never used here.

Adjacent singular values closer than a relative gap eps leave their singular directions defined only up to a rotation, so the result carries a per-direction degeneracy report. Directions are grouped into contiguous blocks that end only at a gap of at least eps; callers attribute such a block as a subspace instead of trusting a single, rotation-dependent direction inside it.

A singular value near zero relative to the top of the spectrum is null rather than near-equal: its singular vector is an arbitrary null-space direction, not a rotation of a comparable neighbour. Null directions are flagged under their own tolerance, null_rtol, keyed to the spectrum’s top value the way torch.linalg.matrix_rank() keys its default tolerance.

Example:

from transformer_lens.model_bridge import TransformerBridge
from transformer_lens.tools.analysis.svd_circuits import decompose_head

model = TransformerBridge.boot_transformers("gpt2", device="cpu")
decomposition = decompose_head(model, layer=0, head=0)
ov = decomposition.OV
for row in ov.rank_report:
    print(f"direction {row.idx}: sigma={row.sigma:.3f} ratio={row.sigma_ratio:.3f}")
class transformer_lens.tools.analysis.svd_circuits.ActivationProjection(head_svd: HeadSVD, coefficients: Float[Tensor, 'pos rank'], str_tokens: List[str])

Bases: object

Per-position coefficients of a head’s actual output in its OV output basis.

head_svd

The OV decomposition this was projected against.

Type:

transformer_lens.tools.analysis.svd_circuits.HeadSVD

coefficients

[pos, rank]; coefficients[:, i] is the signed amount of singular direction i (head_svd.V[:, i]) present in the head’s actual output at each position. Summing coefficients[:, i] * head_svd.V[:, i] over i reconstructs the head’s real per-position output to numerical precision, since V’s columns are orthonormal and this projects onto the exact basis the head writes in.

Type:

jaxtyping.Float[Tensor, ‘pos rank’]

str_tokens

Tokenized prompt, aligned with the position axis, for display.

Type:

List[str]

coefficients: Float[Tensor, 'pos rank']
head_svd: HeadSVD
str_tokens: List[str]
exception transformer_lens.tools.analysis.svd_circuits.DegenerateDirectionError

Bases: ValueError

Raised when per-direction attribution is requested for a degenerate direction.

Subclasses ValueError so callers that already except ValueError keep working, mirroring how the other analysis tools raise ValueError for their input guards.

class transformer_lens.tools.analysis.svd_circuits.HeadDecomposition(layer: int, head: int, QK: HeadSVD | None = None, OV: HeadSVD | None = None)

Bases: object

Container returned by decompose_head().

QK and OV hold the HeadSVD for each requested map, or None when that map was not requested.

OV: HeadSVD | None = None
QK: HeadSVD | None = None
head: int
layer: int
class transformer_lens.tools.analysis.svd_circuits.HeadSVD(which: Literal['QK', 'OV'], layer: int, head: int, U: Float[Tensor, 'd_model rank'], S: Float[Tensor, 'rank'], V: Float[Tensor, 'd_model rank'], rank_report: List[RankReportRow], eps: float, null_rtol: float, folded_ln: bool = False)

Bases: object

Factored SVD of one head map ("QK" or "OV") with a degeneracy report.

which

Which map this decomposes, "QK" or "OV".

Type:

Literal[‘QK’, ‘OV’]

layer

Layer of the decomposed head.

Type:

int

head

Head index within the layer.

Type:

int

U

Left singular vectors, [d_model, rank]: column i is the map’s input direction i (for OV, the residual-stream direction this head’s value computation reads from; for QK, the destination/query-read direction).

Type:

jaxtyping.Float[Tensor, ‘d_model rank’]

S

Singular values, [rank], sorted descending.

Type:

jaxtyping.Float[Tensor, ‘rank’]

V

Right singular vectors, [d_model, rank]: column i is the map’s output direction i for OV (the residual-stream direction this head writes into, the one to project through W_U), or the source/key-read direction for QK (QK produces no write direction). The reconstruction is U @ S.diag() @ V.transpose(-2, -1).

Type:

jaxtyping.Float[Tensor, ‘d_model rank’]

rank_report

Per-direction RankReportRow list, aligned with the columns of U/V.

Type:

List[transformer_lens.tools.analysis.svd_circuits.RankReportRow]

eps

Relative gap below which adjacent directions share a block; every block boundary sits at a gap of at least eps.

Type:

float

null_rtol

Relative-to-top-singular-value tolerance below which a direction is numerically null.

Type:

float

folded_ln

Whether the model’s weights carried a folded final LayerNorm when this was decomposed. enable_compatibility_mode folds ln1 into W_V and centres W_O/W_U, so a decomposition describes the OV map only under the state it was built in; the readout and patch consumers refuse a decomposition whose state no longer matches the model.

Type:

bool

S: Float[Tensor, 'rank']
U: Float[Tensor, 'd_model rank']
V: Float[Tensor, 'd_model rank']
block_of(i: int) → List[int]

Return every direction index sharing direction i’s degeneracy block.

The result is a singleton [i] for an isolated direction and the full run for a degenerate one.

degenerate_blocks() → List[List[int]]

Return every degenerate block’s indices: the subspaces to attribute whole or skip.

eps: float
folded_ln: bool = False
head: int
is_degenerate(i: int) → bool

Whether direction i is refused: rotation-ambiguous inside a block, or null.

layer: int
null_rtol: float
rank_report: List[RankReportRow]
require_isolated(i: int) → None

Raise DegenerateDirectionError unless direction i is attributable alone.

The message names the cause, since a rotation-ambiguous block is still a subspace worth attributing while a null block carries no signal.

which: Literal['QK', 'OV']
class transformer_lens.tools.analysis.svd_circuits.LogitSignature(direction: int, values: Float[Tensor, 'token'])

Bases: object

Rank-1-reconstruction logit effect for one OV direction, per requested token.

direction

Which HeadSVD column this reconstructs.

Type:

int

values

Signed logit contribution, aligned with the requested tokens.

Type:

jaxtyping.Float[Tensor, ‘token’]

direction: int
values: Float[Tensor, 'token']
class transformer_lens.tools.analysis.svd_circuits.PatchResult(head_svd: HeadSVD, retained: List[int], original_metric: float, patched_metric: float, delta_metric: float, baseline_delta_metric: float, gated: bool)

Bases: object

Result of causally patching a head’s output onto a chosen OV singular subspace.

head_svd

The OV decomposition patched against.

Type:

transformer_lens.tools.analysis.svd_circuits.HeadSVD

retained

Direction indices whose span the head’s output was reconstructed onto; the complement was zeroed.

Type:

List[int]

original_metric

Metric value on the unmodified prompt.

Type:

float

patched_metric

Metric value after the subspace reconstruction.

Type:

float

delta_metric

patched_metric - original_metric.

Type:

float

baseline_delta_metric

the mean of the per-draw delta_metric magnitudes over several random control subspaces of the same width as retained, each drawn inside the head’s own OV span span(V) rather than from the full residual stream, so the control is the effect of an arbitrary same-size subspace of this head’s output rather than of an unrelated residual-stream direction. Averaging the magnitudes rather than the signed deltas keeps this a typical control effect that does not shrink when the controls mix sign, so it is non-negative.

Type:

float

gated

whether the retained subspace passed the causal test for the mode it was expressed in, comparing abs(delta_metric) against baseline_delta_metric (or an explicit threshold, if one was passed). For ablate (retain the complement), removing a load-bearing subspace should move the metric more than removing an arbitrary same-size one, so gated is abs(delta_metric) > threshold. For keep (retain only the given subspace), a subspace that reconstructs the head’s behavior should move the metric less than keeping an arbitrary same-size one, so gated is abs(delta_metric) < threshold. A single “moved more than baseline” test cannot answer both, since keep=S and ablate=complement(S) resolve to the same retained set.

Type:

bool

baseline_delta_metric: float
delta_metric: float
gated: bool
head_svd: HeadSVD
original_metric: float
patched_metric: float
retained: List[int]
class transformer_lens.tools.analysis.svd_circuits.RankReportRow(idx: int, sigma: float, sigma_ratio: float, is_degenerate: bool, is_null: bool, block_id: int)

Bases: object

One singular direction’s summary, aligned with column idx of U/V.

idx

Position of the direction, matching the column index in U and V.

Type:

int

sigma

The singular value for this direction.

Type:

float

sigma_ratio

sigma normalized by the largest singular value, in [0, 1].

Type:

float

is_degenerate

True when this direction is not attributable on its own: it shares a block with a neighbour (defined only up to a rotation within that block) or it is numerically null.

Type:

bool

is_null

True when sigma_ratio falls below null_rtol, so the singular vector is an arbitrary direction from the map’s null space.

Type:

bool

block_id

Index of the contiguous block this direction belongs to.

Type:

int

block_id: int
idx: int
is_degenerate: bool
is_null: bool
sigma: float
sigma_ratio: float
transformer_lens.tools.analysis.svd_circuits.decompose_head(model, layer: int, head: int, *, which: Sequence[str] = ('QK', 'OV'), eps: float = 0.01, null_rtol: float | None = None) → HeadDecomposition

Decompose a head’s QK (W_Q W_K^T) and/or OV (W_V W_O) maps via SVD.

Weight-space only: this reads the head’s per-block weights via the bridge’s model.blocks[layer].attn accessors and needs no forward pass and no compatibility mode. The returned factors are detached from the model. Whether the model’s weights carry a folded final LayerNorm is recorded on each returned HeadSVD so the readout and patch consumers can refuse a decomposition taken under a different state.

Parameters:
  • model – A TransformerBridge.

  • layer – Layer of the head to decompose.

  • head – Head index within the layer.

  • which – Which maps to decompose, a non-empty sequence drawn from ("QK", "OV"). A bare string is rejected rather than iterated.

  • eps – Relative gap at which a block of adjacent directions ends; directions closer than this share a block.

  • null_rtol – Relative-to-top-singular-value tolerance below which a direction counts as numerically null. Defaults to None, which resolves to d_model * torch.finfo(S.dtype).eps per map, matching torch.linalg.matrix_rank()’s default tolerance.

Returns:

A HeadDecomposition whose QK/OV fields hold a HeadSVD for each requested map.

Raises:

ValueError – If layer or head is out of range, or which is empty, a bare string, or contains an unknown entry.

transformer_lens.tools.analysis.svd_circuits.logit_signature(model, head_svd: HeadSVD, direction: int, tokens: int | Sequence[int] | Tensor) → LogitSignature

Signed logit effect of one OV direction’s rank-1 reconstruction on the given tokens.

Pure weight-space computation: reconstructs the head’s OV output along a single singular direction (S[direction] * V[:, direction], never U - see the module docstring) and projects it through W_U restricted to tokens. Runs no forward pass and builds no cache.

Parameters:
  • model – A TransformerBridge with the final LayerNorm folded into W_U; only its W_U is read.

  • head_svd – An OV HeadSVD from decompose_head().

  • direction – Column index of the singular direction to reconstruct.

  • tokens – Token id(s) to read the logit effect for.

Returns:

A LogitSignature with one value per requested token.

Raises:
  • ValueError – If head_svd.which != "OV", if direction is not in [0, rank), if head_svd was decomposed under a different folded-LayerNorm state than model now has, or if model is a TransformerBridge whose weights do not carry a folded final LayerNorm.

  • DegenerateDirectionError – If direction is not attributable alone (see HeadSVD.require_isolated()).

transformer_lens.tools.analysis.svd_circuits.patch_along_directions(model, head_svd: HeadSVD, prompt: str | Tensor, metric: Callable[[Tensor], float], *, keep: Sequence[int] | None = None, ablate: Sequence[int] | None = None, threshold: float | None = None, rng: Generator | None = None, n_baseline: int = 32) → PatchResult

Causally validate a claimed OV subfunction by reconstructing the head’s output onto it.

Requires head_svd.which == "OV": this reconstructs the write/output basis V, and QK has no such vector (see the module docstring). Runs the prompt with use_attn_result enabled once unmodified, once with the head’s hook_result slice reconstructed onto span(head_svd.V[:, retained]), and once per random control subspace. Each control subspace is drawn inside the head’s own OV span span(V) (not from the full residual stream, where a width-w random subspace would keep only w/d_model of a head output that itself occupies only rank/d_model of the stream; an in-span control of width w keeps w/rank of it), so a moved metric is compared against the effect of an arbitrary subspace of this head’s output of the same width. The per-draw control delta magnitudes are averaged over n_baseline draws so one lucky or unlucky draw does not decide the gate, and so controls that mix sign do not cancel into a smaller threshold. Averaging narrows but does not remove the seed dependence of the threshold; see the n_baseline entry below. Restores the model’s prior use_attn_result setting afterward.

The gate’s success condition depends on the mode the caller expressed, because keep=S and ablate=complement(S) resolve to the same retained set and a single “moved more than the control” test would answer only the ablate question. See PatchResult.gated.

Parameters:
  • model – A TransformerBridge.

  • head_svd – An OV HeadSVD from decompose_head().

  • prompt – A single prompt: a string or a [1, pos] token tensor.

  • metric – A function from the model’s logits to a scalar.

  • keep – Direction indices to retain; the rest are zeroed. Exactly one of keep/ablate must be given.

  • ablate – Direction indices to zero; the rest are retained.

  • threshold – Explicit gate threshold. Defaults to None, which uses baseline_delta_metric (already a magnitude) instead.

  • rng – Optional generator for the random control subspaces, for reproducibility. Defaults to a generator seeded with _DEFAULT_BASELINE_SEED so a bare call is reproducible rather than drawing from the global RNG.

  • n_baseline – Number of random in-span control subspaces to average the baseline delta over. Must be at least 1. The per-draw control magnitudes are heavy-tailed, so the averaged threshold still varies with the seed and a direction whose delta sits near it can gate either way; pass an explicit rng for a reproducible verdict and raise n_baseline when the verdict is close to the threshold. Each draw costs one forward pass.

Returns:

A PatchResult describing the patched, baseline, and original metrics.

Raises:
  • ValueError – If head_svd.which != "OV", if head_svd was decomposed under a different folded-LayerNorm state than model now has, if keep/ablate are both given or both omitted, if any index is out of [0, rank), if the retained set is empty (keep=[] or ablate over the full rank) or spans the full rank (keep over every direction) and no explicit threshold is supplied, if n_baseline < 1, or if the model is in training mode.

  • NotImplementedError – If the model’s attention adapter exposes no per-head result, so set_use_attn_result(True) cannot fork the attention output.

  • DegenerateDirectionError – If the retained directions split a degenerate block (see _validate_retained_blocks()).

transformer_lens.tools.analysis.svd_circuits.project_activations(model, head_svd: HeadSVD, prompt: str | Tensor) → ActivationProjection

Project a head’s actual per-position output onto its OV singular directions.

Requires head_svd.which == "OV": this projects onto the write/output basis V, and QK has no such vector (see the module docstring). Runs a real forward pass with use_attn_result enabled to read the per-head output (hook_result), then projects it onto head_svd.V. Restores the model’s prior use_attn_result setting afterward, since flipping that config flag as a side effect of a read-only analysis call would surprise a caller who already had hooks or a cache built around its prior state.

Parameters:
  • model – A TransformerBridge.

  • head_svd – An OV HeadSVD from decompose_head().

  • prompt – A single prompt (not a batch): a string or a [pos] or [1, pos] token tensor.

Returns:

An ActivationProjection with the per-position coefficients.

Raises:
  • ValueError – If head_svd.which != "OV", if head_svd was decomposed under a different folded-LayerNorm state than model now has, if prompt is not a single prompt (a tensor that is not 1-D, or is 2-D with a leading dimension greater than one), or if the model is in training mode.

  • NotImplementedError – If the model’s attention adapter exposes no per-head result, so set_use_attn_result(True) cannot fork the attention output.

transformer_lens.tools.analysis.svd_circuits.vocab_readout(model, head_svd: HeadSVD, *, k: int = 10) → Float[Tensor, 'd_vocab k']

Project the top-k OV output directions through the unembedding.

Requires head_svd.which == "OV": QK produces no write direction to project (see the module docstring). On a TransformerBridge, the final LayerNorm must be folded into W_U, which requires the bridge to have processed its weights with fold_ln enabled on an adapter that supports folding.

Does not call head_svd.require_isolated: a degenerate direction’s vocab readout is still a well-defined projection, unlike a per-direction causal claim, so it is not gated here. The contract that no direction is reported without a passing causal patch is enforced by patch_along_directions().

Parameters:
  • model – A TransformerBridge with the final LayerNorm folded into W_U; only its W_U is read.

  • head_svd – An OV HeadSVD from decompose_head().

  • k – Number of top singular directions to project.

Returns:

column i is direction i’s projection through the unembedding.

Return type:

W_U.T @ head_svd.V[:, :k], shape [d_vocab, k]

Raises:

ValueError – If head_svd.which != "OV", if k is not in (0, rank], if head_svd was decomposed under a different folded-LayerNorm state than model now has, or if model is a TransformerBridge whose weights do not carry a folded final LayerNorm.