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
211 changes: 211 additions & 0 deletions tests/unit/test_activation_cache.py
Original file line number Diff line number Diff line change
@@ -1,3 +1,5 @@
from types import SimpleNamespace

import pytest
import torch

Expand All @@ -7,6 +9,215 @@
from transformer_lens.utilities import Slice


def _head_results_reference(z: torch.Tensor, weights: torch.Tensor) -> torch.Tensor:
# Independent per-head matmuls keep the float64 reference small at wide model widths.
return torch.stack(
[z[..., head, :] @ weights[head] for head in range(weights.shape[0])], dim=-2
)


@pytest.mark.parametrize("near_zero", [False, True])
@torch.enable_grad()
def test_compute_head_results_matches_wide_float32_gradients(
monkeypatch: pytest.MonkeyPatch,
near_zero: bool,
) -> None:
device = "cuda" if torch.cuda.is_available() else "cpu"
generator = torch.Generator(device=device).manual_seed(20261005)
monkeypatch.setattr(torch.backends.cuda.matmul, "allow_tf32", False)
z = torch.randn(1, 32, 32, 128, device=device, generator=generator, requires_grad=True)
weights = torch.randn(32, 128, 4096, device=device, generator=generator, requires_grad=True)
model = SimpleNamespace(
cfg=SimpleNamespace(n_layers=1, n_heads=32),
blocks=[SimpleNamespace(attn=SimpleNamespace(W_O=weights))],
)
cache = ActivationCache({"blocks.0.attn.hook_z": z}, model)
# Freeze a head-major cotangent independently of the implementation's output layout.
grad_output = torch.randn(32, 1, 32, 4096, device=device, generator=generator).movedim(0, -2)
if near_zero:
with torch.no_grad():
weights[..., 1::2] = weights[..., 0::2]
grad_output[..., 1::2] = -grad_output[..., 0::2] + 2**-16
reference_z = z.detach().double().cpu().requires_grad_()
reference_weights = weights.detach().double().cpu().requires_grad_()
expected = _head_results_reference(reference_z, reference_weights)
expected_grads = torch.autograd.grad(
expected, (reference_z, reference_weights), grad_output.double().cpu()
)

cache.compute_head_results()

actual = cache["blocks.0.attn.hook_result"]
actual_grads = torch.autograd.grad(actual, (z, weights), grad_output)
torch.testing.assert_close(actual.double().cpu(), expected, rtol=2e-5, atol=2e-4)
# Activation gradients reduce over 4096 terms and can nearly cancel across BLAS backends.
torch.testing.assert_close(
actual_grads[0].double().cpu(), expected_grads[0], rtol=1e-4, atol=1e-3
)
torch.testing.assert_close(
actual_grads[1].double().cpu(), expected_grads[1], rtol=2e-5, atol=2e-4
)
absolute_products = _head_results_reference(
grad_output.double().cpu().abs(), reference_weights.detach().abs().transpose(-1, -2)
)
error = actual_grads[0].double().cpu() - expected_grads[0]
assert error.norm() <= 8 * torch.finfo(torch.float32).eps * absolute_products.norm()


@torch.enable_grad()
def test_compute_head_results_preserves_cancelling_float16_weight_gradients() -> None:
z = torch.full((32, 4, 32), 5000, dtype=torch.float16)
z[16:] = -5000
weights = torch.zeros(4, 32, 4096, dtype=torch.float16, requires_grad=True)
model = SimpleNamespace(
cfg=SimpleNamespace(n_layers=1, n_heads=4),
blocks=[SimpleNamespace(attn=SimpleNamespace(W_O=weights))],
)
cache = ActivationCache({"blocks.0.attn.hook_z": z}, model, has_batch_dim=False)
reference_weights = weights.detach().double().requires_grad_()
expected = _head_results_reference(z.double(), reference_weights)
expected_grad = torch.autograd.grad(expected.sum(), reference_weights)[0]
assert torch.count_nonzero(expected_grad) == 0

cache.compute_head_results()

actual_grad = torch.autograd.grad(cache["blocks.0.attn.hook_result"].sum(), weights)[0]
torch.testing.assert_close(actual_grad.double(), expected_grad, rtol=0, atol=0)


@pytest.mark.parametrize("weight_grad", [False, True])
@pytest.mark.parametrize(
"has_batch_dim,dtype,noncontiguous",
[
(True, torch.float32, False),
(False, torch.float64, True),
(True, torch.float16, True),
(False, torch.bfloat16, False),
],
)
@torch.enable_grad()
def test_compute_head_results_matches_float64_values_and_gradients(
has_batch_dim: bool,
dtype: torch.dtype,
noncontiguous: bool,
weight_grad: bool,
) -> None:
shape = (2, 9, 7, 32) if has_batch_dim else (19, 7, 32)
z = torch.randn(shape, dtype=dtype)
weights = torch.randn(7, 32, 4096, dtype=dtype)
if noncontiguous:
z = z.transpose(-1, -2).contiguous().transpose(-1, -2)
weights = weights.transpose(-1, -2).contiguous().transpose(-1, -2)
z.requires_grad_()
weights.requires_grad_(weight_grad)
model = SimpleNamespace(
cfg=SimpleNamespace(n_layers=1, n_heads=7),
blocks=[SimpleNamespace(attn=SimpleNamespace(W_O=weights))],
)
cache = ActivationCache({"blocks.0.attn.hook_z": z}, model, has_batch_dim=has_batch_dim)
reference_z = z.detach().double().requires_grad_()
reference_weights = weights.detach().double().requires_grad_(weight_grad)
expected = _head_results_reference(reference_z, reference_weights)
tolerances = {
torch.float64: (1e-12, 1e-12),
torch.float32: (2e-5, 2e-4),
torch.float16: (1e-3, 1e-3),
torch.bfloat16: (1e-2, 1e-2),
}
rtol, atol = tolerances[dtype]

cache.compute_head_results()

actual = cache["blocks.0.attn.hook_result"]
assert actual.dtype == dtype
torch.testing.assert_close(actual.double(), expected, rtol=rtol, atol=atol)
grad_output = torch.randn_like(actual)
inputs = (z, weights) if weight_grad else (z,)
reference_inputs = (reference_z, reference_weights) if weight_grad else (reference_z,)
expected_grads = torch.autograd.grad(expected, reference_inputs, grad_output.double())
actual_grads = torch.autograd.grad(actual, inputs, grad_output)
for gradient_index, (actual_grad, expected_grad) in enumerate(
zip(actual_grads, expected_grads)
):
if dtype == torch.float32 and gradient_index == 0:
torch.testing.assert_close(actual_grad.double(), expected_grad, rtol=1e-4, atol=1e-3)
else:
torch.testing.assert_close(actual_grad.double(), expected_grad, rtol=rtol, atol=atol)


@pytest.mark.parametrize(
"z_dtype,weight_dtype",
[
(torch.float16, torch.float32),
(torch.bfloat16, torch.float32),
(torch.float32, torch.float64),
],
)
@pytest.mark.parametrize("has_batch_dim", [False, True])
@torch.enable_grad()
def test_compute_head_results_promotes_mixed_dtypes(
z_dtype: torch.dtype, weight_dtype: torch.dtype, has_batch_dim: bool
) -> None:
shape = (2, 5, 3, 8) if has_batch_dim else (5, 3, 8)
z = torch.randn(shape, dtype=z_dtype, requires_grad=True)
weights = torch.randn(3, 8, 17, dtype=weight_dtype, requires_grad=True)
model = SimpleNamespace(
cfg=SimpleNamespace(n_layers=1, n_heads=3),
blocks=[SimpleNamespace(attn=SimpleNamespace(W_O=weights))],
)
cache = ActivationCache({"blocks.0.attn.hook_z": z}, model, has_batch_dim=has_batch_dim)
reference_z = z.detach().double().requires_grad_()
reference_weights = weights.detach().double().requires_grad_()
expected = _head_results_reference(reference_z, reference_weights)

cache.compute_head_results()

actual = cache["blocks.0.attn.hook_result"]
assert actual.dtype == torch.promote_types(z_dtype, weight_dtype)
torch.testing.assert_close(actual.double(), expected, rtol=2e-5, atol=2e-5)
grad_output = torch.randn_like(actual)
actual_grads = torch.autograd.grad(actual, (z, weights), grad_output)
expected_grads = torch.autograd.grad(
expected, (reference_z, reference_weights), grad_output.double()
)
for actual_grad, expected_grad in zip(actual_grads, expected_grads):
# The original leaf dtype determines gradient rounding after promotion.
torch.testing.assert_close(actual_grad, expected_grad.to(actual_grad.dtype))


@pytest.mark.parametrize(
"n_heads,d_head,n_tokens,grad_enabled",
[
(4, 32, 96, False),
(32, 128, 2, False),
(4, 32, 64, True),
(32, 128, 2, True),
],
)
def test_compute_head_results_limits_temporary_allocation(
n_heads: int, d_head: int, n_tokens: int, grad_enabled: bool
) -> None:
z = torch.randn(1, n_tokens, n_heads, d_head, dtype=torch.float32)
weights = torch.randn(n_heads, d_head, 4096, dtype=torch.float32, requires_grad=True)
model = SimpleNamespace(
cfg=SimpleNamespace(n_layers=1, n_heads=n_heads),
blocks=[SimpleNamespace(attn=SimpleNamespace(W_O=weights))],
)
cache = ActivationCache({"blocks.0.attn.hook_z": z}, model)

with torch.set_grad_enabled(grad_enabled):
with torch.profiler.profile(
activities=[torch.profiler.ProfilerActivity.CPU], profile_memory=True, acc_events=True
) as profile:
cache.compute_head_results()

allocations = [event.cpu_memory_usage for event in profile.events()]
actual = cache["blocks.0.attn.hook_result"]
assert max(allocations) <= actual.numel() * actual.element_size()
expected = _head_results_reference(z.double(), weights.detach().double())
torch.testing.assert_close(actual.double(), expected, rtol=2e-5, atol=2e-4)


@pytest.fixture(scope="module", params=["LN", "RMS"])
def activation_cache(request: pytest.FixtureRequest) -> ActivationCache:
cfg = TransformerBridgeConfig(
Expand Down
23 changes: 10 additions & 13 deletions transformer_lens/ActivationCache.py
Original file line number Diff line number Diff line change
Expand Up @@ -748,7 +748,9 @@ def compute_head_results(

Existing per-head results are preserved, including those without a batch dimension.
Missing or invalid results are recomputed from cached ``hook_z`` and the model's
``W_O`` weights.
``W_O`` weights. A per-head contraction avoids materializing the broadcast product
over both ``d_head`` and ``d_model``. As in the bridge forward pass, the contraction
follows active autocast settings.

TransformerBridge exposes ``blocks[i].attn.W_O`` via its component-mapping
compatibility shim.
Expand All @@ -767,18 +769,13 @@ def compute_head_results(
continue
all_cached = False

# Add singleton dimension to match W_O's shape for broadcasting
z = einops.rearrange(
self[("z", layer, "attn")],
"... head_index d_head -> ... head_index d_head 1",
)

# Element-wise multiplication of z and W_O (with shape [head_index, d_head, d_model])
block = self.model.blocks[layer]
result = z * block.attn.W_O

# Sum over d_head to get the contribution of each head to the residual stream
self.cache_dict[result_key] = result.sum(dim=-2)
z = self[("z", layer, "attn")]
weights = self.model.blocks[layer].attn.W_O
if z.dtype != weights.dtype:
# Preserve broadcast promotion for caches collected under autocast.
dtype = torch.promote_types(z.dtype, weights.dtype)
z, weights = z.to(dtype), weights.to(dtype)
self.cache_dict[result_key] = torch.einsum("...hd,hdm->...hm", z, weights)

if all_cached:
logging.warning("Tried to compute head results when they were already cached")
Expand Down
Loading