From 5f7d644da2a638429468313a5aed1aa30cd26fea Mon Sep 17 00:00:00 2001 From: ArS377 Date: Sun, 2 Aug 2026 02:11:01 -0700 Subject: [PATCH] Fix native bridge state dict round trips --- tests/unit/model_bridge/test_boot_native.py | 32 +++++++++++++++++++++ tests/unit/test_tracr_conversion.py | 12 ++++++++ 2 files changed, 44 insertions(+) diff --git a/tests/unit/model_bridge/test_boot_native.py b/tests/unit/model_bridge/test_boot_native.py index ec998d8a6..6fac13dad 100644 --- a/tests/unit/model_bridge/test_boot_native.py +++ b/tests/unit/model_bridge/test_boot_native.py @@ -3,6 +3,7 @@ import sys +import pytest import torch from transformer_lens.config import TransformerBridgeConfig @@ -71,6 +72,37 @@ def test_boot_native_returns_bridge_over_native_model(): assert isinstance(bridge.original_model, NativeModel) +def test_native_state_dict_round_trip_restores_parameters(): + bridge = TransformerBridge.boot_native(_cfg()) + + saved_state_dict = {key: value.detach().clone() for key, value in bridge.state_dict().items()} + original_parameters = { + name: parameter.detach().clone() for name, parameter in bridge.named_parameters() + } + + with torch.no_grad(): + for parameter in bridge.parameters(): + parameter.zero_() + + result = bridge.load_state_dict(saved_state_dict, strict=True) + + assert result.missing_keys == [] + assert result.unexpected_keys == [] + + for name, parameter in bridge.named_parameters(): + torch.testing.assert_close(parameter, original_parameters[name]) + + +def test_native_state_dict_strict_rejects_unexpected_keys(): + bridge = TransformerBridge.boot_native(_cfg()) + + with pytest.raises(RuntimeError, match="Unexpected key"): + bridge.load_state_dict( + {"not.a.real.weight": torch.zeros(1)}, + strict=True, + ) + + def test_boot_native_accepts_dict_config(): cfg_dict = dict( d_model=32, diff --git a/tests/unit/test_tracr_conversion.py b/tests/unit/test_tracr_conversion.py index 0b9eb34de..cb16bf6ec 100644 --- a/tests/unit/test_tracr_conversion.py +++ b/tests/unit/test_tracr_conversion.py @@ -6,6 +6,7 @@ import pytest import torch +from transformer_lens.model_bridge import TransformerBridge from transformer_lens.utilities.tracr import ( infer_tracr_output_label, make_tracr_categorical_unembed, @@ -121,6 +122,17 @@ def test_bridge_config_matches_tracr_metadata(): assert cfg.attention_dir == "bidirectional" +def test_bridge_state_dict_loads_into_native_bridge(): + model = _fake_tracr_model() + bridge = TransformerBridge.boot_native(make_tracr_transformer_bridge_config(model)) + state_dict = make_tracr_transformer_bridge_state_dict(model, output_label="reverse_1") + + result = bridge.load_state_dict(state_dict, strict=False) + + assert result.unexpected_keys == [] + torch.testing.assert_close(bridge.state_dict()["embed.weight"], state_dict["tok_embed.weight"]) + + def test_bridge_state_dict_transposes_tracr_weights_and_reconstructs_unembed(): model = _fake_tracr_model()