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:
objectPer-position coefficients of a head’s actual output in its OV output basis.
- head_svd¶
The OV decomposition this was projected against.
- coefficients¶
[pos, rank];coefficients[:, i]is the signed amount of singular directioni(head_svd.V[:, i]) present in the head’s actual output at each position. Summingcoefficients[:, i] * head_svd.V[:, i]overireconstructs the head’s real per-position output to numerical precision, sinceV’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']¶
- str_tokens: List[str]¶
- exception transformer_lens.tools.analysis.svd_circuits.DegenerateDirectionError¶
Bases:
ValueErrorRaised when per-direction attribution is requested for a degenerate direction.
Subclasses
ValueErrorso callers that alreadyexcept ValueErrorkeep working, mirroring how the other analysis tools raiseValueErrorfor their input guards.
- class transformer_lens.tools.analysis.svd_circuits.HeadDecomposition(layer: int, head: int, QK: HeadSVD | None = None, OV: HeadSVD | None = None)¶
Bases:
objectContainer returned by
decompose_head().QKandOVhold theHeadSVDfor each requested map, orNonewhen that map was not requested.- 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:
objectFactored 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 throughW_U), or the source/key-read direction for QK (QK produces no write direction). The reconstruction isU @ S.diag() @ V.transpose(-2, -1).- Type:
jaxtyping.Float[Tensor, ‘d_model rank’]
- rank_report¶
Per-direction
RankReportRowlist, aligned with the columns ofU/V.
- 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_modefoldsln1intoW_Vand centresW_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
iis refused: rotation-ambiguous inside a block, or null.
- layer: int¶
- null_rtol: float¶
- rank_report: List[RankReportRow]¶
- require_isolated(i: int) None¶
Raise
DegenerateDirectionErrorunless directioniis 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:
objectRank-1-reconstruction logit effect for one OV direction, per requested token.
- direction¶
Which
HeadSVDcolumn 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:
objectResult of causally patching a head’s output onto a chosen OV singular subspace.
- head_svd¶
The OV decomposition patched against.
- 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_metricmagnitudes over several random control subspaces of the same width asretained, each drawn inside the head’s own OV spanspan(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)againstbaseline_delta_metric(or an explicit threshold, if one was passed). Forablate(retain the complement), removing a load-bearing subspace should move the metric more than removing an arbitrary same-size one, sogatedisabs(delta_metric) > threshold. Forkeep(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, sogatedisabs(delta_metric) < threshold. A single “moved more than baseline” test cannot answer both, sincekeep=Sandablate=complement(S)resolve to the same retained set.- Type:
bool
- baseline_delta_metric: float¶
- delta_metric: float¶
- gated: bool¶
- 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:
objectOne singular direction’s summary, aligned with column
idxofU/V.- idx¶
Position of the direction, matching the column index in
UandV.- Type:
int
- sigma¶
The singular value for this direction.
- Type:
float
- sigma_ratio¶
sigmanormalized 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_ratiofalls belownull_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].attnaccessors 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 returnedHeadSVDso 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 tod_model * torch.finfo(S.dtype).epsper map, matchingtorch.linalg.matrix_rank()’s default tolerance.
- Returns:
A
HeadDecompositionwhoseQK/OVfields hold aHeadSVDfor each requested map.- Raises:
ValueError – If
layerorheadis out of range, orwhichis 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], neverU- see the module docstring) and projects it throughW_Urestricted totokens. Runs no forward pass and builds no cache.- Parameters:
model – A
TransformerBridgewith the final LayerNorm folded intoW_U; only itsW_Uis read.head_svd – An OV
HeadSVDfromdecompose_head().direction – Column index of the singular direction to reconstruct.
tokens – Token id(s) to read the logit effect for.
- Returns:
A
LogitSignaturewith one value per requested token.- Raises:
ValueError – If
head_svd.which != "OV", ifdirectionis not in[0, rank), ifhead_svdwas decomposed under a different folded-LayerNorm state thanmodelnow has, or ifmodelis aTransformerBridgewhose weights do not carry a folded final LayerNorm.DegenerateDirectionError – If
directionis not attributable alone (seeHeadSVD.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 basisV, and QK has no such vector (see the module docstring). Runs the prompt withuse_attn_resultenabled once unmodified, once with the head’shook_resultslice reconstructed ontospan(head_svd.V[:, retained]), and once per random control subspace. Each control subspace is drawn inside the head’s own OV spanspan(V)(not from the full residual stream, where a width-wrandom subspace would keep onlyw/d_modelof a head output that itself occupies onlyrank/d_modelof the stream; an in-span control of widthwkeepsw/rankof 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 overn_baselinedraws 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 then_baselineentry below. Restores the model’s prioruse_attn_resultsetting afterward.The gate’s success condition depends on the mode the caller expressed, because
keep=Sandablate=complement(S)resolve to the same retained set and a single “moved more than the control” test would answer only theablatequestion. SeePatchResult.gated.- Parameters:
model – A
TransformerBridge.head_svd – An OV
HeadSVDfromdecompose_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/ablatemust be given.ablate – Direction indices to zero; the rest are retained.
threshold – Explicit gate threshold. Defaults to
None, which usesbaseline_delta_metric(already a magnitude) instead.rng – Optional generator for the random control subspaces, for reproducibility. Defaults to a generator seeded with
_DEFAULT_BASELINE_SEEDso 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
rngfor a reproducible verdict and raisen_baselinewhen the verdict is close to the threshold. Each draw costs one forward pass.
- Returns:
A
PatchResultdescribing the patched, baseline, and original metrics.- Raises:
ValueError – If
head_svd.which != "OV", ifhead_svdwas decomposed under a different folded-LayerNorm state thanmodelnow has, ifkeep/ablateare both given or both omitted, if any index is out of[0, rank), if the retained set is empty (keep=[]orablateover the full rank) or spans the full rank (keepover every direction) and no explicitthresholdis supplied, ifn_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 basisV, and QK has no such vector (see the module docstring). Runs a real forward pass withuse_attn_resultenabled to read the per-head output (hook_result), then projects it ontohead_svd.V. Restores the model’s prioruse_attn_resultsetting 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
HeadSVDfromdecompose_head().prompt – A single prompt (not a batch): a string or a
[pos]or[1, pos]token tensor.
- Returns:
An
ActivationProjectionwith the per-position coefficients.- Raises:
ValueError – If
head_svd.which != "OV", ifhead_svdwas decomposed under a different folded-LayerNorm state thanmodelnow has, ifpromptis 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 aTransformerBridge, the final LayerNorm must be folded intoW_U, which requires the bridge to have processed its weights withfold_lnenabled 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 bypatch_along_directions().- Parameters:
model – A
TransformerBridgewith the final LayerNorm folded intoW_U; only itsW_Uis read.head_svd – An OV
HeadSVDfromdecompose_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", ifkis not in(0, rank], ifhead_svdwas decomposed under a different folded-LayerNorm state thanmodelnow has, or ifmodelis aTransformerBridgewhose weights do not carry a folded final LayerNorm.