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
101 changes: 101 additions & 0 deletions tests/unit/test_activation_cache.py
Original file line number Diff line number Diff line change
Expand Up @@ -278,6 +278,107 @@ def test_stack_neuron_results_integer_neuron_slice_without_labels(
torch.testing.assert_close(actual, expected)


@pytest.mark.parametrize("pos_slice", [None, (1, 4), [0, 2], -1])
@pytest.mark.parametrize("batch_slice", [None, [0, 2], 1])
@pytest.mark.parametrize("logit_difference", [False, True])
def test_logit_attrs_per_example_targets_match_individual_examples(
activation_cache: ActivationCache,
pos_slice: tuple[int, int] | list[int] | int | None,
batch_slice: list[int] | int | None,
logit_difference: bool,
) -> None:
targets = torch.tensor([10, 11, 12])
incorrect = torch.tensor([13, 14, 15]) if logit_difference else None
batch_indices = torch.arange(3)
if batch_slice is not None:
targets = targets[batch_slice]
batch_indices = batch_indices[batch_slice]
if incorrect is not None:
incorrect = incorrect[batch_slice]

residual = activation_cache.decompose_resid(pos_slice=pos_slice)
assert isinstance(residual, torch.Tensor)
actual = activation_cache.logit_attrs(
residual,
tokens=targets,
incorrect_tokens=incorrect,
pos_slice=pos_slice,
batch_slice=batch_slice,
)

expected_rows = []
for index in batch_indices.reshape(-1).tolist():
single_cache = activation_cache.apply_slice_to_batch_dim(index)
single_residual = single_cache.decompose_resid(pos_slice=pos_slice)
assert isinstance(single_residual, torch.Tensor)
expected_rows.append(
single_cache.logit_attrs(
single_residual,
tokens=10 + index,
incorrect_tokens=13 + index if logit_difference else None,
pos_slice=pos_slice,
has_batch_dim=False,
)
)
expected = (
expected_rows[0] if isinstance(batch_slice, int) else torch.stack(expected_rows, dim=1)
)
assert actual.shape == expected.shape
torch.testing.assert_close(actual, expected)


@pytest.mark.parametrize(
"target_layout",
[
"scalar",
"per-position",
"batchless-per-position",
"shared-per-position",
"single-prompt-per-position",
],
)
@pytest.mark.parametrize("logit_difference", [False, True])
def test_logit_attrs_preserves_other_target_layouts(
activation_cache: ActivationCache,
target_layout: str,
logit_difference: bool,
) -> None:
cache = activation_cache
targets = torch.tensor([[10, 11, 12, 13], [14, 15, 16, 17], [18, 19, 20, 21]])
if target_layout == "scalar":
targets = torch.tensor(10)
elif target_layout == "batchless-per-position":
cache = cache.apply_slice_to_batch_dim(1)
targets = targets[1]
elif target_layout in ("shared-per-position", "single-prompt-per-position"):
targets = targets[1]
if target_layout == "single-prompt-per-position":
cache = cache.apply_slice_to_batch_dim([1])
incorrect = targets + 1 if logit_difference else None
residual = cache.decompose_resid()
assert isinstance(residual, torch.Tensor)
actual = cache.logit_attrs(
residual, targets, incorrect_tokens=incorrect, has_batch_dim=cache.has_batch_dim
)

# A scalar target at each position is independent of target broadcasting.
expected_positions = []
for position in range(4):
position_targets = targets if targets.ndim == 0 else targets[..., position]
expected_positions.append(
cache.logit_attrs(
residual[..., position, :],
position_targets,
incorrect_tokens=position_targets + 1 if logit_difference else None,
pos_slice=position,
has_batch_dim=cache.has_batch_dim,
)
)
expected = torch.stack(expected_positions, dim=-1)
assert actual.shape == expected.shape
torch.testing.assert_close(actual, expected)


@pytest.mark.parametrize("projection_ndim", [1, 2], ids=["vector", "matrix"])
@torch.no_grad()
def test_full_resid_decomposition_projected_ln_uses_mlp_input_scale(
Expand Down
17 changes: 16 additions & 1 deletion transformer_lens/ActivationCache.py
Original file line number Diff line number Diff line change
Expand Up @@ -578,7 +578,12 @@ def logit_attrs(
residual_stack:
Stack of components of residual stream to get logit attributions for.
tokens:
Tokens to compute logit attributions on.
Tokens to compute logit attributions on. A [batch] tensor selects one token
per example across all retained positions. A 1-D tensor whose length differs
from the batch size selects one token per position, shared across examples.
Use [1, position] to share per-position targets when batch and position sizes
are equal; [batch, position] selects targets for each example and position.
Token tensors should already match any batch/position slices.
incorrect_tokens:
If provided, compute attributions on logit difference between tokens and
incorrect_tokens. Must have the same shape as tokens.
Expand Down Expand Up @@ -643,6 +648,16 @@ def logit_attrs(
has_batch_dim=has_batch_dim,
)

if (
has_batch_dim
and batch_slice.mode != "int"
and scaled_residual_stack.ndim == 4
and logit_directions.ndim == 2
and logit_directions.shape[0] == scaled_residual_stack.shape[1]
):
# Per-example directions must broadcast over positions, not across the batch.
logit_directions = logit_directions.unsqueeze(-2)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

A 1-D target tensor whose length doesn't match the batch size used to mean one token per position, and this branch now treats it as per-example. On a single-prompt cache that silently returns a [component, pos, pos] grid, and on larger batches it raises. Could the new branch tell the two apart, and add a single-prompt case to the layouts test?


logit_attrs = (scaled_residual_stack * logit_directions).sum(dim=-1)
return logit_attrs

Expand Down
Loading