diff --git a/tests/unit/test_activation_cache.py b/tests/unit/test_activation_cache.py index 30ff06d6a..f12969b10 100644 --- a/tests/unit/test_activation_cache.py +++ b/tests/unit/test_activation_cache.py @@ -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( diff --git a/transformer_lens/ActivationCache.py b/transformer_lens/ActivationCache.py index 1813adbf4..e7685851a 100644 --- a/transformer_lens/ActivationCache.py +++ b/transformer_lens/ActivationCache.py @@ -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. @@ -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) + logit_attrs = (scaled_residual_stack * logit_directions).sum(dim=-1) return logit_attrs