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

1"""BART adapter and the shared BART-family encoder-decoder base (BART, Marian, 

2MBart, Pegasus, Blenderbot, M2M100/NLLB); per-member differences are declarative.""" 

3 

4from typing import Any, Dict 

5 

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) 

20 

21 

22class BartFamilyArchitectureAdapter(ArchitectureAdapter): 

23 """Shared base for the BART-family encoder-decoder adapters.""" 

24 

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 

36 

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

38 """Validate the config, set family flags, and build the mapping.""" 

39 super().__init__(cfg) 

40 

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 ) 

49 

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 ) 

57 

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 ) 

65 

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 

77 

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 = {} 

84 

85 self.component_mapping = self._build_component_mapping() 

86 

87 def _norm(self, name: str) -> NormalizationBridge: 

88 return NormalizationBridge(name=name, config=self.cfg, use_native_layernorm_autograd=True) 

89 

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 ) 

102 

103 def _encoder_attention(self) -> AttentionBridge: 

104 """Encoder self-attention seam; LED swaps in its Longformer variant.""" 

105 return self._attention("self_attn") 

106 

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 ) 

121 

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 ) 

136 

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 ) 

158 

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 

178 

179 def setup_hook_compatibility(self, bridge: Any) -> None: 

180 """Fold the trained final_logits_bias into the unembed bias. 

181 

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 

188 

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_() 

197 

198 

199class BartArchitectureAdapter(BartFamilyArchitectureAdapter): 

200 """Architecture adapter for BartForConditionalGeneration models. 

201 

202 Post-LN with layernorm_embedding; checkpoints ship scale_embedding=False, 

203 so the family default-on is disabled. 

204 """ 

205 

206 force_scale_embedding = False 

207 has_layernorm_embedding = True