From df09119141d9d8e0d07bb60d017d5821bede1c69 Mon Sep 17 00:00:00 2001 From: kigland Date: Tue, 11 Aug 2026 15:25:16 +0800 Subject: [PATCH] Fix BLOOM residual branch hook semantics --- .../model_bridge/test_bloom_hook_semantics.py | 43 +++++++++++++++++++ .../test_bloom_adapter.py | 5 +++ .../supported_architectures/bloom.py | 4 ++ 3 files changed, 52 insertions(+) create mode 100644 tests/integration/model_bridge/test_bloom_hook_semantics.py diff --git a/tests/integration/model_bridge/test_bloom_hook_semantics.py b/tests/integration/model_bridge/test_bloom_hook_semantics.py new file mode 100644 index 000000000..f2aead104 --- /dev/null +++ b/tests/integration/model_bridge/test_bloom_hook_semantics.py @@ -0,0 +1,43 @@ +"""Integration tests for BLOOM residual-branch hook semantics.""" + +import pytest +import torch + +from transformer_lens.model_bridge import TransformerBridge + +MODEL = "trl-internal-testing/tiny-BloomForCausalLM" + + +@pytest.fixture(scope="module") +def bloom_bridge() -> TransformerBridge: + return TransformerBridge.boot_transformers(MODEL, device="cpu", dtype=torch.float32) + + +def test_residual_branch_hooks_decompose_stream(bloom_bridge: TransformerBridge) -> None: + tokens = bloom_bridge.to_tokens("The capital of France is Paris.") + + with torch.no_grad(): + _, cache = bloom_bridge.run_with_cache(tokens) + + for layer in range(bloom_bridge.cfg.n_layers): + torch.testing.assert_close( + cache[f"blocks.{layer}.hook_resid_mid"], + cache[f"blocks.{layer}.hook_resid_pre"] + cache[f"blocks.{layer}.hook_attn_out"], + ) + torch.testing.assert_close( + cache[f"blocks.{layer}.hook_resid_post"], + cache[f"blocks.{layer}.hook_resid_mid"] + cache[f"blocks.{layer}.hook_mlp_out"], + ) + + +def test_residual_branch_hooks_are_writable(bloom_bridge: TransformerBridge) -> None: + tokens = bloom_bridge.to_tokens("The capital of France is Paris.") + + with torch.no_grad(): + baseline = bloom_bridge(tokens) + ablated = bloom_bridge.run_with_hooks( + tokens, + fwd_hooks=[("blocks.0.hook_attn_out", lambda tensor, hook: torch.zeros_like(tensor))], + ) + + assert not torch.equal(ablated, baseline) diff --git a/tests/unit/model_bridge/supported_architectures/test_bloom_adapter.py b/tests/unit/model_bridge/supported_architectures/test_bloom_adapter.py index 96212d404..a77c22183 100644 --- a/tests/unit/model_bridge/supported_architectures/test_bloom_adapter.py +++ b/tests/unit/model_bridge/supported_architectures/test_bloom_adapter.py @@ -103,6 +103,11 @@ def test_blocks_type_and_name(self, adapter: BloomArchitectureAdapter) -> None: assert isinstance(mapping["blocks"], BloomBlockBridge) assert mapping["blocks"].name == "transformer.h" + def test_residual_branch_hook_aliases(self, adapter: BloomArchitectureAdapter) -> None: + blocks = self._mapping(adapter)["blocks"] + assert blocks.hook_aliases["hook_attn_out"] == "attn.o.hook_out" + assert blocks.hook_aliases["hook_mlp_out"] == "mlp.out.hook_out" + def test_ln_final_type_and_name(self, adapter: BloomArchitectureAdapter) -> None: mapping = self._mapping(adapter) assert isinstance(mapping["ln_final"], NormalizationBridge) diff --git a/transformer_lens/model_bridge/supported_architectures/bloom.py b/transformer_lens/model_bridge/supported_architectures/bloom.py index 517f05376..2744d6d2e 100644 --- a/transformer_lens/model_bridge/supported_architectures/bloom.py +++ b/transformer_lens/model_bridge/supported_architectures/bloom.py @@ -69,6 +69,10 @@ def __init__(self, cfg: Any) -> None: "blocks": BloomBlockBridge( name="transformer.h", config=self.cfg, + hook_alias_overrides={ + "hook_attn_out": "attn.o.hook_out", + "hook_mlp_out": "mlp.out.hook_out", + }, submodules={ "ln1": NormalizationBridge(name="input_layernorm", config=self.cfg), "ln2": NormalizationBridge(name="post_attention_layernorm", config=self.cfg),