From 3ed421b6932fe02f81f74ba5b318b6a81321e14d Mon Sep 17 00:00:00 2001 From: lightyagami6669 Date: Sun, 2 Aug 2026 17:39:40 +0530 Subject: [PATCH 1/2] fix(bridge): make native TransformerBridge state_dict()/load_state_dict() true inverses state_dict() emits TL-renamed keys, but load_state_dict() only matched raw native names, so a round trip silently loaded nothing and strict=True was silently downgraded to strict=False. Adds the inverse key mapping (including aliased parameters reachable via multiple attribute paths, e.g. GPT-2's split q/k/v views into c_attn) and proper missing/unexpected key accounting that raises under strict=True. Fixes #1587 --- .../test_state_dict_round_trip.py | 131 ++++++++++++++++++ .../model_bridge/transformer_bridge.py | 77 ++++++++-- 2 files changed, 195 insertions(+), 13 deletions(-) create mode 100644 tests/unit/model_bridge/test_state_dict_round_trip.py diff --git a/tests/unit/model_bridge/test_state_dict_round_trip.py b/tests/unit/model_bridge/test_state_dict_round_trip.py new file mode 100644 index 000000000..c938cacf1 --- /dev/null +++ b/tests/unit/model_bridge/test_state_dict_round_trip.py @@ -0,0 +1,131 @@ +"""Regression tests for TransformerBridge.state_dict()/load_state_dict() round-tripping (#1587). + +state_dict() emits TL-renamed keys (e.g. "blocks.0.attn.q.weight"), but +load_state_dict() only matched raw native parameter names, so a +state_dict() -> load_state_dict() round trip silently loaded nothing and +strict=True was silently downgraded to strict=False. +""" +from __future__ import annotations + +import pytest +import torch + +from transformer_lens.config import TransformerBridgeConfig +from transformer_lens.model_bridge import TransformerBridge + + +def _native_cfg(**overrides) -> TransformerBridgeConfig: + base = dict( + d_model=32, + d_head=16, + n_heads=2, + n_layers=2, + n_ctx=8, + d_vocab=16, + d_mlp=64, + act_fn="gelu", + normalization_type="LN", + seed=0, + ) + base.update(overrides) + return TransformerBridgeConfig(**base) + + +def test_native_round_trip_overwrites_params_not_a_noop(): + bridge = TransformerBridge.boot_native(_native_cfg()) + + sd = {k: v.clone() for k, v in bridge.state_dict().items()} + assert sd, "state_dict() returned no TL-format keys" + + with torch.no_grad(): + for p in bridge.parameters(): + p.zero_() + assert all((p == 0).all() for p in bridge.parameters()) + + bridge.load_state_dict(sd, strict=True) + + # Compare against the snapshot directly rather than asserting "not all + # zero" - LayerNorm bias legitimately initializes to all-zero, so that + # check would pass even for a param that never got reloaded. + reloaded = bridge.state_dict() + for key, value in sd.items(): + assert torch.equal(reloaded[key], value), f"{key} did not round-trip" + + +def test_native_strict_true_raises_on_missing_key(): + bridge = TransformerBridge.boot_native(_native_cfg()) + sd = bridge.state_dict() + incomplete = dict(sd) + incomplete.pop(next(iter(incomplete))) + + with pytest.raises(RuntimeError, match="Missing key"): + bridge.load_state_dict(incomplete, strict=True) + + +def test_native_strict_true_raises_on_unexpected_key(): + bridge = TransformerBridge.boot_native(_native_cfg()) + sd = dict(bridge.state_dict()) + sd["totally.bogus.key"] = torch.zeros(1) + + with pytest.raises(RuntimeError, match="Unexpected key"): + bridge.load_state_dict(sd, strict=True) + + +def test_native_strict_false_does_not_raise_on_partial_dict(): + bridge = TransformerBridge.boot_native(_native_cfg()) + sd = bridge.state_dict() + first_key = next(iter(sd)) + partial = {first_key: sd[first_key]} + + result = bridge.load_state_dict(partial, strict=False) + assert result.unexpected_keys == [] + assert len(result.missing_keys) > 0 + + +def test_native_raw_keys_still_load_tracr_style(): + """boot_native's own raw parameter names must keep loading directly, + mirroring tracr's make_tracr_transformer_bridge_state_dict compatibility + contract (utilities/tracr.py).""" + bridge = TransformerBridge.boot_native(_native_cfg()) + raw_sd = {k: v.clone() for k, v in bridge.original_model.state_dict().items()} + + with torch.no_grad(): + for p in bridge.parameters(): + p.zero_() + + bridge.load_state_dict(raw_sd, strict=True) + + reloaded_raw = bridge.original_model.state_dict() + for key, value in raw_sd.items(): + assert torch.equal(reloaded_raw[key], value), f"{key} did not round-trip" + + +@pytest.mark.slow +def test_boot_transformers_round_trip_matches_forward_pass(): + """GPT-2's Conv1D-combined attention makes the bridge's q/k/v components + storage-sharing VIEWS into c_attn, not independent parameters - so this is + the case that actually exercises convert_hf_key_to_tl_key's HF-name + renaming, not just identity passthrough like boot_native does.""" + bridge = TransformerBridge.boot_transformers("gpt2", device="cpu") + bridge.eval() + + torch.manual_seed(0) + tokens = torch.randint(0, 1000, (1, 8)) + with torch.no_grad(): + logits_before = bridge(tokens).clone() + + sd = {k: v.clone() for k, v in bridge.state_dict().items()} + + with torch.no_grad(): + for p in bridge.parameters(): + p.zero_() + + bridge.load_state_dict(sd, strict=True) + + with torch.no_grad(): + logits_after = bridge(tokens).clone() + + max_diff = (logits_before - logits_after).abs().max().item() + assert torch.allclose( + logits_before, logits_after, atol=1e-5 + ), f"round trip did not restore forward-pass output: max diff={max_diff:.3e}" diff --git a/transformer_lens/model_bridge/transformer_bridge.py b/transformer_lens/model_bridge/transformer_bridge.py index b5c6a677a..64efac6d2 100644 --- a/transformer_lens/model_bridge/transformer_bridge.py +++ b/transformer_lens/model_bridge/transformer_bridge.py @@ -3481,9 +3481,39 @@ def state_dict(self, destination=None, prefix="", keep_vars=False): return tl_state_dict + def _tl_key_to_actual_keys(self) -> dict[str, list[str]]: + """Inverse of the renaming state_dict() applies: map each TL-format key + back to every raw parameter/buffer path that represents it. + + Mirrors the filtering and key-conversion in state_dict() exactly, except + it keeps every raw key for a given TL key instead of only the first-seen + one. Bridge components frequently expose the same underlying parameter + through more than one attribute path (e.g. GPT-2's split q/k/v weights + are views into the wrapped module's combined c_attn weight, reachable + both via a block-level shortcut and via the nested _original_component + chain) - all of those aliases must be written for the round trip to + actually change what forward() reads, not just what state_dict() shows. + """ + mapping: dict[str, list[str]] = {} + for actual_key in self.original_model.state_dict(): + if actual_key == "_original_component" or actual_key.startswith("_original_component."): + continue + clean_key = actual_key.replace("._original_component", "") + if not self._is_valid_bridge_path(clean_key): + continue + hf_key = self._normalize_bridge_key_to_hf(clean_key) + tl_key = self.adapter.convert_hf_key_to_tl_key(hf_key) + mapping.setdefault(tl_key, []).append(actual_key) + return mapping + def load_state_dict(self, state_dict, strict=True, assign=False): """Load state dict into the model, handling both clean keys and original keys with _original_component references. + Accepts three key formats: TL-format keys as emitted by state_dict() + (e.g. "blocks.0.attn.q.weight"), raw native parameter paths (e.g. for + ``boot_native`` / tracr-style loading), and raw paths with + "_original_component" segments stripped. + Args: state_dict: Dictionary containing a whole state of the module strict: Whether to strictly enforce that the keys in state_dict match the keys returned by this module's state_dict() function @@ -3494,26 +3524,47 @@ def load_state_dict(self, state_dict, strict=True, assign=False): """ current_state_dict = self.original_model.state_dict() clean_to_actual = {} - actual_to_clean = {} for actual_key in current_state_dict.keys(): if actual_key != "_original_component": - clean_key = actual_key.replace("._original_component", "") - clean_to_actual[clean_key] = actual_key - actual_to_clean[actual_key] = clean_key + clean_to_actual[actual_key.replace("._original_component", "")] = actual_key + + tl_to_actual = self._tl_key_to_actual_keys() + mapped_state_dict = {} + unexpected_keys = [] for input_key, value in state_dict.items(): if input_key in current_state_dict: mapped_state_dict[input_key] = value - else: - if input_key in clean_to_actual: - actual_key = clean_to_actual[input_key] + elif input_key in clean_to_actual: + mapped_state_dict[clean_to_actual[input_key]] = value + elif input_key in tl_to_actual: + for actual_key in tl_to_actual[input_key]: mapped_state_dict[actual_key] = value - else: - mapped_state_dict[input_key] = value - effective_strict = strict and len(mapped_state_dict) == len(current_state_dict) - return self.original_model.load_state_dict( - mapped_state_dict, strict=effective_strict, assign=assign - ) + else: + unexpected_keys.append(input_key) + + required_actual_keys = {key for keys in tl_to_actual.values() for key in keys} + missing_keys = sorted(required_actual_keys - mapped_state_dict.keys()) + + if strict and (missing_keys or unexpected_keys): + error_msgs = [] + if unexpected_keys: + error_msgs.append( + "Unexpected key(s) in state_dict: " + + ", ".join(f'"{k}"' for k in sorted(unexpected_keys)) + ) + if missing_keys: + error_msgs.append( + "Missing key(s) in state_dict: " + ", ".join(f'"{k}"' for k in missing_keys) + ) + raise RuntimeError( + "Error(s) in loading state_dict for {}:\n\t{}".format( + type(self.original_model).__name__, "\n\t".join(error_msgs) + ) + ) + + result = self.original_model.load_state_dict(mapped_state_dict, strict=False, assign=assign) + return type(result)(missing_keys=missing_keys, unexpected_keys=unexpected_keys) def get_params(self): """Access to model parameters in the format expected by SVDInterpreter. From 468094ff2ec703047f0e86d165aebf6502c2e080 Mon Sep 17 00:00:00 2001 From: lightyagami6669 Date: Mon, 3 Aug 2026 22:22:35 +0530 Subject: [PATCH 2/2] feat(utilities): add one-time converter for legacy TL-format checkpoints HookedTransformer checkpoints saved before TransformerBridge existed (OthelloGPT, grokking demos, ARENA content) use property-style keys (blocks.0.attn.W_Q, embed.W_E) and per-head tensor shapes that load_state_dict doesn't recognize natively. convert_tl_checkpoint converts these once into the key/tensor format TransformerBridge.boot_native(cfg).load_state_dict accepts, without teaching load_state_dict a second key convention. Validates per-head attention shapes against cfg before reshaping, since merging/splitting head dims produces a validly-shaped result for any head count, so a mismatched cfg would otherwise silently mis-group heads rather than trip a shape error. Handles GQA's leading-underscore _W_K/_W_V naming. Fixes #1588 --- docs/source/content/migrating_to_v3.md | 24 ++ .../test_tl_checkpoint_conversion.py | 211 ++++++++++++++++++ .../utilities/tl_checkpoint_conversion.py | 159 +++++++++++++ 3 files changed, 394 insertions(+) create mode 100644 tests/unit/model_bridge/test_tl_checkpoint_conversion.py create mode 100644 transformer_lens/utilities/tl_checkpoint_conversion.py diff --git a/docs/source/content/migrating_to_v3.md b/docs/source/content/migrating_to_v3.md index ddf48ee67..3b202fbc4 100644 --- a/docs/source/content/migrating_to_v3.md +++ b/docs/source/content/migrating_to_v3.md @@ -318,6 +318,30 @@ for position in range(tokens.shape[1]): HuggingFace model. It is not a `TransformerLensKeyValueCache`, and code should not depend on the latter's layout or methods. +### Load a legacy TL-format checkpoint + +Historical training-run checkpoints (OthelloGPT, grokking demos, ARENA +content) were saved via `HookedTransformer.state_dict()` before the bridge +existed, using property-style keys (`blocks.0.attn.W_Q`, `embed.W_E`, ...) +and per-head tensor shapes that `bridge.load_state_dict` doesn't recognize +natively. `convert_tl_checkpoint` is a one-time converter for exactly this: +convert once, load, then re-save in bridge format — `load_state_dict` itself +stays native-only rather than carrying a second, permanent key convention. + +```python +from transformer_lens.config import TransformerBridgeConfig +from transformer_lens.model_bridge import TransformerBridge +from transformer_lens.utilities.tl_checkpoint_conversion import convert_tl_checkpoint + +cfg = TransformerBridgeConfig(...) # same hyperparameters the checkpoint was trained under +legacy_state_dict = torch.load("othello_gpt.pth") + +bridge = TransformerBridge.boot_native(cfg) +bridge.load_state_dict(convert_tl_checkpoint(legacy_state_dict, cfg), strict=True) + +torch.save(bridge.state_dict(), "othello_gpt_bridge_format.pth") # re-save once, done +``` + ### Type helpers for both model classes Use the structural protocol when a helper should accept either a diff --git a/tests/unit/model_bridge/test_tl_checkpoint_conversion.py b/tests/unit/model_bridge/test_tl_checkpoint_conversion.py new file mode 100644 index 000000000..3cc9d9304 --- /dev/null +++ b/tests/unit/model_bridge/test_tl_checkpoint_conversion.py @@ -0,0 +1,211 @@ +"""Tests for the legacy TL-property-format checkpoint converter (#1588). + +Historical HookedTransformer checkpoints (OthelloGPT, grokking, ARENA content) +are saved with the old property-style keys ("blocks.0.attn.W_Q", "embed.W_E", +...) and per-head tensor shapes. `convert_tl_checkpoint` maps those onto the +key/tensor format `TransformerBridge.boot_native(cfg).load_state_dict` accepts +natively, so these checkpoints can be loaded once and re-saved in bridge +format without teaching `load_state_dict` a second key convention. +""" +from __future__ import annotations + +import pytest +import torch + +from transformer_lens import HookedTransformer +from transformer_lens.config import HookedTransformerConfig, TransformerBridgeConfig +from transformer_lens.model_bridge import TransformerBridge +from transformer_lens.utilities.tl_checkpoint_conversion import convert_tl_checkpoint + + +def _cfg_kwargs(**overrides): + base = dict( + d_model=32, + d_head=16, + n_heads=2, + n_layers=2, + n_ctx=8, + d_vocab=16, + d_mlp=64, + act_fn="gelu", + normalization_type="LN", + seed=0, + ) + base.update(overrides) + return base + + +def _ht_and_bridge_cfg(**overrides): + kwargs = _cfg_kwargs(**overrides) + return HookedTransformerConfig(**kwargs), TransformerBridgeConfig(**kwargs) + + +def test_convert_tl_checkpoint_loads_strict_into_native_bridge(): + ht_cfg, bridge_cfg = _ht_and_bridge_cfg() + ht = HookedTransformer(ht_cfg) + + converted = convert_tl_checkpoint(ht.state_dict(), bridge_cfg) + + bridge = TransformerBridge.boot_native(bridge_cfg) + result = bridge.load_state_dict(converted, strict=True) + + assert result.missing_keys == [] + assert result.unexpected_keys == [] + + +def test_convert_tl_checkpoint_matches_source_forward_pass(): + ht_cfg, bridge_cfg = _ht_and_bridge_cfg() + ht = HookedTransformer(ht_cfg) + + converted = convert_tl_checkpoint(ht.state_dict(), bridge_cfg) + bridge = TransformerBridge.boot_native(bridge_cfg) + bridge.load_state_dict(converted, strict=True) + + tokens = torch.randint(0, ht_cfg.d_vocab, (1, 4)) + with torch.no_grad(): + ht_logits = ht(tokens) + bridge_logits = bridge(tokens) + + torch.testing.assert_close(bridge_logits, ht_logits, atol=1e-4, rtol=1e-4) + + +def test_convert_tl_checkpoint_places_qkvo_in_correct_head_slots(): + """Independent check that per-head Q/K/V/O land in the right slots: read + the converted+loaded bridge back out through its own W_Q/W_K/W_V/W_O + properties (implemented separately from the converter) and compare + directly against the source HookedTransformer's per-head weights, rather + than trusting the converter's own reshape math.""" + ht_cfg, bridge_cfg = _ht_and_bridge_cfg() + ht = HookedTransformer(ht_cfg) + + converted = convert_tl_checkpoint(ht.state_dict(), bridge_cfg) + bridge = TransformerBridge.boot_native(bridge_cfg) + bridge.load_state_dict(converted, strict=True) + + torch.testing.assert_close(bridge.W_Q, ht.W_Q) + torch.testing.assert_close(bridge.W_K, ht.W_K) + torch.testing.assert_close(bridge.W_V, ht.W_V) + torch.testing.assert_close(bridge.W_O, ht.W_O) + torch.testing.assert_close(bridge.b_Q, ht.b_Q) + torch.testing.assert_close(bridge.b_K, ht.b_K) + torch.testing.assert_close(bridge.b_V, ht.b_V) + torch.testing.assert_close(bridge.b_O, ht.b_O) + + +def test_convert_tl_checkpoint_raises_on_cfg_mismatch(): + """A wrong cfg can't be caught by a downstream shape-mismatch error -- + merging per-head dims produces a validly-shaped result for any head + count, since d_model == n_heads * d_head for any factoring of it. The + converter must catch this itself.""" + ht_cfg, _ = _ht_and_bridge_cfg() + ht = HookedTransformer(ht_cfg) + + wrong_cfg = TransformerBridgeConfig( + **_cfg_kwargs(n_heads=4, d_head=8) + ) # same d_model, wrong split + + with pytest.raises(ValueError, match="attn.W_Q"): + convert_tl_checkpoint(ht.state_dict(), wrong_cfg) + + +def test_convert_tl_checkpoint_raises_on_unrecognized_key(): + _, bridge_cfg = _ht_and_bridge_cfg() + with pytest.raises(ValueError, match="not a recognized"): + convert_tl_checkpoint({"blocks.0.attn.totally_unknown_param": torch.zeros(1)}, bridge_cfg) + + +def test_convert_tl_checkpoint_supports_gqa(): + ht_cfg, bridge_cfg = _ht_and_bridge_cfg(n_heads=4, d_head=8, n_key_value_heads=2) + ht = HookedTransformer(ht_cfg) + + converted = convert_tl_checkpoint(ht.state_dict(), bridge_cfg) + bridge = TransformerBridge.boot_native(bridge_cfg) + result = bridge.load_state_dict(converted, strict=True) + + assert result.missing_keys == [] + assert result.unexpected_keys == [] + + tokens = torch.randint(0, ht_cfg.d_vocab, (1, 4)) + with torch.no_grad(): + ht_logits = ht(tokens) + bridge_logits = bridge(tokens) + torch.testing.assert_close(bridge_logits, ht_logits, atol=1e-4, rtol=1e-4) + + +def test_convert_tl_checkpoint_supports_lnpre(): + """OthelloGPT (this converter's motivating use case, #1588) uses + normalization_type="LNPre" -- param-free pre-norm, so ln1/ln2/ln_final + have no weight/bias keys at all in the state dict for this converter to + handle; this just confirms the round trip still works end to end.""" + ht_cfg, bridge_cfg = _ht_and_bridge_cfg(normalization_type="LNPre") + ht = HookedTransformer(ht_cfg) + + converted = convert_tl_checkpoint(ht.state_dict(), bridge_cfg) + bridge = TransformerBridge.boot_native(bridge_cfg) + result = bridge.load_state_dict(converted, strict=True) + + assert result.missing_keys == [] + assert result.unexpected_keys == [] + + tokens = torch.randint(0, ht_cfg.d_vocab, (1, 4)) + with torch.no_grad(): + ht_logits = ht(tokens) + bridge_logits = bridge(tokens) + torch.testing.assert_close(bridge_logits, ht_logits, atol=1e-4, rtol=1e-4) + + +def test_convert_tl_checkpoint_supports_attn_only(): + ht_cfg, bridge_cfg = _ht_and_bridge_cfg(attn_only=True) + ht = HookedTransformer(ht_cfg) + + converted = convert_tl_checkpoint(ht.state_dict(), bridge_cfg) + bridge = TransformerBridge.boot_native(bridge_cfg) + result = bridge.load_state_dict(converted, strict=True) + + assert result.missing_keys == [] + assert result.unexpected_keys == [] + + tokens = torch.randint(0, ht_cfg.d_vocab, (1, 4)) + with torch.no_grad(): + ht_logits = ht(tokens) + bridge_logits = bridge(tokens) + torch.testing.assert_close(bridge_logits, ht_logits, atol=1e-4, rtol=1e-4) + + +def test_convert_tl_checkpoint_drops_attention_buffers(): + ht_cfg, bridge_cfg = _ht_and_bridge_cfg() + ht = HookedTransformer(ht_cfg) + + converted = convert_tl_checkpoint(ht.state_dict(), bridge_cfg) + + assert not any(key.endswith((".mask", ".IGNORE")) for key in converted) + + +def test_convert_tl_checkpoint_supports_gated_mlp_and_rms_norm(): + """The native bridge's gated MLP has no bias parameter at all for + gate/in/out (matching how real gated-MLP HF architectures like Llama are + built) while HookedTransformer's gated MLP keeps live b_in/b_out + parameters (pre-existing mismatch between the two implementations, + unrelated to this converter). convert_tl_checkpoint still faithfully + translates those keys since they're real HT parameters; load_state_dict + is the one that should refuse them under strict=True. Here they're + exactly zero (freshly constructed, untrained model) so dropping them via + strict=False is lossless and the forward pass still matches exactly. + """ + ht_cfg, bridge_cfg = _ht_and_bridge_cfg(gated_mlp=True, normalization_type="RMS", act_fn="silu") + ht = HookedTransformer(ht_cfg) + + converted = convert_tl_checkpoint(ht.state_dict(), bridge_cfg) + bridge = TransformerBridge.boot_native(bridge_cfg) + result = bridge.load_state_dict(converted, strict=False) + + assert result.missing_keys == [] + assert set(result.unexpected_keys) == { + f"blocks.{i}.mlp.{part}.bias" for i in range(ht_cfg.n_layers) for part in ("in", "out") + } + + tokens = torch.randint(0, ht_cfg.d_vocab, (1, 4)) + with torch.no_grad(): + ht_logits = ht(tokens) + bridge_logits = bridge(tokens) + torch.testing.assert_close(bridge_logits, ht_logits, atol=1e-4, rtol=1e-4) diff --git a/transformer_lens/utilities/tl_checkpoint_conversion.py b/transformer_lens/utilities/tl_checkpoint_conversion.py new file mode 100644 index 000000000..81a6d2f37 --- /dev/null +++ b/transformer_lens/utilities/tl_checkpoint_conversion.py @@ -0,0 +1,159 @@ +"""One-time converter for legacy TL-property-format checkpoints (#1588). + +Historical training runs (OthelloGPT, grokking demos, ARENA content) were +saved via ``HookedTransformer.state_dict()`` before ``TransformerBridge`` +existed, using property-style keys ("blocks.0.attn.W_Q", "embed.W_E", ...) +and per-head tensor shapes. ``convert_tl_checkpoint`` maps those onto the +key/tensor format ``TransformerBridge.boot_native(cfg).load_state_dict`` +accepts natively, so these checkpoints can be converted once and re-saved in +bridge format. This is deliberately a standalone converter rather than a +second key convention taught to ``load_state_dict`` itself: convert once, +``bridge.load_state_dict(converted)``, then re-save with ``bridge.state_dict()``. +""" + +from __future__ import annotations + +from typing import Callable, Optional + +import einops +import torch + +from transformer_lens.config.transformer_bridge_config import TransformerBridgeConfig + +# Buffers that live on HookedTransformer's attention blocks but have no +# Parameter counterpart on the bridge side (causal mask, IGNORE sentinel). +_DROPPED_BUFFER_SUFFIXES = (".mask", ".IGNORE") + +TensorConvert = Callable[[torch.Tensor, TransformerBridgeConfig, str], torch.Tensor] + + +def _validate_shape(tensor: torch.Tensor, expected: tuple[int, ...], key: str) -> None: + # Merging/splitting per-head dims (unlike a plain transpose) produces a + # validly-shaped result for *any* head count, since d_model == n_heads * + # d_head for any factoring of it — a wrong cfg silently mis-groups heads + # without ever tripping a downstream shape-mismatch error. Check the + # untouched per-head shape explicitly before reshaping. + if tuple(tensor.shape) != expected: + raise ValueError( + f"convert_tl_checkpoint: {key!r} has shape {tuple(tensor.shape)}, " + f"expected {expected} for the given cfg. The checkpoint may not " + "match this cfg (n_heads/n_key_value_heads/d_head/d_model)." + ) + + +def _kv_heads(cfg: TransformerBridgeConfig) -> int: + return cfg.n_key_value_heads or cfg.n_heads + + +def _convert_w_q(t: torch.Tensor, cfg: TransformerBridgeConfig, key: str) -> torch.Tensor: + _validate_shape(t, (cfg.n_heads, cfg.d_model, cfg.d_head), key) + return einops.rearrange(t, "n_heads d_model d_head -> (n_heads d_head) d_model") + + +def _convert_w_kv(t: torch.Tensor, cfg: TransformerBridgeConfig, key: str) -> torch.Tensor: + _validate_shape(t, (_kv_heads(cfg), cfg.d_model, cfg.d_head), key) + return einops.rearrange(t, "n_heads d_model d_head -> (n_heads d_head) d_model") + + +def _convert_w_o(t: torch.Tensor, cfg: TransformerBridgeConfig, key: str) -> torch.Tensor: + _validate_shape(t, (cfg.n_heads, cfg.d_head, cfg.d_model), key) + return einops.rearrange(t, "n_heads d_head d_model -> d_model (n_heads d_head)") + + +def _convert_b_q(t: torch.Tensor, cfg: TransformerBridgeConfig, key: str) -> torch.Tensor: + _validate_shape(t, (cfg.n_heads, cfg.d_head), key) + return einops.rearrange(t, "n_heads d_head -> (n_heads d_head)") + + +def _convert_b_kv(t: torch.Tensor, cfg: TransformerBridgeConfig, key: str) -> torch.Tensor: + _validate_shape(t, (_kv_heads(cfg), cfg.d_head), key) + return einops.rearrange(t, "n_heads d_head -> (n_heads d_head)") + + +def _identity(t: torch.Tensor, cfg: TransformerBridgeConfig, key: str) -> torch.Tensor: + return t + + +def _transpose(t: torch.Tensor, cfg: TransformerBridgeConfig, key: str) -> torch.Tensor: + return t.T.contiguous() + + +# Old TL-property suffix -> (new bridge-key suffix, tensor conversion). +# Checked in order, longest/most-specific first, so e.g. ".b_Q" is matched +# before the generic ".b" LayerNorm-bias fallback. +_SUFFIX_CONVERSIONS: list[tuple[str, str, TensorConvert]] = [ + (".W_Q", ".q.weight", _convert_w_q), + # GQA stores K/V under a leading-underscore name (the raw Parameter); + # plain ".W_K"/".W_V" become expanding (non-Parameter) properties instead. + ("._W_K", ".k.weight", _convert_w_kv), + ("._W_V", ".v.weight", _convert_w_kv), + (".W_K", ".k.weight", _convert_w_kv), + (".W_V", ".v.weight", _convert_w_kv), + (".W_O", ".o.weight", _convert_w_o), + (".b_Q", ".q.bias", _convert_b_q), + ("._b_K", ".k.bias", _convert_b_kv), + ("._b_V", ".v.bias", _convert_b_kv), + (".b_K", ".k.bias", _convert_b_kv), + (".b_V", ".v.bias", _convert_b_kv), + (".b_O", ".o.bias", _identity), + (".W_in", ".in.weight", _transpose), + (".b_in", ".in.bias", _identity), + (".W_out", ".out.weight", _transpose), + (".b_out", ".out.bias", _identity), + (".W_gate", ".gate.weight", _transpose), + (".b_gate", ".gate.bias", _identity), + (".W_U", ".weight", _transpose), + (".b_U", ".bias", _identity), + (".W_E", ".weight", _identity), + (".W_pos", ".weight", _identity), + (".w", ".weight", _identity), + (".b", ".bias", _identity), +] + + +def _convert_key_and_tensor( + key: str, tensor: torch.Tensor, cfg: TransformerBridgeConfig +) -> Optional[tuple[str, torch.Tensor]]: + for old_suffix, new_suffix, convert in _SUFFIX_CONVERSIONS: + if key.endswith(old_suffix): + new_key = key[: -len(old_suffix)] + new_suffix + return new_key, convert(tensor, cfg, key) + return None + + +def convert_tl_checkpoint( + state_dict: dict[str, torch.Tensor], + cfg: TransformerBridgeConfig, +) -> dict[str, torch.Tensor]: + """Convert a legacy TL-property-format state dict to the key/tensor + format ``TransformerBridge.boot_native(cfg).load_state_dict`` accepts. + + Args: + state_dict: A state dict in the old ``HookedTransformer`` convention + (e.g. from ``HookedTransformer.state_dict()``), with keys like + ``"blocks.0.attn.W_Q"`` and per-head tensor shapes. + cfg: The config the checkpoint was trained/saved under. Used both to + reshape per-head attention weights and to validate that the + checkpoint's per-head shapes actually match this cfg — a + mismatched cfg would otherwise silently mis-group heads without + ever tripping a shape error, since d_model == n_heads * d_head + holds for any wrong factoring too. + + Returns: + A state dict with modern bridge keys (e.g. ``"blocks.0.attn.q.weight"``) + and flat ``nn.Linear``-oriented tensor shapes, ready for + ``bridge.load_state_dict(converted, strict=True)``. + """ + converted: dict[str, torch.Tensor] = {} + for key, tensor in state_dict.items(): + if key.endswith(_DROPPED_BUFFER_SUFFIXES): + continue + result = _convert_key_and_tensor(key, tensor, cfg) + if result is None: + raise ValueError( + f"convert_tl_checkpoint: don't know how to convert key {key!r} " + "(not a recognized TL-property parameter or buffer suffix)." + ) + new_key, new_tensor = result + converted[new_key] = new_tensor + return converted