Coverage for transformer_lens/model_bridge/supported_architectures/neel_solu_old.py: 32%
26 statements
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
1"""Neel Solu Old architecture adapter."""
3from typing import Any
5from transformer_lens.conversion_utils.conversion_steps import RearrangeTensorConversion
6from transformer_lens.conversion_utils.param_processing_conversion import (
7 ParamProcessingConversion,
8)
9from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
10from transformer_lens.model_bridge.generalized_components import (
11 AttentionBridge,
12 BlockBridge,
13 EmbeddingBridge,
14 MLPBridge,
15 NormalizationBridge,
16 PosEmbedBridge,
17 UnembeddingBridge,
18)
21class NeelSoluOldArchitectureAdapter(ArchitectureAdapter):
22 """Architecture adapter for Neel's SOLU models (old style)."""
24 def __init__(self, cfg: Any) -> None:
25 """Initialize the Neel SOLU old-style architecture adapter.
27 Args:
28 cfg: The configuration object.
29 """
30 super().__init__(cfg)
32 self.weight_processing_conversions = {
33 "blocks.{i}.attn.q.weight": ParamProcessingConversion(
34 tensor_conversion=RearrangeTensorConversion(
35 "d_model n_head d_head -> n_head d_model d_head"
36 ),
37 ),
38 "blocks.{i}.attn.k.weight": ParamProcessingConversion(
39 tensor_conversion=RearrangeTensorConversion(
40 "d_model n_head d_head -> n_head d_model d_head"
41 ),
42 ),
43 "blocks.{i}.attn.v.weight": ParamProcessingConversion(
44 tensor_conversion=RearrangeTensorConversion(
45 "d_model n_head d_head -> n_head d_model d_head"
46 ),
47 ),
48 "blocks.{i}.attn.o.weight": ParamProcessingConversion(
49 tensor_conversion=RearrangeTensorConversion(
50 "n_head d_head d_model -> n_head d_head d_model"
51 ),
52 ),
53 }
54 self.component_mapping = {
55 "embed": EmbeddingBridge(name="wte"),
56 "pos_embed": PosEmbedBridge(name="wpe"),
57 "blocks": BlockBridge(
58 name="blocks",
59 submodules={
60 "ln1": NormalizationBridge(name="ln1", config=self.cfg),
61 "attn": AttentionBridge(name="attn", config=self.cfg),
62 "ln2": NormalizationBridge(name="ln2", config=self.cfg),
63 "mlp": MLPBridge(name="mlp"),
64 },
65 ),
66 "ln_final": NormalizationBridge(name="ln_f", config=self.cfg),
67 "unembed": UnembeddingBridge(name="unembed"),
68 }
71def convert_neel_solu_old_weights(state_dict: dict, cfg: Any):
72 """
73 Converts the weights of my old SoLU models to the HookedTransformer format.
74 Takes as input a state dict, *not* a model object.
76 There are a bunch of dumb bugs in the original code, sorry!
78 Models 1L, 2L, 4L and 6L have left facing weights (ie, weights have shape
79 [dim_out, dim_in]) while HookedTransformer does right facing (ie [dim_in,
80 dim_out]).
82 8L has *just* a left facing W_pos, the rest right facing.
84 And some models were trained with
85 """
86 # Early models have left facing W_pos
87 reverse_pos = cfg.n_layers <= 8
89 # Models prior to 8L have left facing everything (8L has JUST left facing W_pos - sorry! Stupid bug)
90 reverse_weights = cfg.n_layers <= 6
92 new_state_dict = {}
93 for k, v in state_dict.items():
94 k = k.replace("norm", "ln")
95 if k.startswith("ln."):
96 k = k.replace("ln.", "ln_final.")
97 new_state_dict[k] = v
99 if reverse_pos:
100 new_state_dict["pos_embed.W_pos"] = new_state_dict["pos_embed.W_pos"].T
101 if reverse_weights:
102 for k, v in new_state_dict.items():
103 if "W_" in k and "W_pos" not in k:
104 new_state_dict[k] = v.transpose(-2, -1)
105 return new_state_dict