Coverage for transformer_lens/model_bridge/supported_architectures/bert.py: 100%

29 statements  

« prev     ^ index     » next       coverage.py v7.10.1, created at 2026-09-21 19:27 +0000

1"""BERT architecture adapter. 

2 

3This module provides the architecture adapter for BERT models. 

4""" 

5 

6from typing import Any 

7 

8from transformer_lens.conversion_utils.conversion_steps import RearrangeTensorConversion 

9from transformer_lens.conversion_utils.param_processing_conversion import ( 

10 ParamProcessingConversion, 

11) 

12from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter 

13from transformer_lens.model_bridge.generalized_components import ( 

14 AttentionBridge, 

15 BertPoolerBridge, 

16 BlockBridge, 

17 EmbeddingBridge, 

18 LinearBridge, 

19 MLPBridge, 

20 NormalizationBridge, 

21 PosEmbedBridge, 

22 UnembeddingBridge, 

23) 

24 

25 

26class BertArchitectureAdapter(ArchitectureAdapter): 

27 """Architecture adapter for BERT models.""" 

28 

29 supports_generation: bool = False 

30 

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

32 """Initialize the BERT architecture adapter. 

33 

34 Args: 

35 cfg: The configuration object. 

36 """ 

37 super().__init__(cfg) 

38 

39 # Set config variables for weight processing 

40 self.cfg.normalization_type = "LN" 

41 self.cfg.positional_embedding_type = "standard" 

42 self.cfg.final_rms = False 

43 self.cfg.gated_mlp = False 

44 self.cfg.attn_only = False 

45 

46 # BERT uses post-LN (LayerNorm after residual, not before sublayer). 

47 # fold_ln assumes pre-LN (LN before sublayer) and folds ln1 into attention 

48 # QKV and ln2 into MLP. For post-LN, ln1 output feeds MLP (not attention) 

49 # and ln2 output feeds next block's attention (not MLP), so folding into 

50 # the wrong sublayer produces incorrect results. 

51 self.supports_fold_ln = False 

52 

53 n_heads = self.cfg.n_heads 

54 

55 self.weight_processing_conversions = { 

56 "blocks.{i}.attn.q.weight": ParamProcessingConversion( 

57 tensor_conversion=RearrangeTensorConversion( 

58 "(h d_head) d_model -> h d_model d_head", h=n_heads 

59 ), 

60 ), 

61 "blocks.{i}.attn.k.weight": ParamProcessingConversion( 

62 tensor_conversion=RearrangeTensorConversion( 

63 "(h d_head) d_model -> h d_model d_head", h=n_heads 

64 ), 

65 ), 

66 "blocks.{i}.attn.v.weight": ParamProcessingConversion( 

67 tensor_conversion=RearrangeTensorConversion( 

68 "(h d_head) d_model -> h d_model d_head", h=n_heads 

69 ), 

70 ), 

71 "blocks.{i}.attn.q.bias": ParamProcessingConversion( 

72 tensor_conversion=RearrangeTensorConversion("(h d_head) -> h d_head", h=n_heads), 

73 ), 

74 "blocks.{i}.attn.k.bias": ParamProcessingConversion( 

75 tensor_conversion=RearrangeTensorConversion("(h d_head) -> h d_head", h=n_heads), 

76 ), 

77 "blocks.{i}.attn.v.bias": ParamProcessingConversion( 

78 tensor_conversion=RearrangeTensorConversion("(h d_head) -> h d_head", h=n_heads), 

79 ), 

80 "blocks.{i}.attn.o.weight": ParamProcessingConversion( 

81 tensor_conversion=RearrangeTensorConversion( 

82 "d_model (h d_head) -> h d_head d_model", h=n_heads 

83 ), 

84 ), 

85 } 

86 

87 # Set up component mapping 

88 # MLM defaults; prepare_model() adjusts for other task heads (e.g., NSP). 

89 self.component_mapping = { 

90 "embed": EmbeddingBridge(name="bert.embeddings.word_embeddings"), 

91 "token_type_embed": EmbeddingBridge(name="bert.embeddings.token_type_embeddings"), 

92 "pos_embed": PosEmbedBridge(name="bert.embeddings.position_embeddings"), 

93 "embed_ln": NormalizationBridge( 

94 name="bert.embeddings.LayerNorm", 

95 config=self.cfg, 

96 use_native_layernorm_autograd=True, 

97 ), 

98 "blocks": BlockBridge( 

99 name="bert.encoder.layer", 

100 # BERT has no single MLP module (intermediate.dense and output.dense 

101 # are siblings in BertLayer), so the MLPBridge forward is never called 

102 # and mlp.hook_out never fires. Redirect hook_mlp_out to the actual 

103 # MLP output hook (output of the "out" linear layer). 

104 hook_alias_overrides={ 

105 "hook_mlp_out": "mlp.out.hook_out", 

106 "hook_mlp_in": "mlp.in.hook_in", 

107 }, 

108 submodules={ 

109 "ln1": NormalizationBridge( 

110 name="attention.output.LayerNorm", 

111 config=self.cfg, 

112 use_native_layernorm_autograd=True, 

113 ), 

114 "ln2": NormalizationBridge( 

115 name="output.LayerNorm", 

116 config=self.cfg, 

117 use_native_layernorm_autograd=True, 

118 ), 

119 "attn": AttentionBridge( 

120 name="attention", 

121 config=self.cfg, 

122 submodules={ 

123 "q": LinearBridge(name="self.query"), 

124 "k": LinearBridge(name="self.key"), 

125 "v": LinearBridge(name="self.value"), 

126 "o": LinearBridge(name="output.dense"), 

127 }, 

128 ), 

129 "mlp": MLPBridge( 

130 name=None, 

131 config=self.cfg, 

132 submodules={ 

133 "in": LinearBridge(name="intermediate.dense"), 

134 "out": LinearBridge(name="output.dense"), 

135 }, 

136 ), 

137 }, 

138 ), 

139 "mlm_head": LinearBridge(name="cls.predictions.transform.dense"), 

140 "unembed": UnembeddingBridge(name="cls.predictions.decoder"), 

141 "ln_final": NormalizationBridge( 

142 name="cls.predictions.transform.LayerNorm", 

143 config=self.cfg, 

144 use_native_layernorm_autograd=True, 

145 ), 

146 } 

147 

148 def prepare_model(self, hf_model: Any) -> None: 

149 """Adjust component mapping based on the actual HF model variant. 

150 

151 BertForMaskedLM has cls.predictions (MLM head). 

152 BertForNextSentencePrediction has cls.seq_relationship (NSP head) 

153 and no MLM-specific LayerNorm. 

154 """ 

155 if getattr(getattr(hf_model, "bert", None), "pooler", None) is not None: 

156 # Wrap the pooler itself, not its inner Linear: HF applies tanh after 

157 # the projection, so hook_out here is the pooled [CLS] that 

158 # HookedEncoder's BertPooler exposes as hook_pooler_out. The dense 

159 # stays hookable as a submodule for the pre-activation projection. 

160 self.components["pooler"] = BertPoolerBridge( 

161 name="bert.pooler", 

162 submodules={"dense": LinearBridge(name="dense")}, 

163 ) 

164 

165 has_predictions = hasattr(getattr(hf_model, "cls", None), "predictions") 

166 has_nsp_head = hasattr(getattr(hf_model, "cls", None), "seq_relationship") 

167 if has_nsp_head and has_predictions: 

168 self.components["nsp_head"] = LinearBridge(name="cls.seq_relationship") 

169 elif has_nsp_head: 

170 # NSP-only model — swap head components. 

171 self.components["unembed"] = UnembeddingBridge(name="cls.seq_relationship") 

172 self.components.pop("mlm_head", None) 

173 self.components.pop("ln_final", None)