transformer_lens.tools.analysis package¶
Submodules¶
- transformer_lens.tools.analysis.direct_logit_attribution module
- transformer_lens.tools.analysis.direct_path_patching module
- transformer_lens.tools.analysis.jacobian_lens module
JacobianLensJacobianLens.jacobiansJacobianLens.n_promptsJacobianLens.d_modelJacobianLens.metadataJacobianLens.ablation_hooks()JacobianLens.clear_device_cache()JacobianLens.fit()JacobianLens.from_pretrained()JacobianLens.lens_vectors()JacobianLens.load()JacobianLens.merge()JacobianLens.readout()JacobianLens.save()JacobianLens.source_layersJacobianLens.steering_hooks()JacobianLens.swap_hooks()JacobianLens.transport()JacobianLens.validate_model()
JacobianLensReadoutJacobianLensReadout.lens_topk_valuesJacobianLensReadout.lens_topk_indicesJacobianLensReadout.model_topk_valuesJacobianLensReadout.model_topk_indicesJacobianLensReadout.lens_logitsJacobianLensReadout.model_logitsJacobianLensReadout.tokensJacobianLensReadout.positionsJacobianLensReadout.use_jacobianJacobianLensReadout.lens_logitsJacobianLensReadout.lens_topk_indicesJacobianLensReadout.lens_topk_valuesJacobianLensReadout.model_logitsJacobianLensReadout.model_topk_indicesJacobianLensReadout.model_topk_valuesJacobianLensReadout.positionsJacobianLensReadout.tokensJacobianLensReadout.top_tokens()JacobianLensReadout.use_jacobian
Module contents¶
Analysis tools for TransformerLens.
This subpackage collects high-level, single-call interpretability analyses that
sit on top of the hook/cache system. Model support is documented per tool;
new analyses may target the TransformerBridge API exclusively.
- Tools:
direct_logit_attribution: Direct Logit Attribution (DLA) over components, layers, or attention heads.
direct_path_patching: Direct path patching for head-to-head circuit analysis.
jacobian_lens: The Jacobian lens (J-lens) — per-layer causal transport to the output vocabulary basis, with loading of published lens artifacts, native fitting, readouts, and interventions.
- class transformer_lens.tools.analysis.DirectLogitAttribution(attribution: Float[Tensor, 'component *batch_and_pos'], labels: List[str], unit: str)¶
Bases:
objectResult of a
direct_logit_attribution()call.- attribution¶
Tensor of logit (or logit-difference) attributions with shape
[component, *batch_and_pos]. The leading axis is aligned withlabels. Whenposselects a single position (the default) the position axis is dropped, leaving[component, batch]— or[component]if the cache had its batch dimension removed.- Type:
jaxtyping.Float[Tensor, ‘component *batch_and_pos’]
- labels¶
Human-readable name for each component, aligned with the leading axis of
attribution(e.g."embed","0_attn_out","L3H7").- Type:
List[str]
- unit¶
The decomposition unit used (“component”, “layer”, or “head”).
- Type:
str
- attribution: Float[Tensor, 'component *batch_and_pos']¶
- labels: List[str]¶
- top(k: int = 5) List[tuple]¶
Return the
khighest-attribution(label, value)pairs.Attribution is reduced to a scalar per component by meaning over any remaining batch/position dimensions, so this is most meaningful when a single position was selected.
- unit: str¶
- class transformer_lens.tools.analysis.JacobianLens(jacobians: Dict[int, Float[Tensor, 'd_model d_model']], *, n_prompts: int, d_model: int, metadata: Dict[str, Any] | None = None)¶
Bases:
objectA fitted Jacobian lens: one transport matrix per source layer.
Layer convention (matching the reference implementation and the published artifacts): index
lrefers to the output of blocklat the Bridge-native hookblocks.{l}.hook_out.J[l]maps that activation to the final block’s output, pre final norm. The final layer itself is never fitted (its transport is the identity), sosource_layers == [0, ..., n_layers - 2]for a full fit.- jacobians¶
{layer: [d_model, d_model]}transport matrices, fp32, CPU.
- n_prompts¶
Number of prompts averaged into the fit.
- d_model¶
Residual stream width the lens was fitted for.
- metadata¶
Optional provenance (model name, TransformerLens version, fit hyperparameters). Preserved by
save()/load(); artifacts from the reference implementation load with empty metadata.
- ablation_hooks(model: Any, tokens: str | int | Sequence[str | int], layers: Sequence[int], *, positions: Sequence[int] | None = None) List[Tuple[str, Any]]¶
Hooks that project token directions out of the residual stream.
For each token’s unit lens vector
v̂:h <- h - (h·v̂) v̂, applied sequentially when several tokens are given.- Parameters:
model – The model the hooks will run on.
tokens – Concept token(s) to suppress.
layers – Layers to intervene at.
positions – Chunk-local positions to ablate (negative indices allowed and normalized on every hook invocation). Defaults to all.
- Returns:
[(hook_name, fn), ...]formodel.hooks(fwd_hooks=...).
- clear_device_cache() None¶
Release lazily cached Jacobian copies on accelerator devices.
- classmethod fit(model: Any, prompts: Sequence[str], *, corpus: str, source_layers: Sequence[int] | None = None, dim_batch: int = 8, max_seq_len: int = 128, skip_first_positions: int = 16, show_progress: bool = True, metadata: Dict[str, Any] | None = None) JacobianLens¶
Fit a Jacobian lens on a hooked model.
Implements the reference estimator exactly. For each prompt: one forward pass (the prompt replicated
dim_batchtimes along the batch axis), thenceil(d_model / dim_batch)backward passes. Each backward plants a one-hot cotangent for one output dimension at every valid target position simultaneously — causal attention guarantees the gradient at source positiontis then the sum over target positionst' >= twith no explicit masking. Rows are averaged over valid source positions (the firstskip_first_positionsand the final position are excluded), and prompts contribute equally to the final mean. There is no randomness: the computation is deterministic given the prompts.The reference implementation reports that fit quality saturates quickly — on the order of 100 prompts of 128 tokens is usable; the published lenses use up to 1000. Use
merge()to parallelize across prompt slices.- Parameters:
model – A raw
TransformerBridge. Model parameters are temporarily frozen (requires_grad=False) during fitting and restored after. Cotangents and activation gradients use the model dtype; fit with a float32 model for the highest-fidelity estimator.prompts – Prompt strings. Prompts too short to contain a valid position (
seq_len <= skip_first_positions + 1) are skipped with a warning and do not count towardn_prompts.corpus – Stable identifier for the prompt corpus or slice, recorded in artifact provenance.
source_layers – Layers to fit. Defaults to every layer below the final layer. Negative indices count from
n_layers.dim_batch – Output dimensions per backward pass. Higher is faster but replicates the prompt
dim_batchtimes in memory; total backward FLOPs are unchanged.max_seq_len – Prompts are truncated to this many tokens.
skip_first_positions – Leading positions excluded from the source average.
show_progress – Show a tqdm progress bar over prompts.
metadata – Extra provenance merged into
metadata.
- Returns:
The fitted
JacobianLens.- Raises:
TypeError – If model is not a
TransformerBridge.ValueError – On compatibility mode, invalid provenance or layer indices, or if no prompt was long enough to fit on.
- classmethod from_pretrained(name_or_path: str, *, filename: str = 'lens.pt', revision: str | None = None, model: Any = None) JacobianLens¶
Load a lens from a local path, a short model name, or a Hub repo.
Resolution order¶
Local file — if name_or_path is an existing
.ptfile, load it directly.Local directory — if name_or_path is a directory, load
<name_or_path>/<filename>.Registry short name or HF model ID — if name_or_path matches a key or alias in the bundled
jacobian_lens_registry.json(e.g."gemma-2-2b"or"google/gemma-2-2b"), the corresponding artifact inneuronpedia/jacobian-lensis fetched automatically. The filename argument is ignored in this case because the registry already encodes the correct subpath.Explicit Hub repo — otherwise name_or_path is treated as a Hub repo id and filename is used as-is, preserving full backward compatibility (e.g.
from_pretrained("neuronpedia/jacobian-lens", filename="gpt2-small/jlens/...")).
- param name_or_path:
A local
.ptfile, a local directory, a short model name such as"gemma-2-2b"or"llama3.1-8b", a Hugging Face model ID such as"google/gemma-2-2b", or an explicit Hub repo id paired with filename.- param filename:
File (or subpath) inside a local directory or an explicit Hub repo. Ignored when name_or_path resolves via the registry.
- param revision:
Optional Hub revision (branch, tag, or commit) to pin. When omitted, the Hub repository’s mutable default branch is followed; pin a commit hash for reproducible analyses.
- param model:
If given,
validate_model()is called so dimension or weight-processing mismatches fail here rather than at first use.- returns:
The loaded (and, if
modelwas given, validated) lens.
Examples:
# Short model name — no need to remember the HF subpath lens = JacobianLens.from_pretrained("gemma-2-2b", model=model) # HF model ID also works lens = JacobianLens.from_pretrained("google/gemma-2-2b", model=model) # Explicit Hub repo + subpath (backward-compatible) lens = JacobianLens.from_pretrained( "neuronpedia/jacobian-lens", filename="gpt2-small/jlens/Salesforce-wikitext/gpt2_jacobian_lens.pt", model=model, )
- lens_vectors(model: Any, tokens: str | int | Sequence[str | int], layer: int) Float[Tensor, 'n d_model']¶
Residual-stream directions for vocabulary tokens at a layer.
The J-lens vector for token
tis rowtofW_U J[layer]expressed in layer-layerresidual coordinates:v_t = J[layer]^T W_U[:, t].- Parameters:
model – The model supplying
W_U.tokens – A token string / id, or a sequence of them. Strings must encode to a single token.
layer – Source layer for the vectors.
- Returns:
One vector per token, fp32, on the model’s device.
- classmethod load(path: str) JacobianLens¶
Load a lens artifact or fit checkpoint saved in a supported schema.
Two file schemas are accepted:
Artifact (the reference format, written by
save()or the Anthropic reference package): must contain aJkey mapping layer indices to transport matrices, plusn_prompts,d_model, and an optionalmetadatadict.Fit checkpoint (running-sum format, written by the reference implementation’s
write_checkpoint()during fitting): must contain ajacobian_sumkey mapping layer indices to running-sum matrices (i.e. the sum over prompts, not yet divided by the prompt count), plusn_done.d_modelis inferred from the first matrix’s shape; no explicitd_modelkey is required or expected. The per-layer means are reconstructed on load. Aconverted_from: "jacobian_lens_checkpoint"key is added to metadata somerge()refuses to silently combine checkpoints with natively TL-fitted lenses. Fit-reserved provenance keys (transformer_lens_fit, etc.) are stripped; scalar fields that can be serialised underweights_only=Trueare preserved. Tensor-valued metadata fields that would fail_validate_metadata()are recorded by name in adropped_fieldslist.Fit checkpoint schema (reference
write_checkpoint()format)¶The reference implementation writes exactly six top-level keys; all other keys in the payload are ignored:
{ "jacobian_sum": {<layer int>: <float32 tensor [d, d]>, ...}, "n_done": <int>, # prompts accumulated into jacobian_sum "next_idx": <int>, # next prompt index (informational) "source_layers": [<int>, ...], # documented layer indices (informational) "target_layer": <int>, # target layer — harvested into metadata "skip_first": <int>, # leading positions skipped (informational) # optional flat provenance accepted from alternative checkpoint writers: "model_name": <str>, "model_revision": <str>, "corpus": <str>, # optional nested provenance accepted from alternative writers: "metadata": {<str>: <scalar/list/dict>, ...}, }
- param path:
Path to the
.ptfile.- raises ValueError:
If the file lacks both a
Jkey (artifact) and ajacobian_sumkey (checkpoint), or if a checkpoint records a non-positiven_prompts.
- classmethod merge(lenses: Sequence[JacobianLens]) JacobianLens¶
Combine lenses fitted on disjoint prompt slices.
The per-layer matrices are averaged weighted by each lens’s
n_prompts, matching the reference implementation, so fitting can be parallelized across processes or machines and merged afterwards. Provenance must match across shards (apart fromn_prompts), so a merge cannot silently relabel matrices fitted with different models, corpora, dtypes, or estimator settings. The merged count replaces the per-shard count.- Parameters:
lenses – Lenses that agree exactly on
source_layersandd_model.- Raises:
ValueError – On an empty sequence or mismatched lenses.
- readout(model: Any, input: str | Int[Tensor, 'batch seq'], *, layers: Sequence[int] | None = None, positions: Sequence[int] | None = None, use_jacobian: bool = True, top_k: int = 10, return_full_logits: bool = False) JacobianLensReadout¶
Read per-layer vocabulary logits for a prompt.
Runs the model once with caching, transports the residual stream at each requested layer through
J[layer](or the identity whenuse_jacobian=False— the logit lens), and applies the model’s own final norm, unembedding, architecture logit scaling, and logit soft cap.- Parameters:
model – A raw
TransformerBridge.input – A prompt string, or a
[1, seq]token tensor.layers – Layers to read. Defaults to every fitted layer plus the final layer. The final layer (
n_layers - 1) is always read with the identity transport — by construction its lens equals the model’s own output distribution.positions – Token positions to read (negative indices allowed). Defaults to all positions.
use_jacobian – Apply the Jacobian transport.
Falsegives the logit-lens baseline through the identical code path.top_k – Number of values and vocabulary ids retained per layer and position. Defaults to 10.
return_full_logits – Also retain full vocabulary tensors on CPU. This is opt-in because a 64-token Gemma readout across all layers is roughly 1.7 GB.
- Returns:
- Raises:
ValueError – If the model fails
validate_model(),inputis batched,top_kis invalid, or a requested layer has no transport matrix.
- save(path: str, *, dtype: dtype = torch.float16) None¶
Save the lens in the reference implementation’s artifact format.
The four official keys (
J,n_prompts,source_layers,d_model) are written unchanged so the file stays loadable by the reference package; TransformerLens provenance is stored under an additivemetadatakey.- Parameters:
path – Destination
.ptpath.dtype – Storage dtype. Defaults to fp16 like the reference implementation — Jacobian entries are order-one, so the smaller dtype costs little precision and halves the artifact on disk.
- property source_layers: List[int]¶
Sorted list of layers this lens has transport matrices for.
- steering_hooks(model: Any, token: str | int, layers: Sequence[int], *, alpha: float = 4.0, positions: Sequence[int] | None = None) List[Tuple[str, Any]]¶
Hooks that steer the residual stream along a token’s J-lens vector.
At each layer the unit-normalized lens vector is added, scaled by
alphatimes the activation’s median per-position residual norm:h <- h + alpha * median||h|| * v̂. This norm-matched parameterization follows the steering description in the reference implementation’s experiment protocols; the paper’s minimal form is the unscaledh <- h + alpha * v_t, recoverable by passing the rawlens_vectors()output to your own hook. The median (not mean) is used so attention-sink positions — whose residual norms run orders of magnitude above typical positions — do not inflate the scale.- Parameters:
model – The model the hooks will run on.
token – The concept token to steer toward.
layers – Layers to intervene at.
alpha – Steering strength scalar;
0disables. Because of the norm-matched scale, values of order 1 already perturb the stream by roughly its own magnitude.positions – Chunk-local positions to steer (negative indices allowed and normalized on every hook invocation). Defaults to all.
- Returns:
[(hook_name, fn), ...]formodel.hooks(fwd_hooks=...)ormodel.run_with_hooks(fwd_hooks=...).
- swap_hooks(model: Any, source_token: str | int, target_token: str | int, layers: Sequence[int], *, alpha: float = 1.0, positions: Sequence[int] | None = None) List[Tuple[str, Any]]¶
Hooks that swap two concepts’ coordinates in lens space.
The paper’s patching-in-lens-coordinates intervention: with
V = [v_s, v_t]and lens coordinatesc = V⁺ h(pseudoinverse), the update ish <- h + alpha * V (sigma(c) - c)wheresigmaexchanges the two coordinates. The component ofhorthogonal tospan{v_s, v_t}is untouched.alpha=2is the paper’s “double-strength” swap.- Parameters:
model – The model the hooks will run on.
source_token – The concept to remove (e.g.
" France").target_token – The concept to install (e.g.
" China").layers – Layers to intervene at (the paper clamps the swap across an intermediate-layer band).
alpha – Swap strength.
positions – Chunk-local positions to swap (negative indices allowed and normalized on every hook invocation). Defaults to all.
- Returns:
[(hook_name, fn), ...]formodel.hooks(fwd_hooks=...).
- transport(residual: Float[Tensor, '... d_model'], layer: int) Float[Tensor, '... d_model']¶
Map layer-
layeractivations into the final block’s output basis.Computes
J[layer] @ hper activation vector, in fp32.- Parameters:
residual – Activations from the output of block
layer.layer – Source layer index.
- validate_model(model: Any) JacobianLens¶
Check that
modelmatches this lens; raise loudly if not.Requires a raw causal
TransformerBridgewith the standard direct final-norm/unembed path, verifies recorded model provenance, residual width and layer range, and enforces the published final-block target convention.- Parameters:
model – A raw
TransformerBridge.- Returns:
self, for chaining.- Raises:
TypeError – If model is not a
TransformerBridge.ValueError – On model provenance or
d_modelmismatch, out-of-range source layers, compatibility mode, unsupported attention/output paths, or a non-final target convention.
- class transformer_lens.tools.analysis.JacobianLensReadout(lens_topk_values: Dict[int, Float[Tensor, 'pos k']], lens_topk_indices: Dict[int, Int[Tensor, 'pos k']], model_topk_values: Float[Tensor, 'pos k'], model_topk_indices: Int[Tensor, 'pos k'], tokens: Int[Tensor, 'seq'], positions: List[int], use_jacobian: bool = True, lens_logits: Dict[int, Float[Tensor, 'pos d_vocab']] | None = None, model_logits: Float[Tensor, 'pos d_vocab'] | None = None)¶
Bases:
objectResult of a
JacobianLens.readout()call.- lens_topk_values¶
Per-layer retained top-k pre-softmax values, on CPU.
- Type:
Dict[int, jaxtyping.Float[Tensor, ‘pos k’]]
- lens_topk_indices¶
Per-layer retained top-k vocabulary ids, on CPU.
- Type:
Dict[int, jaxtyping.Int[Tensor, ‘pos k’]]
- model_topk_values¶
The model output’s retained top-k pre-softmax values, on CPU.
- Type:
jaxtyping.Float[Tensor, ‘pos k’]
- model_topk_indices¶
The model output’s retained top-k vocabulary ids, on CPU.
- Type:
jaxtyping.Int[Tensor, ‘pos k’]
- lens_logits¶
Optional full per-layer logits, on CPU. Present only when
readout(return_full_logits=True)was requested.- Type:
Dict[int, jaxtyping.Float[Tensor, ‘pos d_vocab’]] | None
- model_logits¶
Optional full model logits, on CPU. Present only when
readout(return_full_logits=True)was requested.- Type:
jaxtyping.Float[Tensor, ‘pos d_vocab’] | None
- tokens¶
The token ids of the run prompt,
[seq].- Type:
jaxtyping.Int[Tensor, ‘seq’]
- positions¶
The (normalized, non-negative) positions the readout covers, aligned with the
posaxis of retained top-k and optional full logits.- Type:
List[int]
- use_jacobian¶
Whether the Jacobian transport was applied (
False= logit lens).- Type:
bool
- lens_logits: Dict[int, Float[Tensor, 'pos d_vocab']] | None = None¶
- lens_topk_indices: Dict[int, Int[Tensor, 'pos k']]¶
- lens_topk_values: Dict[int, Float[Tensor, 'pos k']]¶
- model_logits: Float[Tensor, 'pos d_vocab'] | None = None¶
- model_topk_indices: Int[Tensor, 'pos k']¶
- model_topk_values: Float[Tensor, 'pos k']¶
- positions: List[int]¶
- tokens: Int[Tensor, 'seq']¶
- top_tokens(tokenizer: Any, k: int = 5) Dict[int, List[List[str]]]¶
Decode the top-
ktokens per layer and position.- Parameters:
tokenizer – The model’s tokenizer (
model.tokenizer).k – Number of top tokens to decode per (layer, position).
- Returns:
{layer: [ [top-k strings] per position ]}, positions aligned withpositions.
- use_jacobian: bool = True¶
- transformer_lens.tools.analysis.direct_logit_attribution(model, input: str | List[str] | Tensor | None = None, answer_tokens: str | int | Tensor | None = None, incorrect_tokens: str | int | Tensor | None = None, *, unit: str = 'component', pos: int | Tuple[int] | Tuple[int, int] | Tuple[int, int, int] | List[int] | Tensor | ndarray | None = -1, cache: ActivationCache | None = None) DirectLogitAttribution¶
Compute Direct Logit Attribution for a prompt.
Decomposes the contribution of model components to the logit of
answer_tokens(or, ifincorrect_tokensis given, to the logit differenceanswer - incorrectalong theW_Udirection, which is usually what you want for circuit analysis).The model is run once with caching unless a precomputed
cacheis passed. Works with bothHookedTransformerandTransformerBridge.Note that DLA attributes only the part of a logit that comes from the residual stream through the unembedding direction; the unembedding bias
b_Uis a per-token constant that no component produces. So a complete decomposition reconstructslogit[token] - b_U[token]rather than the raw logit.On a
TransformerBridge, compatibility mode must be enabled (so the final LayerNorm is folded intoW_U) — otherwise the projection direction is wrong and DLA returns silently incorrect numbers. Hybrid architectures (Mamba/SSM/Mixer/LinearAttention) are not yet supported becausedecompose_residonly understands theattn_out + mlp_outblock layout; both conditions raise an explicit error at call time.- Parameters:
model – A
HookedTransformerorTransformerBridge(the latter withenable_compatibility_mode()already called).input – Prompt to run — a string, list of strings, or token tensor. Optional only when a precomputed
cacheis supplied.answer_tokens – The correct token(s) to attribute, as a string, id, or tensor. A string is converted with
model.to_single_token.incorrect_tokens – Optional baseline token(s). When given, attribution is computed for the
answer - incorrectresidual direction. Must broadcast to the same shape asanswer_tokens.unit –
Decomposition granularity:
"component"(default): embedding + each layer’s attention and MLP output (viadecompose_resid)."layer": cumulative residual stream after each sublayer, i.e. logit-lens trajectory (viaaccumulated_resid)."head": each attention head individually, plus a remainder term for everything else (viastack_head_results).
pos – Sequence position(s) to attribute. Defaults to
-1(the final token, the usual choice for next-token DLA). PassNoneto keep every position (the result then has a trailing position axis).cache – Optional precomputed
ActivationCacheto reuse instead of running the model again.
- Returns:
A
DirectLogitAttributionwithattribution(shape[component, *batch_and_pos]) and alignedlabels.- Raises:
ValueError – If
unitis invalid,answer_tokensisNone, neitherinputnorcacheis provided, or aTransformerBridgeis passed without compatibility mode enabled.NotImplementedError – If a
TransformerBridgereports a hybrid block layout (Mamba/SSM/Mixer/LinearAttention).
- transformer_lens.tools.analysis.get_act_patch_direct_path(model: HookedTransformer | TransformerBridge, corrupted_tokens: Tensor, clean_cache: ActivationCache, corrupted_cache: ActivationCache, patching_metric: Callable[[Tensor], Tensor], src_layer: int, src_head: int, component: Literal['q', 'k', 'v'] = 'q', verbose: bool = True) Float[Tensor, 'n_layers n_heads']¶
Sweep direct path patches from one source head to all downstream heads.
For every destination head B = (dst_layer, dst_head) where dst_layer > src_layer, patch the contribution of source head A = (src_layer, src_head) into B’s query (or key / value) input, and record the patching metric.
The patch is a linear approximation:
delta_resid = clean_A_result - corrupted_A_result [batch, pos, d_model] delta_B_comp = (delta_resid / ln1_scale) @ W_comp[dst_head]
where W_comp is W_Q, W_K, or W_V according to component.
- Parameters:
model – A HookedTransformer or TransformerBridge instance.
corrupted_tokens – Token IDs for the corrupted input, shape [batch, seq_len].
clean_cache – Cached activations from the clean (unpatched) run.
corrupted_cache – Cached activations from the corrupted run (needed for ln1 scale).
patching_metric – A function mapping the model’s logits tensor to a scalar.
src_layer – Layer index of the source attention head.
src_head – Head index of the source attention head.
component – Which input to patch at the destination head — “q” (default), “k”, or “v”.
verbose – Whether to show a tqdm progress bar.
- Returns:
results – results[dst_layer, dst_head] is the patching metric when the direct path A → B is patched in. Entries for dst_layer <= src_layer are left as 0.0 (no causal path from A to those layers).
- Return type:
Float[Tensor, “n_layers n_heads”]
- transformer_lens.tools.analysis.get_act_patch_direct_path_all_sources(model: HookedTransformer | TransformerBridge, corrupted_tokens: Tensor, clean_cache: ActivationCache, corrupted_cache: ActivationCache, patching_metric: Callable[[Tensor], Tensor], component: Literal['q', 'k', 'v'] = 'q', verbose: bool = True) Float[Tensor, 'n_layers n_heads n_layers n_heads']¶
Full sweep: all (src_layer, src_head) → (dst_layer, dst_head) direct paths.
Returns a 4-D tensor of shape [n_layers, n_heads, n_layers, n_heads]. result[sl, sh, dl, dh] = patching metric when head (sl,sh)’s output is patched directly into head (dl,dh)’s query/key/value input.
Entries where dl <= sl are 0 (no causal path).
This runs O(n_layers * n_heads * n_layers * n_heads) forward passes and is intended for small models or targeted sub-sweeps. For large models prefer calling get_act_patch_direct_path per source head.