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

1"""Neel Solu Old architecture adapter.""" 

2 

3from typing import Any 

4 

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) 

19 

20 

21class NeelSoluOldArchitectureAdapter(ArchitectureAdapter): 

22 """Architecture adapter for Neel's SOLU models (old style).""" 

23 

24 def __init__(self, cfg: Any) -> None: 

25 """Initialize the Neel SOLU old-style architecture adapter. 

26 

27 Args: 

28 cfg: The configuration object. 

29 """ 

30 super().__init__(cfg) 

31 

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 } 

69 

70 

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. 

75 

76 There are a bunch of dumb bugs in the original code, sorry! 

77 

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]). 

81 

82 8L has *just* a left facing W_pos, the rest right facing. 

83 

84 And some models were trained with 

85 """ 

86 # Early models have left facing W_pos 

87 reverse_pos = cfg.n_layers <= 8 

88 

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 

91 

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 

98 

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