From 9b2fc0cc81ff921d555cd08a4b4675608ba7606f Mon Sep 17 00:00:00 2001 From: Almog Tavor Date: Sun, 5 Jul 2026 23:13:14 +0300 Subject: [PATCH 1/2] Mistral: use fork-capable PositionEmbeddingsAttentionBridge Mistral has separate q/k/v/o projections with RoPE and GQA, structurally identical to Qwen2, but its adapter used the plain AttentionBridge, which delegates q/k/v to HF and exposes no per-receiver fork point. As a result set_use_attn_in() raised and blocks.{i}.hook_attn_in did not exist, so per-head input interventions (activation/path patching on attention inputs) were impossible on Mistral. Switch the Mistral adapter to PositionEmbeddingsAttentionBridge (the same fork-capable bridge Qwen2/Qwen3 use) and pass config to its BlockBridge so hook_mlp_in is wired too. Add a test booting the tiny random Mistral that checks the attention bridge is fork-capable, hook_attn_in fires at [batch, pos, n_heads, d_model], and zeroing a single head's forked input actually changes the logits. --- .../model_bridge/test_mistral_attn_in_fork.py | 56 +++++++++++++++++++ .../supported_architectures/mistral.py | 5 +- 2 files changed, 59 insertions(+), 2 deletions(-) create mode 100644 tests/unit/model_bridge/test_mistral_attn_in_fork.py diff --git a/tests/unit/model_bridge/test_mistral_attn_in_fork.py b/tests/unit/model_bridge/test_mistral_attn_in_fork.py new file mode 100644 index 000000000..6bb7dd2ae --- /dev/null +++ b/tests/unit/model_bridge/test_mistral_attn_in_fork.py @@ -0,0 +1,56 @@ +"""Mistral adapter uses a fork-capable attention bridge (per-receiver ``hook_attn_in``). + +Mistral has separate q/k/v/o projections with RoPE + GQA, identical in structure to Qwen2, but +its adapter previously used the plain ``AttentionBridge`` (which delegates q/k/v to HF and exposes +no fork point), so ``set_use_attn_in`` raised and ``blocks.{i}.hook_attn_in`` did not exist. It now +uses ``PositionEmbeddingsAttentionBridge`` like Qwen2, enabling the input fork the circuit-analysis +tooling relies on. These tests boot the tiny random Mistral so they stay cheap. +""" +from __future__ import annotations + +import pytest +import torch + +from transformer_lens.model_bridge import TransformerBridge +from transformer_lens.model_bridge.generalized_components.position_embeddings_attention import ( + PositionEmbeddingsAttentionBridge, +) + +_MODEL = "hf-internal-testing/tiny-random-MistralForCausalLM" + + +@pytest.fixture(scope="module") +def bridge() -> TransformerBridge: + return TransformerBridge.boot_transformers(_MODEL, device="cpu") + + +@pytest.mark.slow +def test_mistral_uses_fork_capable_attention(bridge: TransformerBridge) -> None: + assert isinstance(bridge.blocks[0].attn, PositionEmbeddingsAttentionBridge) + + +@pytest.mark.slow +def test_mistral_hook_attn_in_shape_and_intervenes(bridge: TransformerBridge) -> None: + toks = torch.arange(1, 7).unsqueeze(0) + baseline = bridge(toks) + + bridge.set_use_attn_in(True) # previously raised on the plain AttentionBridge + + captured: dict = {} + + def _grab(tensor: torch.Tensor, hook: object) -> torch.Tensor: + captured["shape"] = tuple(tensor.shape) + return tensor + + bridge.run_with_hooks(toks, fwd_hooks=[("blocks.0.hook_attn_in", _grab)]) + # [batch, pos, n_heads, d_model] + assert captured["shape"] == (1, 6, bridge.cfg.n_heads, bridge.cfg.d_model) + + def _zero_head0(tensor: torch.Tensor, hook: object) -> torch.Tensor: + tensor = tensor.clone() + tensor[:, :, 0, :] = 0.0 + return tensor + + ablated = bridge.run_with_hooks(toks, fwd_hooks=[("blocks.0.hook_attn_in", _zero_head0)]) + # zeroing a single head's forked input must actually change the output + assert (ablated - baseline).abs().max().item() > 1e-6 diff --git a/transformer_lens/model_bridge/supported_architectures/mistral.py b/transformer_lens/model_bridge/supported_architectures/mistral.py index e0a373f02..a22789515 100644 --- a/transformer_lens/model_bridge/supported_architectures/mistral.py +++ b/transformer_lens/model_bridge/supported_architectures/mistral.py @@ -4,11 +4,11 @@ from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter from transformer_lens.model_bridge.generalized_components import ( - AttentionBridge, BlockBridge, EmbeddingBridge, GatedMLPBridge, LinearBridge, + PositionEmbeddingsAttentionBridge, RMSNormalizationBridge, RotaryEmbeddingBridge, UnembeddingBridge, @@ -49,10 +49,11 @@ def __init__(self, cfg: Any) -> None: "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb", config=self.cfg), "blocks": BlockBridge( name="model.layers", + config=self.cfg, submodules={ "ln1": RMSNormalizationBridge(name="input_layernorm", config=self.cfg), "ln2": RMSNormalizationBridge(name="post_attention_layernorm", config=self.cfg), - "attn": AttentionBridge( + "attn": PositionEmbeddingsAttentionBridge( name="self_attn", config=self.cfg, requires_position_embeddings=True, From 388540b3876fb5ad1c7c1edff7b7fa4f8627074b Mon Sep 17 00:00:00 2001 From: Almog Tavor Date: Thu, 23 Jul 2026 18:27:16 +0300 Subject: [PATCH 2/2] Skip non-fp params when picking attention compute dtype PositionEmbeddingsAttentionBridge used the first param's dtype, which is integer storage on quantized models (uint8 bnb, int32 GPTQ), so fp hidden states got cast to it. Ported the guard from AttentionBridge. Also updates the two Mistral tests that still assumed a plain AttentionBridge. The quantized one reads the dtype at q_proj now, since the reimplemented bridge never calls the HF attention forward. --- .../model_bridge/test_bridge_integration.py | 44 +++++++------------ .../test_mistral_adapter.py | 11 ++--- .../position_embeddings_attention.py | 15 ++++--- 3 files changed, 31 insertions(+), 39 deletions(-) diff --git a/tests/integration/model_bridge/test_bridge_integration.py b/tests/integration/model_bridge/test_bridge_integration.py index f26d2adcd..3e2a8a79e 100644 --- a/tests/integration/model_bridge/test_bridge_integration.py +++ b/tests/integration/model_bridge/test_bridge_integration.py @@ -704,22 +704,18 @@ def test_AttentionBridge_preserves_fp_input_when_first_param_is_quantized(): producing gibberish logits on every quantized model. Fakes a "quantized first parameter" by replacing q_proj.weight with a - uint8 tensor, then runs a forward and asserts the input the original - component receives is still floating-point. + uint8 tensor, then runs a forward and asserts the hidden states reaching + q_proj are still floating-point. """ from transformer_lens.model_bridge.generalized_components.attention import ( AttentionBridge, ) - # Use tiny Mistral — it's a plain AttentionBridge (not JointQKV). bridge: TransformerBridge = TransformerBridge.boot_transformers( # type: ignore "trl-internal-testing/tiny-MistralForCausalLM-0.2", device="cpu" ) attn_bridge = bridge.blocks[0].attn # type: ignore[attr-defined] - assert ( - type(attn_bridge).__name__ == "AttentionBridge" - ), f"Expected plain AttentionBridge, got {type(attn_bridge).__name__}" assert isinstance(attn_bridge, AttentionBridge) original = attn_bridge.original_component @@ -734,42 +730,32 @@ def test_AttentionBridge_preserves_fp_input_when_first_param_is_quantized(): next(original.parameters()).dtype == torch.uint8 ), "Test setup: first param should be uint8 to trigger the bug condition" - # Capture what dtype reaches the original component's forward. + # Capture what dtype reaches q_proj. The uint8 weight makes the matmul + # itself fail, which is fine: the pre-hook already recorded the dtype. received_dtype: list = [] - orig_forward = original.forward - - def capture(*args, **kwargs): - if "hidden_states" in kwargs: - received_dtype.append(kwargs["hidden_states"].dtype) - elif args: - received_dtype.append(args[0].dtype) - # Don't actually run forward — fake-quantized weight would error. - # Return a shape-compatible dummy. HF Mistral attention returns a tuple. - bsz, seq, d_model = (kwargs.get("hidden_states", args[0] if args else None)).shape - n_heads = bridge.cfg.n_heads # type: ignore[attr-defined] - return ( - torch.zeros(bsz, seq, d_model, dtype=torch.float32), - torch.zeros(bsz, n_heads, seq, seq, dtype=torch.float32), - ) - original.forward = capture # type: ignore[method-assign] + def capture(module, inputs): + if inputs and isinstance(inputs[0], torch.Tensor): + received_dtype.append(inputs[0].dtype) + + handle = original.q_proj.register_forward_pre_hook(capture) try: test_input = torch.tensor([[1, 2, 3, 4, 5]]) with torch.no_grad(): try: bridge(test_input) except Exception: - pass # downstream may fail; we only care what reached attn forward + pass # downstream may fail; we only care what reached q_proj finally: - original.forward = orig_forward # type: ignore[method-assign] + handle.remove() original.q_proj.weight = fp_weight - assert len(received_dtype) > 0, "Original attention forward never called" + assert len(received_dtype) > 0, "q_proj was never called" for dt in received_dtype: assert dt.is_floating_point, ( - f"Bridge passed dtype={dt} to original attention forward, but it must be " - f"floating point. Regression of the AttentionBridge dtype-cast bug — " - f"target_dtype must skip non-fp (quantized-storage) parameters." + f"Bridge passed dtype={dt} into q_proj, but it must be floating point. " + f"Regression of the AttentionBridge dtype-cast bug: target_dtype must " + f"skip non-fp (quantized-storage) parameters." ) diff --git a/tests/unit/model_bridge/supported_architectures/test_mistral_adapter.py b/tests/unit/model_bridge/supported_architectures/test_mistral_adapter.py index 34b286b67..69de1d7fd 100644 --- a/tests/unit/model_bridge/supported_architectures/test_mistral_adapter.py +++ b/tests/unit/model_bridge/supported_architectures/test_mistral_adapter.py @@ -112,12 +112,13 @@ def test_block_submodule_types(self, adapter: MistralArchitectureAdapter) -> Non assert isinstance(blocks.submodules["attn"], AttentionBridge) assert isinstance(blocks.submodules["mlp"], GatedMLPBridge) - def test_attn_is_not_position_embeddings_subclass( - self, adapter: MistralArchitectureAdapter - ) -> None: - """Mistral uses plain AttentionBridge, not PositionEmbeddingsAttentionBridge.""" + def test_attn_is_position_embeddings_bridge(self, adapter: MistralArchitectureAdapter) -> None: + """Mistral uses PositionEmbeddingsAttentionBridge, like Qwen2. + + That bridge is what gives Mistral hook_attn_in and the q/k/v forks. + """ attn = adapter.component_mapping["blocks"].submodules["attn"] - assert not isinstance(attn, PositionEmbeddingsAttentionBridge) + assert isinstance(attn, PositionEmbeddingsAttentionBridge) def test_block_submodule_hf_paths(self, adapter: MistralArchitectureAdapter) -> None: blocks = adapter.component_mapping["blocks"] diff --git a/transformer_lens/model_bridge/generalized_components/position_embeddings_attention.py b/transformer_lens/model_bridge/generalized_components/position_embeddings_attention.py index 5e43abf4e..9cba82b24 100644 --- a/transformer_lens/model_bridge/generalized_components/position_embeddings_attention.py +++ b/transformer_lens/model_bridge/generalized_components/position_embeddings_attention.py @@ -274,12 +274,17 @@ def forward(self, *args: Any, **kwargs: Any) -> Any: # Apply input hook hidden_states = self.hook_in(hidden_states) - # Match dtype of HF module + # Match dtype of HF module. Skip non-fp params: quantized weights (bnb + # uint8/int8, GPTQ/AWQ int32, HQQ, torchao) are stored in integer dtypes + # and dequantized internally during matmul. The compute dtype must come + # from a fp parameter; casting fp inputs to an integer storage dtype + # destroys precision. target_dtype = None - try: - target_dtype = next(hf_attn.parameters()).dtype - except StopIteration: - pass + for p in hf_attn.parameters(): + if not p.dtype.is_floating_point: + continue + target_dtype = p.dtype + break if target_dtype is not None and hidden_states.is_floating_point(): hidden_states = hidden_states.to(dtype=target_dtype)