from unittest import mock
from tests.typecheck_errors import TYPECHECK_ERRORS
from transformer_lens.hook_points import HookPoint
def setup_hook_point_and_hook():
hook_point = HookPoint()
def hook(activation, hook):
return activation
return hook_point, hook
@mock.patch("torch.utils.hooks.RemovableHandle", autospec=True)
def test_add_hook_forward(mock_handle):
mock_handle.return_value.id = 0
hook_point, hook = setup_hook_point_and_hook()
hook_point.add_hook(hook, dir="fwd")
assert len(hook_point.fwd_hooks) == 1
@mock.patch("torch.utils.hooks.RemovableHandle", autospec=True)
def test_add_hook_backward(mock_handle):
mock_handle.return_value.id = 0
hook_point, hook = setup_hook_point_and_hook()
hook_point.add_hook(hook, dir="bwd")
assert len(hook_point.bwd_hooks) == 1
@mock.patch("torch.utils.hooks.RemovableHandle", autospec=True)
def test_add_hook_permanent(mock_handle):
mock_handle.return_value.id = 0
hook_point, hook = setup_hook_point_and_hook()
hook_point.add_hook(hook, dir="fwd", is_permanent=True)
assert hook_point.fwd_hooks[0].is_permanent
@mock.patch("torch.utils.hooks.RemovableHandle", autospec=True)
def test_add_hook_with_level(mock_handle):
mock_handle.return_value.id = 0
hook_point, hook = setup_hook_point_and_hook()
hook_point.add_hook(hook, dir="fwd", level=5)
assert hook_point.fwd_hooks[0].context_level == 5
@mock.patch("transformer_lens.hook_points.LensHandle")
@mock.patch("torch.utils.hooks.RemovableHandle")
def test_add_hook_prepend(mock_handle, mock_lens_handle):
mock_handle.id = 0
mock_handle.next_id = 1
hook_point, _ = setup_hook_point_and_hook()
def hook1(activation, hook):
return activation
def hook2(activation, hook):
return activation
class _LensHandleBox:
def __init__(self, handle, is_permanent, context_level, user_hook=None):
self.hook = handle
self.is_permanent = is_permanent
self.context_level = context_level
self.user_hook = user_hook
mock_lens_handle.side_effect = _LensHandleBox
next_id = {"val": 1}
def fake_register_forward_hook(fn, prepend=False):
handle = mock.MagicMock()
handle.id = next_id["val"]
next_id["val"] += 1
return handle
hook_point.register_forward_hook = fake_register_forward_hook
hook_point.add_hook(hook1, dir="fwd")
hook_point.add_hook(hook2, dir="fwd", prepend=True)
assert len(hook_point.fwd_hooks) == 2
assert hook_point.fwd_hooks[0].hook.id == 2
assert hook_point.fwd_hooks[1].hook.id == 1
def test_enable_reshape():
"""Test that enable_reshape sets the hook conversion correctly."""
from transformer_lens.conversion_utils.conversion_steps.base_tensor_conversion import (
BaseTensorConversion,
)
class TestHookConversion(BaseTensorConversion):
def handle_conversion(self, input_value, *full_context):
return input_value * 2
def revert(self, input_value, *full_context):
return input_value + 1
hook_point = HookPoint()
conversion = TestHookConversion()
hook_point.enable_reshape(conversion)
assert hook_point.hook_conversion is conversion
def test_enable_reshape_with_none():
"""Test that enable_reshape works with None values."""
hook_point = HookPoint()
hook_point.enable_reshape(None)
assert hook_point.hook_conversion is None
def test_reshape_functionality_integration():
"""Test that hook conversion works in an integration context."""
import torch
from transformer_lens.conversion_utils.conversion_steps.base_tensor_conversion import (
BaseTensorConversion,
)
class TestHookConversion(BaseTensorConversion):
def handle_conversion(self, input_value, *full_context):
return input_value * 2
def revert(self, input_value, *full_context):
return input_value + 10
class TestModule(torch.nn.Module):
def __init__(self):
super().__init__()
self.hook_point = HookPoint()
def forward(self, x):
return self.hook_point(x)
module = TestModule()
conversion = TestHookConversion()
module.hook_point.enable_reshape(conversion)
def test_hook(activation, hook):
return activation + 1
module.hook_point.add_hook(test_hook, dir="fwd")
test_input = torch.tensor([1.0, 2.0, 3.0])
result = module(test_input)
expected = torch.tensor([13.0, 15.0, 17.0])
assert torch.equal(result, expected)
def test_reshape_functionality_hook_returns_none_integration():
"""Test that output revert is not applied when hook returns None."""
import torch
from transformer_lens.conversion_utils.conversion_steps.base_tensor_conversion import (
BaseTensorConversion,
)
class TestHookConversion(BaseTensorConversion):
def handle_conversion(self, input_value, *full_context):
return input_value * 2
def revert(self, input_value, *full_context):
return input_value + 10
class TestModule(torch.nn.Module):
def __init__(self):
super().__init__()
self.hook_point = HookPoint()
def forward(self, x):
return self.hook_point(x)
module = TestModule()
conversion = TestHookConversion()
module.hook_point.enable_reshape(conversion)
def test_hook(activation, hook):
return None
module.hook_point.add_hook(test_hook, dir="fwd")
test_input = torch.tensor([1.0, 2.0, 3.0])
result = module(test_input)
assert torch.equal(result, test_input)
def test_alias_hook_preserves_earlier_replacement_when_later_alias_returns_none():
"""A later no-op alias must not discard an earlier alias replacement."""
import torch
hook_point = HookPoint()
seen = []
def selective_hook(activation, hook):
seen.append((hook.name, activation.item()))
if hook.name == "replace":
return activation + 10
return None
hook_point.add_hook(selective_hook, alias_names=["replace", "observe"])
result = hook_point(torch.tensor(1.0))
assert seen == [("replace", 1.0), ("observe", 11.0)]
assert result.item() == 11.0
def test_alias_hook_reverts_conversion_after_earlier_replacement():
"""Alias replacement remains eligible for output conversion reversion."""
import torch
from transformer_lens.conversion_utils.conversion_steps.base_tensor_conversion import (
BaseTensorConversion,
)
class ScaleThenShift(BaseTensorConversion):
def handle_conversion(self, input_value, *full_context):
return input_value * 2
def revert(self, input_value, *full_context):
return input_value + 100
hook_point = HookPoint()
hook_point.enable_reshape(ScaleThenShift())
def selective_hook(activation, hook):
if hook.name == "replace":
return activation + 10
return None
hook_point.add_hook(selective_hook, alias_names=["replace", "observe"])
result = hook_point(torch.tensor(1.0))
assert result.item() == 112.0
def test_backward_alias_hook_preserves_earlier_gradient_replacement():
"""A later no-op alias must not discard an earlier gradient replacement."""
import torch
hook_point = HookPoint()
seen = []
def selective_hook(gradient, hook):
seen.append((hook.name, gradient.item()))
if hook.name == "replace":
return gradient + 10
return None
hook_point.add_hook(
selective_hook,
dir="bwd",
alias_names=["replace", "observe"],
)
input_value = torch.tensor(2.0, requires_grad=True)
(hook_point(input_value) * 3).backward()
assert seen == [("replace", 3.0), ("observe", 13.0)]
assert input_value.grad is not None
assert input_value.grad.item() == 13.0
class TestHookPointHasHooks:
"""Comprehensive test suite for HookPoint.has_hooks method."""
def setup_method(self):
"""Set up fresh HookPoint and sample hook for each test."""
self.hook_point = HookPoint()
def sample_hook(activation, hook):
return activation
self.sample_hook = sample_hook
def test_no_hooks_returns_false(self):
"""Test that has_hooks returns False when no hooks are present."""
assert not self.hook_point.has_hooks()
assert not self.hook_point.has_hooks(dir="fwd")
assert not self.hook_point.has_hooks(dir="bwd")
assert not self.hook_point.has_hooks(dir="both")
def test_forward_hook_detection(self):
"""Test detection of forward hooks."""
self.hook_point.add_hook(self.sample_hook, dir="fwd")
assert self.hook_point.has_hooks()
assert self.hook_point.has_hooks(dir="fwd")
assert self.hook_point.has_hooks(dir="both")
assert not self.hook_point.has_hooks(dir="bwd")
def test_backward_hook_detection(self):
"""Test detection of backward hooks."""
self.hook_point.add_hook(self.sample_hook, dir="bwd")
assert self.hook_point.has_hooks()
assert self.hook_point.has_hooks(dir="bwd")
assert self.hook_point.has_hooks(dir="both")
assert not self.hook_point.has_hooks(dir="fwd")
def test_both_direction_hooks(self):
"""Test detection when both forward and backward hooks are present."""
self.hook_point.add_hook(self.sample_hook, dir="fwd")
self.hook_point.add_hook(self.sample_hook, dir="bwd")
assert self.hook_point.has_hooks()
assert self.hook_point.has_hooks(dir="fwd")
assert self.hook_point.has_hooks(dir="bwd")
assert self.hook_point.has_hooks(dir="both")
def test_permanent_hook_detection(self):
"""Test detection of permanent hooks."""
self.hook_point.add_hook(self.sample_hook, dir="fwd", is_permanent=True)
assert self.hook_point.has_hooks()
assert self.hook_point.has_hooks(including_permanent=True)
assert not self.hook_point.has_hooks(including_permanent=False)
def test_non_permanent_hook_detection(self):
"""Test detection of non-permanent hooks."""
self.hook_point.add_hook(self.sample_hook, dir="fwd", is_permanent=False)
assert self.hook_point.has_hooks()
assert self.hook_point.has_hooks(including_permanent=True)
assert self.hook_point.has_hooks(including_permanent=False)
def test_mixed_permanent_hooks(self):
"""Test detection with mix of permanent and non-permanent hooks."""
self.hook_point.add_hook(self.sample_hook, dir="fwd", is_permanent=True)
self.hook_point.add_hook(self.sample_hook, dir="fwd", is_permanent=False)
assert self.hook_point.has_hooks(including_permanent=True)
assert self.hook_point.has_hooks(including_permanent=False)
def test_only_permanent_hooks(self):
"""Test detection when only permanent hooks are present."""
self.hook_point.add_hook(self.sample_hook, dir="fwd", is_permanent=True)
self.hook_point.add_hook(self.sample_hook, dir="bwd", is_permanent=True)
assert self.hook_point.has_hooks(including_permanent=True)
assert self.hook_point.has_hooks(dir="fwd", including_permanent=True)
assert self.hook_point.has_hooks(dir="bwd", including_permanent=True)
assert not self.hook_point.has_hooks(including_permanent=False)
assert not self.hook_point.has_hooks(dir="fwd", including_permanent=False)
assert not self.hook_point.has_hooks(dir="bwd", including_permanent=False)
def test_context_level_filtering(self):
"""Test context level filtering functionality."""
self.hook_point.add_hook(self.sample_hook, dir="fwd", level=0)
self.hook_point.add_hook(self.sample_hook, dir="fwd", level=1)
self.hook_point.add_hook(self.sample_hook, dir="bwd", level=2)
assert self.hook_point.has_hooks(level=0)
assert self.hook_point.has_hooks(level=1)
assert self.hook_point.has_hooks(level=2)
assert not self.hook_point.has_hooks(level=3)
assert not self.hook_point.has_hooks(level=-1)
assert self.hook_point.has_hooks(level=None)
def test_context_level_with_direction(self):
"""Test context level filtering combined with direction filtering."""
self.hook_point.add_hook(self.sample_hook, dir="fwd", level=0)
self.hook_point.add_hook(self.sample_hook, dir="bwd", level=1)
assert self.hook_point.has_hooks(dir="fwd", level=0)
assert self.hook_point.has_hooks(dir="bwd", level=1)
assert not self.hook_point.has_hooks(dir="fwd", level=1)
assert not self.hook_point.has_hooks(dir="bwd", level=0)
def test_context_level_with_permanent_flags(self):
"""Test context level filtering combined with permanent hook filtering."""
self.hook_point.add_hook(self.sample_hook, dir="fwd", level=0, is_permanent=True)
self.hook_point.add_hook(self.sample_hook, dir="fwd", level=1, is_permanent=False)
assert self.hook_point.has_hooks(level=0, including_permanent=True)
assert not self.hook_point.has_hooks(level=0, including_permanent=False)
assert self.hook_point.has_hooks(level=1, including_permanent=True)
assert self.hook_point.has_hooks(level=1, including_permanent=False)
def test_all_parameters_combined(self):
"""Test all parameters combined in various ways."""
self.hook_point.add_hook(self.sample_hook, dir="fwd", level=0, is_permanent=True)
self.hook_point.add_hook(self.sample_hook, dir="fwd", level=1, is_permanent=False)
self.hook_point.add_hook(self.sample_hook, dir="bwd", level=0, is_permanent=False)
self.hook_point.add_hook(self.sample_hook, dir="bwd", level=2, is_permanent=True)
assert self.hook_point.has_hooks(dir="fwd", level=0, including_permanent=True)
assert not self.hook_point.has_hooks(dir="fwd", level=0, including_permanent=False)
assert self.hook_point.has_hooks(dir="fwd", level=1, including_permanent=False)
assert self.hook_point.has_hooks(dir="bwd", level=0, including_permanent=False)
assert not self.hook_point.has_hooks(dir="bwd", level=1, including_permanent=False)
assert self.hook_point.has_hooks(dir="bwd", level=2, including_permanent=True)
def test_invalid_direction_raises_error(self):
"""Test that invalid direction parameter raises error (caught by type checking)."""
import pytest
with pytest.raises(TYPECHECK_ERRORS):
self.hook_point.has_hooks(dir="invalid")
def test_multiple_hooks_same_criteria(self):
"""Test detection when multiple hooks match the same criteria."""
self.hook_point.add_hook(self.sample_hook, dir="fwd", level=0, is_permanent=False)
self.hook_point.add_hook(self.sample_hook, dir="fwd", level=0, is_permanent=False)
self.hook_point.add_hook(self.sample_hook, dir="fwd", level=0, is_permanent=False)
assert self.hook_point.has_hooks(dir="fwd", level=0, including_permanent=False)
def test_hook_removal_affects_detection(self):
"""Test that removing hooks affects detection."""
self.hook_point.add_hook(self.sample_hook, dir="fwd")
assert self.hook_point.has_hooks()
self.hook_point.remove_hooks(dir="both")
assert not self.hook_point.has_hooks()
def test_default_parameter_values(self):
"""Test that default parameter values work correctly."""
self.hook_point.add_hook(self.sample_hook, dir="fwd", is_permanent=True, level=0)
self.hook_point.add_hook(self.sample_hook, dir="bwd", is_permanent=False, level=1)
assert self.hook_point.has_hooks()
assert self.hook_point.has_hooks(dir="both", including_permanent=True, level=None)
def test_edge_case_empty_after_filtering(self):
"""Test edge case where hooks exist but are filtered out."""
self.hook_point.add_hook(self.sample_hook, dir="fwd", is_permanent=True, level=5)
assert not self.hook_point.has_hooks(including_permanent=False)
assert not self.hook_point.has_hooks(dir="bwd")
assert not self.hook_point.has_hooks(level=0)
assert not self.hook_point.has_hooks(dir="bwd", level=5, including_permanent=True)
def test_functional_hook_execution_still_works(self):
"""Test that has_hooks doesn't interfere with actual hook functionality."""
import torch
results = []
def test_hook(activation, hook):
results.append("hook_called")
return activation
self.hook_point.add_hook(test_hook, dir="fwd")
assert self.hook_point.has_hooks()
test_input = torch.tensor([1.0, 2.0, 3.0])
output = self.hook_point(test_input)
assert torch.equal(output, test_input)
assert "hook_called" in results
def test_hook_point_with_conversions(self):
"""Test has_hooks with hook conversions if they exist."""
import torch
def simple_hook(activation, hook):
return activation * 2
self.hook_point.add_hook(simple_hook, dir="fwd")
assert self.hook_point.has_hooks()
assert self.hook_point.has_hooks(dir="fwd")
test_input = torch.tensor([1.0, 2.0])
output = self.hook_point(test_input)
expected = torch.tensor([2.0, 4.0])
assert torch.allclose(output, expected)