Coverage for transformer_lens/model_bridge/supported_architectures/bart.py: 100%
81 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"""BART adapter and the shared BART-family encoder-decoder base (BART, Marian,
2MBart, Pegasus, Blenderbot, M2M100/NLLB); per-member differences are declarative."""
4from typing import Any, Dict
6from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
7from transformer_lens.model_bridge.generalized_components import (
8 AttentionBridge,
9 BlockBridge,
10 EmbeddingBridge,
11 LinearBridge,
12 MLPBridge,
13 NormalizationBridge,
14 PosEmbedBridge,
15 UnembeddingBridge,
16)
17from transformer_lens.model_bridge.generalized_components.base import (
18 GeneralizedComponent,
19)
22class BartFamilyArchitectureAdapter(ArchitectureAdapter):
23 """Shared base for the BART-family encoder-decoder adapters."""
25 # Blenderbot ships asymmetric stacks (e.g. 2 encoder / 24 decoder layers) and
26 # follows the decoder; every other member requires symmetric stacks.
27 require_symmetric_layers: bool = True
28 n_layers_from: str = "encoder"
29 # BART checkpoints don't scale embeddings; the rest default scale_embedding on.
30 force_scale_embedding: bool = True
31 # layernorm_embedding after the token+position embeds (BART, MBart).
32 has_layernorm_embedding: bool = False
33 # Trailing per-stack layer_norm — the pre-LN members (MBart, Pegasus,
34 # Blenderbot, M2M100).
35 has_final_stack_norm: bool = False
37 def __init__(self, cfg: Any) -> None:
38 """Validate the config, set family flags, and build the mapping."""
39 super().__init__(cfg)
41 name = type(self).__name__
42 encoder_layers = getattr(self.cfg, "encoder_layers", self.cfg.n_layers)
43 decoder_layers = getattr(self.cfg, "decoder_layers", self.cfg.n_layers)
44 if self.require_symmetric_layers and encoder_layers != decoder_layers:
45 raise ValueError(
46 f"{name} only supports symmetric configs for now: "
47 f"encoder_layers={encoder_layers}, decoder_layers={decoder_layers}."
48 )
50 encoder_heads = getattr(self.cfg, "encoder_attention_heads", self.cfg.n_heads)
51 decoder_heads = getattr(self.cfg, "decoder_attention_heads", self.cfg.n_heads)
52 if encoder_heads != decoder_heads:
53 raise ValueError(
54 f"{name} only supports symmetric attention heads for now: "
55 f"encoder_attention_heads={encoder_heads}, decoder_attention_heads={decoder_heads}."
56 )
58 encoder_ffn_dim = getattr(self.cfg, "encoder_ffn_dim", self.cfg.d_mlp)
59 decoder_ffn_dim = getattr(self.cfg, "decoder_ffn_dim", self.cfg.d_mlp)
60 if encoder_ffn_dim != decoder_ffn_dim:
61 raise ValueError(
62 f"{name} only supports symmetric FFN dims for now: "
63 f"encoder_ffn_dim={encoder_ffn_dim}, decoder_ffn_dim={decoder_ffn_dim}."
64 )
66 self.cfg.n_layers = decoder_layers if self.n_layers_from == "decoder" else encoder_layers
67 self.cfg.n_heads = encoder_heads
68 self.cfg.d_head = self.cfg.d_model // encoder_heads
69 self.cfg.d_mlp = encoder_ffn_dim
70 self.cfg.normalization_type = "LN"
71 self.cfg.positional_embedding_type = "standard"
72 self.cfg.final_rms = False
73 self.cfg.gated_mlp = False
74 self.cfg.attn_only = False
75 if self.force_scale_embedding and self.cfg.scale_embedding is None:
76 self.cfg.scale_embedding = True
78 # Post-LN members break fold-LN's pre-LN assumption; pre-LN members keep
79 # the family-wide conservative default (per-stack final norms + embed
80 # scaling sit outside the folding machinery).
81 self.supports_fold_ln = False
82 self.supports_center_writing_weights = False
83 self.weight_processing_conversions = {}
85 self.component_mapping = self._build_component_mapping()
87 def _norm(self, name: str) -> NormalizationBridge:
88 return NormalizationBridge(name=name, config=self.cfg, use_native_layernorm_autograd=True)
90 def _attention(self, name: str, *, is_cross_attention: bool = False) -> AttentionBridge:
91 return AttentionBridge(
92 name=name,
93 config=self.cfg,
94 submodules={
95 "q": LinearBridge(name="q_proj"),
96 "k": LinearBridge(name="k_proj"),
97 "v": LinearBridge(name="v_proj"),
98 "o": LinearBridge(name="out_proj"),
99 },
100 is_cross_attention=is_cross_attention,
101 )
103 def _encoder_attention(self) -> AttentionBridge:
104 """Encoder self-attention seam; LED swaps in its Longformer variant."""
105 return self._attention("self_attn")
107 def _mlp(self) -> MLPBridge:
108 """MLPBridge(name=None), not SymbolicBridge: the latter exposes no
109 hook_pre/hook_post. fc1/fc2 sit directly on the block with no MLP
110 container, and component setup promotes on `name is None` rather than
111 on the bridge type, so the mlp.in/mlp.out weight paths are unchanged
112 and MLPBridge.forward is never invoked. Same shape as BERT."""
113 return MLPBridge(
114 name=None,
115 config=self.cfg,
116 submodules={
117 "in": LinearBridge(name="fc1"),
118 "out": LinearBridge(name="fc2"),
119 },
120 )
122 def _encoder_block(self) -> BlockBridge:
123 return BlockBridge(
124 name="model.encoder.layers",
125 hook_alias_overrides={
126 "hook_mlp_in": "mlp.in.hook_in",
127 "hook_mlp_out": "mlp.out.hook_out",
128 },
129 submodules={
130 "attn": self._encoder_attention(),
131 "ln1": self._norm("self_attn_layer_norm"),
132 "ln2": self._norm("final_layer_norm"),
133 "mlp": self._mlp(),
134 },
135 )
137 def _decoder_block(self) -> BlockBridge:
138 return BlockBridge(
139 name="model.decoder.layers",
140 hook_alias_overrides={
141 "hook_attn_in": "self_attn.hook_attn_in",
142 "hook_attn_out": "self_attn.hook_out",
143 "hook_q_input": "self_attn.hook_q_input",
144 "hook_k_input": "self_attn.hook_k_input",
145 "hook_v_input": "self_attn.hook_v_input",
146 "hook_mlp_in": "mlp.in.hook_in",
147 "hook_mlp_out": "mlp.out.hook_out",
148 },
149 submodules={
150 "self_attn": self._attention("self_attn"),
151 "ln1": self._norm("self_attn_layer_norm"),
152 "cross_attn": self._attention("encoder_attn", is_cross_attention=True),
153 "ln2": self._norm("encoder_attn_layer_norm"),
154 "ln3": self._norm("final_layer_norm"),
155 "mlp": self._mlp(),
156 },
157 )
159 def _build_component_mapping(self) -> Dict[str, GeneralizedComponent]:
160 mapping: Dict[str, GeneralizedComponent] = {
161 "embed": EmbeddingBridge(name="model.encoder.embed_tokens"),
162 "pos_embed": PosEmbedBridge(name="model.encoder.embed_positions"),
163 }
164 if self.has_layernorm_embedding:
165 mapping["embed_ln"] = self._norm("model.encoder.layernorm_embedding")
166 mapping["encoder_blocks"] = self._encoder_block()
167 if self.has_final_stack_norm:
168 mapping["encoder_ln_final"] = self._norm("model.encoder.layer_norm")
169 mapping["decoder_embed"] = EmbeddingBridge(name="model.decoder.embed_tokens")
170 mapping["decoder_pos_embed"] = PosEmbedBridge(name="model.decoder.embed_positions")
171 if self.has_layernorm_embedding:
172 mapping["decoder_embed_ln"] = self._norm("model.decoder.layernorm_embedding")
173 mapping["decoder_blocks"] = self._decoder_block()
174 if self.has_final_stack_norm:
175 mapping["decoder_ln_final"] = self._norm("model.decoder.layer_norm")
176 mapping["unembed"] = UnembeddingBridge(name="lm_head")
177 return mapping
179 def setup_hook_compatibility(self, bridge: Any) -> None:
180 """Fold the trained final_logits_bias into the unembed bias.
182 HF adds the buffer after lm_head, so b_U would read fabricated zeros
183 and unembed.hook_out would fire pre-bias (Marian opus-mt trains it).
184 Moving it into the bias UnembeddingBridge injects is numerically
185 identity; zeroing the buffer keeps the (re-run) fold idempotent.
186 """
187 import torch
189 model = getattr(bridge, "original_model", None)
190 buf = getattr(model, "final_logits_bias", None)
191 lm_head = getattr(model, "lm_head", None)
192 if buf is None or lm_head is None or getattr(lm_head, "bias", None) is None:
193 return
194 with torch.no_grad():
195 lm_head.bias.add_(buf.reshape(-1).to(lm_head.bias.dtype))
196 buf.zero_()
199class BartArchitectureAdapter(BartFamilyArchitectureAdapter):
200 """Architecture adapter for BartForConditionalGeneration models.
202 Post-LN with layernorm_embedding; checkpoints ship scale_embedding=False,
203 so the family default-on is disabled.
204 """
206 force_scale_embedding = False
207 has_layernorm_embedding = True