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
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-21 19:27 +0000
1"""BERT architecture adapter.
3This module provides the architecture adapter for BERT models.
4"""
6from typing import Any
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)
26class BertArchitectureAdapter(ArchitectureAdapter):
27 """Architecture adapter for BERT models."""
29 supports_generation: bool = False
31 def __init__(self, cfg: Any) -> None:
32 """Initialize the BERT architecture adapter.
34 Args:
35 cfg: The configuration object.
36 """
37 super().__init__(cfg)
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
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
53 n_heads = self.cfg.n_heads
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 }
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 }
148 def prepare_model(self, hf_model: Any) -> None:
149 """Adjust component mapping based on the actual HF model variant.
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 )
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)