Skip to content
Merged
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
733 changes: 201 additions & 532 deletions demos/Jacobian_Lens_Demo.ipynb

Large diffs are not rendered by default.

128 changes: 128 additions & 0 deletions tests/unit/tools/test_jacobian_lens.py
Original file line number Diff line number Diff line change
Expand Up @@ -23,6 +23,7 @@
_ToyBlock,
_ToyBridge,
)
from transformer_lens.ActivationCache import ActivationCache
from transformer_lens.model_bridge import TransformerBridge
from transformer_lens.model_bridge.generalized_components import AltUpBlockBridge
from transformer_lens.model_bridge.supported_architectures.deepseek_v4 import (
Expand Down Expand Up @@ -926,6 +927,133 @@ def test_swap_leaves_orthogonal_complement_unchanged(
assert orthogonal_part.abs().max().item() < 1e-4


def test_swap_clamp_holds_each_layer_at_exchanged_clean_coordinates(
Comment thread
jlarson4 marked this conversation as resolved.
toy_model: _ToyBridge, fitted_lens: JacobianLens
) -> None:
layers = [1, 2]
tokens = toy_model.to_tokens("a toy prompt")
_, clean_cache_dict = toy_model.run_with_cache(tokens)
clean_cache = ActivationCache(clean_cache_dict, toy_model)
hooks = fitted_lens.swap_clamp_hooks(toy_model, 3, 5, layers, clean_cache)

with toy_model.hooks(fwd_hooks=hooks):
_, clamped_cache = toy_model.run_with_cache(tokens)

for layer in layers:
name = f"blocks.{layer}.hook_out"
vectors = fitted_lens.lens_vectors(toy_model, [3, 5], layer)
pinv = torch.linalg.pinv(vectors.T)
clean_coords = clean_cache[name].float() @ pinv.T
clamped_coords = clamped_cache[name].float() @ pinv.T
torch.testing.assert_close(clamped_coords, clean_coords[..., [1, 0]], atol=1e-5, rtol=1e-5)


def test_swap_clamp_accepts_plain_cache_dict(
toy_model: _ToyBridge, fitted_lens: JacobianLens
) -> None:
layer = 1
tokens = toy_model.to_tokens("a toy prompt")
_, clean_cache_dict = toy_model.run_with_cache(tokens, return_cache_object=False)
assert isinstance(clean_cache_dict, dict)
hooks = fitted_lens.swap_clamp_hooks(toy_model, 3, 5, [layer], clean_cache_dict)

with toy_model.hooks(fwd_hooks=hooks):
_, clamped_cache = toy_model.run_with_cache(tokens)

name = f"blocks.{layer}.hook_out"
vectors = fitted_lens.lens_vectors(toy_model, [3, 5], layer)
pinv = torch.linalg.pinv(vectors.T)
clean_coords = clean_cache_dict[name].float() @ pinv.T
clamped_coords = clamped_cache[name].float() @ pinv.T
torch.testing.assert_close(clamped_coords, clean_coords[..., [1, 0]], atol=1e-5, rtol=1e-5)


def test_swap_clamp_requires_each_clean_activation(toy_model: _ToyBridge) -> None:
with pytest.raises(ValueError, match="clean_cache is missing"):
_lens().swap_clamp_hooks(
toy_model, 3, 5, layers=[0], clean_cache=ActivationCache({}, toy_model)
)


def test_swap_clamp_rejects_identical_tokens(toy_model: _ToyBridge) -> None:
tokens = toy_model.to_tokens("a toy prompt")
_, clean_cache_dict = toy_model.run_with_cache(tokens)
clean_cache = ActivationCache(clean_cache_dict, toy_model)
with pytest.raises(ValueError, match="same token|identical|distinct"):
_lens().swap_clamp_hooks(toy_model, 3, 3, layers=[0], clean_cache=clean_cache)


def test_swap_clamp_warns_for_near_parallel_vectors() -> None:
model = _ToyBridge()
with torch.no_grad():
source = model.unembed.weight[3]
noise = torch.randn_like(source)
noise -= noise.dot(source) / source.square().sum() * source
noise *= 0.05 * source.norm() / noise.norm()
model.unembed.weight[5].copy_(source + noise)
tokens = model.to_tokens("a toy prompt")
_, clean_cache_dict = model.run_with_cache(tokens)
clean_cache = ActivationCache(clean_cache_dict, model)
with pytest.warns(UserWarning, match="parallel|ill-conditioned|poorly conditioned"):
hooks = _lens().swap_clamp_hooks(model, 3, 5, layers=[0], clean_cache=clean_cache)
assert hooks[0][0] == "blocks.0.hook_out"


def test_swap_clamp_validates_clean_activation_shape(toy_model: _ToyBridge) -> None:
tokens = toy_model.to_tokens("a toy prompt")
_, clean_cache_dict = toy_model.run_with_cache(tokens)
clean_cache = ActivationCache(clean_cache_dict, toy_model)
name = "blocks.0.hook_out"
clean_cache.cache_dict[name] = clean_cache[name][..., :-1]
with pytest.raises(ValueError, match="d_model|shape"):
_lens().swap_clamp_hooks(toy_model, 3, 5, layers=[0], clean_cache=clean_cache)


def test_swap_clamp_validates_clean_live_batch_shape(toy_model: _ToyBridge) -> None:
tokens = toy_model.to_tokens("a toy prompt")
_, clean_cache_dict = toy_model.run_with_cache(tokens)
clean_cache = ActivationCache(clean_cache_dict, toy_model)
name = "blocks.0.hook_out"
clean_cache.cache_dict[name] = clean_cache[name].repeat(2, 1, 1)
hooks = _lens().swap_clamp_hooks(toy_model, 3, 5, layers=[0], clean_cache=clean_cache)
with pytest.raises(ValueError, match="incompatible with the live activation"):
with toy_model.hooks(fwd_hooks=hooks):
toy_model(tokens)


def test_swap_clamp_positions_and_orthogonal_complement(
toy_model: _ToyBridge, fitted_lens: JacobianLens
) -> None:
layer = 2
name = f"blocks.{layer}.hook_out"
positions = [1, -1]
tokens = toy_model.to_tokens("a toy prompt")
_, clean_cache_dict = toy_model.run_with_cache(tokens)
clean_cache = ActivationCache(clean_cache_dict, toy_model)
hooks = fitted_lens.swap_clamp_hooks(
toy_model, 3, 5, layers=[layer], clean_cache=clean_cache, positions=positions
)
with toy_model.hooks(fwd_hooks=hooks):
_, clamped_cache = toy_model.run_with_cache(tokens)

normalized = [1, clean_cache[name].shape[1] - 1]
untouched = [index for index in range(clean_cache[name].shape[1]) if index not in normalized]
torch.testing.assert_close(clamped_cache[name][:, untouched], clean_cache[name][:, untouched])

vectors = fitted_lens.lens_vectors(toy_model, [3, 5], layer)
pinv = torch.linalg.pinv(vectors.T)
clean_coords = clean_cache[name][:, normalized].float() @ pinv.T
clamped_coords = clamped_cache[name][:, normalized].float() @ pinv.T
torch.testing.assert_close(clamped_coords, clean_coords[..., [1, 0]], atol=1e-5, rtol=1e-5)

delta = (clamped_cache[name][:, normalized] - clean_cache[name][:, normalized]).reshape(
-1, D_MODEL
)
basis, _ = torch.linalg.qr(vectors.T)
orthogonal_part = delta - (delta @ basis) @ basis.T
assert orthogonal_part.abs().max().item() < 1e-4


def test_exports() -> None:
from transformer_lens.tools import analysis

Expand Down
111 changes: 108 additions & 3 deletions transformer_lens/tools/analysis/jacobian_lens.py
Original file line number Diff line number Diff line change
Expand Up @@ -68,6 +68,7 @@
Callable,
Dict,
List,
Mapping,
MutableMapping,
Optional,
Sequence,
Expand All @@ -79,6 +80,7 @@
from jaxtyping import Float, Int
from tqdm.auto import tqdm

from transformer_lens.ActivationCache import ActivationCache
from transformer_lens.tools.analysis.jacobian_lens_coordinate_patch import (
CoordinatePatch,
solve_coordinate_patch,
Expand Down Expand Up @@ -1335,7 +1337,7 @@ def swap_hooks(
alpha: float = 1.0,
positions: Optional[Sequence[int]] = None,
) -> List[Tuple[str, Any]]:
"""Hooks that swap two concepts' coordinates in lens space.
"""Hooks that swap two concepts' *live* coordinates in lens space.

The paper's patching-in-lens-coordinates intervention: with
``V = [v_s, v_t]`` and lens coordinates ``c = V⁺ h`` (pseudoinverse),
Expand All @@ -1344,12 +1346,17 @@ def swap_hooks(
``span{v_s, v_t}`` is untouched. ``alpha=2`` is the paper's
"double-strength" swap.

This transform re-reads ``c`` from the activation seen by every hook.
It is therefore an involution when applied repeatedly in a subspace
whose coordinates are preserved between layers: a second application
can undo the first. For the paper's multi-layer clamp protocol, use
:meth:`swap_clamp_hooks` with activations cached from the clean run.

Args:
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).
layers: Layers to intervene at.
alpha: Swap strength.
positions: Chunk-local positions to swap (negative indices allowed
and normalized on every hook invocation). Defaults to all.
Expand Down Expand Up @@ -1398,6 +1405,104 @@ def transform(
)
return hooks

def swap_clamp_hooks(
Comment thread
jlarson4 marked this conversation as resolved.
self,
model: Any,
source_token: TokenInput,
target_token: TokenInput,
layers: Sequence[int],
clean_cache: Union[ActivationCache, Mapping[str, torch.Tensor]],
*,
positions: Optional[Sequence[int]] = None,
) -> List[Tuple[str, Any]]:
"""Hooks that clamp lens coordinates to their clean-run exchange.

For each layer, this projects the corresponding activation from
``clean_cache`` into that layer's lens basis, exchanges its source and
target coordinates once, and holds the live activation at that fixed
target. Unlike :meth:`swap_hooks`, the update is idempotent at each
layer: ``h <- h + V (c_target - V⁺h)``.

Args:
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.
clean_cache: Activations from an unmodified ``run_with_cache`` at
each requested layer's ``blocks.{layer}.hook_out`` name. Either
the ``ActivationCache`` it returns by default or the plain dict
from ``return_cache_object=False`` is accepted.
positions: Chunk-local positions to clamp (negative indices allowed
and normalized against the clean activations). Defaults to all.

Returns:
``[(hook_name, fn), ...]`` for ``model.hooks(fwd_hooks=...)``.
"""
self.validate_model(model)
source_id, target_id = _to_token_ids(model, [source_token, target_token])
if source_id == target_id:
raise ValueError(
"source_token and target_token resolve to the same token id; "
"a coordinate clamp would be a silent no-op"
)

hooks = []
for layer in [_normalize_layer(layer, model.cfg.n_layers) for layer in layers]:
hook_name = _resid_post_hook_name(layer)
if hook_name not in clean_cache:
raise ValueError(f"clean_cache is missing activation {hook_name!r}")
clean = clean_cache[hook_name]
_validate_residual_activation(
clean, d_model=model.cfg.d_model, hook_name=f"clean_cache[{hook_name!r}]"
)
normalized = _normalize_positions(positions, clean.shape[1])
clean_selected = clean if positions is None else clean[:, normalized, :]

vectors = self.lens_vectors(model, [source_id, target_id], layer)
units = _unit_rows(vectors, layer=layer)
_diagnose_intervention_pair(
units, description=f"swap vectors at layer {layer}", stacklevel=3
)
basis = vectors.T # [d, 2]
pinv = torch.linalg.pinv(basis) # [2, d]
target_coords = (clean_selected.float() @ pinv.T)[..., [1, 0]]
device_basis: Dict[torch.device, torch.Tensor] = {}
device_pinv: Dict[torch.device, torch.Tensor] = {}
device_targets: Dict[torch.device, torch.Tensor] = {}

def transform(
selected: Float[torch.Tensor, "batch pos d_model"],
basis: torch.Tensor = basis,
pinv: torch.Tensor = pinv,
target_coords: torch.Tensor = target_coords,
device_basis: Dict[torch.device, torch.Tensor] = device_basis,
device_pinv: Dict[torch.device, torch.Tensor] = device_pinv,
device_targets: Dict[torch.device, torch.Tensor] = device_targets,
) -> Float[torch.Tensor, "batch pos d_model"]:
local_basis = _cached_on_device(basis, device_basis, selected.device)
local_pinv = _cached_on_device(pinv, device_pinv, selected.device)
local_targets = _cached_on_device(target_coords, device_targets, selected.device)
if local_targets.shape[1] != selected.shape[1] or local_targets.shape[0] not in (
1,
selected.shape[0],
):
raise ValueError(
"clean_cache activation shape is incompatible with the live activation: "
f"target coordinates have shape {tuple(local_targets.shape)}, "
f"live activation has shape {tuple(selected.shape)}"
)
coords = selected.float() @ local_pinv.T
delta = (local_targets - coords) @ local_basis.T
return selected.float() + delta

hooks.append(
(
hook_name,
_make_intervention_hook(transform, positions, model.cfg.d_model),
)
)
return hooks

def coordinate_patch_hooks(
self,
model: Any,
Expand Down
Loading