Coverage for transformer_lens/factories/architecture_adapter_factory.py: 96%
38 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"""Architecture adapter factory.
3This module provides a factory for creating architecture adapters, including
4support for external registration and entry-point discovery.
5"""
7import warnings
8from importlib.metadata import entry_points
10from transformer_lens.config import TransformerBridgeConfig
11from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
12from transformer_lens.model_bridge.supported_architectures import (
13 AfmoeArchitectureAdapter,
14 ApertusArchitectureAdapter,
15 ArceeArchitectureAdapter,
16 ASTArchitectureAdapter,
17 AudioFlamingo3ArchitectureAdapter,
18 BaichuanArchitectureAdapter,
19 BambaArchitectureAdapter,
20 BartArchitectureAdapter,
21 BD3LMArchitectureAdapter,
22 BertArchitectureAdapter,
23 BitNetArchitectureAdapter,
24 BlenderbotArchitectureAdapter,
25 BloomArchitectureAdapter,
26 CodeGenArchitectureAdapter,
27 Cohere2ArchitectureAdapter,
28 CohereArchitectureAdapter,
29 DeepSeekV2ArchitectureAdapter,
30 DeepSeekV3ArchitectureAdapter,
31 DeepSeekV4ArchitectureAdapter,
32 DreamArchitectureAdapter,
33 Emu3ArchitectureAdapter,
34 Ernie4_5_MoeArchitectureAdapter,
35 Ernie4_5ArchitectureAdapter,
36 Exaone4ArchitectureAdapter,
37 ExaoneArchitectureAdapter,
38 FalconArchitectureAdapter,
39 FalconH1ArchitectureAdapter,
40 FalconMambaArchitectureAdapter,
41 FlexOlmoArchitectureAdapter,
42 Florence2ArchitectureAdapter,
43 Gemma1ArchitectureAdapter,
44 Gemma2ArchitectureAdapter,
45 Gemma3ArchitectureAdapter,
46 Gemma3MultimodalArchitectureAdapter,
47 Gemma3nArchitectureAdapter,
48 Gemma4ArchitectureAdapter,
49 Gemma4TextArchitectureAdapter,
50 GiddArchitectureAdapter,
51 Glm4ArchitectureAdapter,
52 Glm4MoeArchitectureAdapter,
53 Glm4MoeLiteArchitectureAdapter,
54 Glm4vArchitectureAdapter,
55 GlmArchitectureAdapter,
56 GlmAsrArchitectureAdapter,
57 GlmMoeDsaArchitectureAdapter,
58 GPT2ArchitectureAdapter,
59 Gpt2LmHeadCustomArchitectureAdapter,
60 GPTBigCodeArchitectureAdapter,
61 GptjArchitectureAdapter,
62 GPTOSSArchitectureAdapter,
63 GraniteArchitectureAdapter,
64 GraniteMoeArchitectureAdapter,
65 GraniteMoeHybridArchitectureAdapter,
66 HrmTextArchitectureAdapter,
67 HubertArchitectureAdapter,
68 HunYuanDenseV1ArchitectureAdapter,
69 HyenaDNAArchitectureAdapter,
70 Idefics3ArchitectureAdapter,
71 InternLM2ArchitectureAdapter,
72 Jais2ArchitectureAdapter,
73 JambaArchitectureAdapter,
74 JetMoeArchitectureAdapter,
75 LagunaArchitectureAdapter,
76 LEDArchitectureAdapter,
77 Lfm2ArchitectureAdapter,
78 Lfm2MoeArchitectureAdapter,
79 LLaDA2MoeArchitectureAdapter,
80 LLaDAArchitectureAdapter,
81 Llama4ArchitectureAdapter,
82 Llama4MultimodalArchitectureAdapter,
83 LlamaArchitectureAdapter,
84 LlavaArchitectureAdapter,
85 LlavaNextArchitectureAdapter,
86 LlavaOnevisionArchitectureAdapter,
87 LongT5ArchitectureAdapter,
88 M2M100ArchitectureAdapter,
89 Mamba2ArchitectureAdapter,
90 MambaArchitectureAdapter,
91 MarianArchitectureAdapter,
92 MBartArchitectureAdapter,
93 MingptArchitectureAdapter,
94 MiniMaxM2ArchitectureAdapter,
95 Ministral3ArchitectureAdapter,
96 Mistral3ArchitectureAdapter,
97 MistralArchitectureAdapter,
98 MixtralArchitectureAdapter,
99 ModernBertDecoderArchitectureAdapter,
100 MPTArchitectureAdapter,
101 MusicFlamingoArchitectureAdapter,
102 NanoChatArchitectureAdapter,
103 NanogptArchitectureAdapter,
104 NativeArchitectureAdapter,
105 NeelSoluOldArchitectureAdapter,
106 NemotronArchitectureAdapter,
107 NemotronHArchitectureAdapter,
108 NeoArchitectureAdapter,
109 NeoxArchitectureAdapter,
110 Olmo2ArchitectureAdapter,
111 Olmo3ArchitectureAdapter,
112 OlmoArchitectureAdapter,
113 OlmoeArchitectureAdapter,
114 OlmoHybridArchitectureAdapter,
115 OpenAIGPTArchitectureAdapter,
116 OpenElmArchitectureAdapter,
117 OptArchitectureAdapter,
118 OuroArchitectureAdapter,
119 PegasusArchitectureAdapter,
120 Phi3ArchitectureAdapter,
121 PhiArchitectureAdapter,
122 PhiMoEArchitectureAdapter,
123 PretrainArchitectureAdapter,
124 Qwen2_5_VLArchitectureAdapter,
125 Qwen2ArchitectureAdapter,
126 Qwen2AudioArchitectureAdapter,
127 Qwen2MoeArchitectureAdapter,
128 Qwen3_5ArchitectureAdapter,
129 Qwen3_5MoeArchitectureAdapter,
130 Qwen3_5MoeMultimodalArchitectureAdapter,
131 Qwen3_5MultimodalArchitectureAdapter,
132 Qwen3ArchitectureAdapter,
133 Qwen3MoeArchitectureAdapter,
134 Qwen3NextArchitectureAdapter,
135 Qwen3VLArchitectureAdapter,
136 Qwen3VLMoeArchitectureAdapter,
137 QwenArchitectureAdapter,
138 RavenArchitectureAdapter,
139 RecurrentGemmaArchitectureAdapter,
140 RWKV7ArchitectureAdapter,
141 RwkvArchitectureAdapter,
142 SeedOssArchitectureAdapter,
143 SmolLM3ArchitectureAdapter,
144 StableLmArchitectureAdapter,
145 Starcoder2ArchitectureAdapter,
146 SwitchTransformersArchitectureAdapter,
147 T5ArchitectureAdapter,
148 T5Gemma2ArchitectureAdapter,
149 T5GemmaArchitectureAdapter,
150 VaultGemmaArchitectureAdapter,
151 ViTArchitectureAdapter,
152 XGLMArchitectureAdapter,
153 YoutuArchitectureAdapter,
154 Zamba2ArchitectureAdapter,
155)
157# Export supported architectures
158SUPPORTED_ARCHITECTURES = {
159 "AfmoeForCausalLM": AfmoeArchitectureAdapter,
160 "ApertusForCausalLM": ApertusArchitectureAdapter,
161 "ArceeForCausalLM": ArceeArchitectureAdapter,
162 "ASTForAudioClassification": ASTArchitectureAdapter,
163 "BaiChuanForCausalLM": BaichuanArchitectureAdapter,
164 "BaichuanForCausalLM": BaichuanArchitectureAdapter,
165 "BambaForCausalLM": BambaArchitectureAdapter,
166 "BartForConditionalGeneration": BartArchitectureAdapter,
167 "BD3LM": BD3LMArchitectureAdapter,
168 "DreamModel": DreamArchitectureAdapter,
169 "Emu3ForConditionalGeneration": Emu3ArchitectureAdapter,
170 "AudioFlamingo3ForConditionalGeneration": AudioFlamingo3ArchitectureAdapter,
171 "FlexOlmoForCausalLM": FlexOlmoArchitectureAdapter,
172 "GiddForDiffusionLM": GiddArchitectureAdapter,
173 "HyenaDNAForCausalLM": HyenaDNAArchitectureAdapter,
174 "LLaDA2MoeModelLM": LLaDA2MoeArchitectureAdapter,
175 "Jais2ForCausalLM": Jais2ArchitectureAdapter,
176 "JetMoeForCausalLM": JetMoeArchitectureAdapter,
177 # jetmoe-8b checkpoints predate the native port and use the remote-code capitalization
178 "JetMoEForCausalLM": JetMoeArchitectureAdapter,
179 "LagunaForCausalLM": LagunaArchitectureAdapter,
180 "Ministral3ForCausalLM": Ministral3ArchitectureAdapter,
181 "VaultGemmaForCausalLM": VaultGemmaArchitectureAdapter,
182 "YoutuForCausalLM": YoutuArchitectureAdapter,
183 "ModernBertDecoderForCausalLM": ModernBertDecoderArchitectureAdapter,
184 "MusicFlamingoForConditionalGeneration": MusicFlamingoArchitectureAdapter,
185 "NanoChatForCausalLM": NanoChatArchitectureAdapter,
186 "RwkvForCausalLM": RwkvArchitectureAdapter,
187 "SwitchTransformersForConditionalGeneration": SwitchTransformersArchitectureAdapter,
188 "BertForMaskedLM": BertArchitectureAdapter,
189 "BertLMHeadModel": BertArchitectureAdapter,
190 "BitNetForCausalLM": BitNetArchitectureAdapter,
191 "BlenderbotForConditionalGeneration": BlenderbotArchitectureAdapter,
192 "BloomForCausalLM": BloomArchitectureAdapter,
193 "BloomModel": BloomArchitectureAdapter,
194 "CodeGenForCausalLM": CodeGenArchitectureAdapter,
195 "Cohere2ForCausalLM": Cohere2ArchitectureAdapter,
196 "CohereForCausalLM": CohereArchitectureAdapter,
197 "DeepseekV2ForCausalLM": DeepSeekV2ArchitectureAdapter,
198 "DeepseekV3ForCausalLM": DeepSeekV3ArchitectureAdapter,
199 "Ernie4_5ForCausalLM": Ernie4_5ArchitectureAdapter,
200 "Ernie4_5_MoeForCausalLM": Ernie4_5_MoeArchitectureAdapter,
201 "ExaoneForCausalLM": ExaoneArchitectureAdapter,
202 "Exaone4ForCausalLM": Exaone4ArchitectureAdapter,
203 "DeepseekV4ForCausalLM": DeepSeekV4ArchitectureAdapter,
204 "FalconForCausalLM": FalconArchitectureAdapter,
205 "FalconH1ForCausalLM": FalconH1ArchitectureAdapter,
206 "FalconMambaForCausalLM": FalconMambaArchitectureAdapter,
207 "Florence2ForConditionalGeneration": Florence2ArchitectureAdapter,
208 "GemmaForCausalLM": Gemma1ArchitectureAdapter, # Default to Gemma1 as it's the original version
209 "Gemma1ForCausalLM": Gemma1ArchitectureAdapter,
210 "Gemma2ForCausalLM": Gemma2ArchitectureAdapter,
211 "Gemma3ForCausalLM": Gemma3ArchitectureAdapter,
212 "Gemma3ForConditionalGeneration": Gemma3MultimodalArchitectureAdapter,
213 "Gemma3nForConditionalGeneration": Gemma3nArchitectureAdapter,
214 "Gemma4ForConditionalGeneration": Gemma4ArchitectureAdapter,
215 # The unified (encoder-free) 12B variant's text decoder is a strict structural
216 # subset of gemma4 (no PLE, no MoE — both optional in the adapter), with the
217 # same module paths. Requires transformers >= 5.10 to load.
218 "Gemma4UnifiedForConditionalGeneration": Gemma4ArchitectureAdapter,
219 "Gemma4ForCausalLM": Gemma4TextArchitectureAdapter,
220 "GraniteForCausalLM": GraniteArchitectureAdapter,
221 "GraniteMoeForCausalLM": GraniteMoeArchitectureAdapter,
222 "GraniteMoeHybridForCausalLM": GraniteMoeHybridArchitectureAdapter,
223 "GlmForCausalLM": GlmArchitectureAdapter,
224 "Glm4ForCausalLM": Glm4ArchitectureAdapter,
225 "Glm4vForConditionalGeneration": Glm4vArchitectureAdapter,
226 "GlmMoeDsaForCausalLM": GlmMoeDsaArchitectureAdapter,
227 "Glm4MoeForCausalLM": Glm4MoeArchitectureAdapter,
228 "Glm4MoeLiteForCausalLM": Glm4MoeLiteArchitectureAdapter,
229 "GlmAsrForConditionalGeneration": GlmAsrArchitectureAdapter,
230 "GPT2LMHeadModel": GPT2ArchitectureAdapter,
231 "GPTBigCodeForCausalLM": GPTBigCodeArchitectureAdapter,
232 "GptOssForCausalLM": GPTOSSArchitectureAdapter,
233 "GPT2LMHeadCustomModel": Gpt2LmHeadCustomArchitectureAdapter,
234 "GPTJForCausalLM": GptjArchitectureAdapter,
235 "HrmTextForCausalLM": HrmTextArchitectureAdapter,
236 "HubertForCTC": HubertArchitectureAdapter,
237 "HubertModel": HubertArchitectureAdapter,
238 "HunYuanDenseV1ForCausalLM": HunYuanDenseV1ArchitectureAdapter,
239 "Idefics3ForConditionalGeneration": Idefics3ArchitectureAdapter,
240 "InternLM2ForCausalLM": InternLM2ArchitectureAdapter,
241 "JambaForCausalLM": JambaArchitectureAdapter,
242 "LEDForConditionalGeneration": LEDArchitectureAdapter,
243 "LLaDAModelLM": LLaDAArchitectureAdapter,
244 "LlamaForCausalLM": LlamaArchitectureAdapter,
245 "Llama4ForCausalLM": Llama4ArchitectureAdapter,
246 "Llama4ForConditionalGeneration": Llama4MultimodalArchitectureAdapter,
247 "LlavaForConditionalGeneration": LlavaArchitectureAdapter,
248 "LlavaNextForConditionalGeneration": LlavaNextArchitectureAdapter,
249 "LlavaOnevisionForConditionalGeneration": LlavaOnevisionArchitectureAdapter,
250 "Lfm2ForCausalLM": Lfm2ArchitectureAdapter,
251 "Lfm2MoeForCausalLM": Lfm2MoeArchitectureAdapter,
252 "LongT5ForConditionalGeneration": LongT5ArchitectureAdapter,
253 "M2M100ForConditionalGeneration": M2M100ArchitectureAdapter,
254 "Mamba2ForCausalLM": Mamba2ArchitectureAdapter,
255 "MambaForCausalLM": MambaArchitectureAdapter,
256 "MarianMTModel": MarianArchitectureAdapter,
257 "MBartForConditionalGeneration": MBartArchitectureAdapter,
258 "MiniMaxM2ForCausalLM": MiniMaxM2ArchitectureAdapter,
259 "NemotronForCausalLM": NemotronArchitectureAdapter,
260 "NemotronHForCausalLM": NemotronHArchitectureAdapter,
261 "MixtralForCausalLM": MixtralArchitectureAdapter,
262 "MistralForCausalLM": MistralArchitectureAdapter,
263 "Mistral3ForConditionalGeneration": Mistral3ArchitectureAdapter,
264 "MPTForCausalLM": MPTArchitectureAdapter,
265 "MptForCausalLM": MPTArchitectureAdapter,
266 "NeoForCausalLM": NeoArchitectureAdapter,
267 "NeoXForCausalLM": NeoxArchitectureAdapter,
268 "NeelSoluOldForCausalLM": NeelSoluOldArchitectureAdapter,
269 "OlmoForCausalLM": OlmoArchitectureAdapter,
270 "Olmo2ForCausalLM": Olmo2ArchitectureAdapter,
271 "Olmo3ForCausalLM": Olmo3ArchitectureAdapter,
272 "OlmoeForCausalLM": OlmoeArchitectureAdapter,
273 "OlmoHybridForCausalLM": OlmoHybridArchitectureAdapter,
274 "OpenAIGPTLMHeadModel": OpenAIGPTArchitectureAdapter,
275 "OpenELMForCausalLM": OpenElmArchitectureAdapter,
276 "OPTForCausalLM": OptArchitectureAdapter,
277 "PegasusForConditionalGeneration": PegasusArchitectureAdapter,
278 "OuroForCausalLM": OuroArchitectureAdapter,
279 "PhiForCausalLM": PhiArchitectureAdapter,
280 "Phi3ForCausalLM": Phi3ArchitectureAdapter,
281 "PhiMoEForCausalLM": PhiMoEArchitectureAdapter,
282 "QwenForCausalLM": QwenArchitectureAdapter,
283 "Qwen2ForCausalLM": Qwen2ArchitectureAdapter,
284 "Qwen2_5_VLForConditionalGeneration": Qwen2_5_VLArchitectureAdapter,
285 "Qwen2AudioForConditionalGeneration": Qwen2AudioArchitectureAdapter,
286 "Qwen2MoeForCausalLM": Qwen2MoeArchitectureAdapter,
287 "Qwen3ForCausalLM": Qwen3ArchitectureAdapter,
288 "Qwen3VLForConditionalGeneration": Qwen3VLArchitectureAdapter,
289 "Qwen3VLMoeForConditionalGeneration": Qwen3VLMoeArchitectureAdapter,
290 "Qwen3MoeForCausalLM": Qwen3MoeArchitectureAdapter,
291 "Qwen3NextForCausalLM": Qwen3NextArchitectureAdapter,
292 "Qwen3_5ForCausalLM": Qwen3_5ArchitectureAdapter,
293 "Qwen3_5ForConditionalGeneration": Qwen3_5MultimodalArchitectureAdapter,
294 "Qwen3_5MoeForCausalLM": Qwen3_5MoeArchitectureAdapter,
295 "Qwen3_5MoeForConditionalGeneration": Qwen3_5MoeMultimodalArchitectureAdapter,
296 "RavenForCausalLM": RavenArchitectureAdapter,
297 "RecurrentGemmaForCausalLM": RecurrentGemmaArchitectureAdapter,
298 "SeedOssForCausalLM": SeedOssArchitectureAdapter,
299 "RWKV7ForCausalLM": RWKV7ArchitectureAdapter,
300 "SmolLM3ForCausalLM": SmolLM3ArchitectureAdapter,
301 "StableLmForCausalLM": StableLmArchitectureAdapter,
302 "Starcoder2ForCausalLM": Starcoder2ArchitectureAdapter,
303 "T5ForConditionalGeneration": T5ArchitectureAdapter,
304 "MT5ForConditionalGeneration": T5ArchitectureAdapter,
305 "T5WithLMHeadModel": T5ArchitectureAdapter,
306 "T5GemmaForConditionalGeneration": T5GemmaArchitectureAdapter,
307 "T5Gemma2ForConditionalGeneration": T5Gemma2ArchitectureAdapter,
308 "XGLMForCausalLM": XGLMArchitectureAdapter,
309 "Zamba2ForCausalLM": Zamba2ArchitectureAdapter,
310 "NanoGPTForCausalLM": NanogptArchitectureAdapter,
311 "TransformerLensNative": NativeArchitectureAdapter,
312 "TransformerLensPretrain": PretrainArchitectureAdapter,
313 "MinGPTForCausalLM": MingptArchitectureAdapter,
314 "GPTNeoForCausalLM": NeoArchitectureAdapter,
315 "GPTNeoXForCausalLM": NeoxArchitectureAdapter,
316 "ViTModel": ViTArchitectureAdapter,
317 "ViTForImageClassification": ViTArchitectureAdapter,
318 "DeiTModel": ViTArchitectureAdapter,
319 "DeiTForImageClassification": ViTArchitectureAdapter, # "DeiTForImageClassificationWithTeacher" unsupported for now
320}
323class ArchitectureAdapterFactory:
324 """Factory for creating architecture adapters.
326 Supports external registration via `register_adapter()` and automatic
327 discovery of adapters from installed packages via entry points.
328 """
330 _adapters = dict(SUPPORTED_ARCHITECTURES)
331 _entry_points_discovered = False
333 @classmethod
334 def register_adapter(
335 cls, architecture_name: str, adapter_class: type["ArchitectureAdapter"]
336 ) -> None:
337 """Register a custom architecture adapter at runtime.
339 This allows users to add their own architecture adapters without
340 modifying TransformerLens source code.
342 Args:
343 architecture_name: The HuggingFace architecture class name
344 (e.g. ``"Qwen3ForCausalLM"``).
345 adapter_class: The adapter class to register.
347 Example:
348 >>> from transformer_lens.config import TransformerBridgeConfig
349 >>> from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
350 >>> from transformer_lens.factories.architecture_adapter_factory import ArchitectureAdapterFactory
351 >>> class MyAdapter(ArchitectureAdapter):
352 ... def __init__(self, cfg):
353 ... super().__init__(cfg)
354 >>> ArchitectureAdapterFactory.register_adapter("MyModelForCausalLM", MyAdapter)
355 >>> cfg = TransformerBridgeConfig(
356 ... d_model=512, d_head=64, n_layers=6, n_ctx=1024,
357 ... architecture="MyModelForCausalLM",
358 ... )
359 >>> adapter = ArchitectureAdapterFactory.select_architecture_adapter(cfg)
360 >>> isinstance(adapter, MyAdapter)
361 True
362 """
363 cls._adapters[architecture_name] = adapter_class
365 @classmethod
366 def discover_entry_points(cls) -> None:
367 """Discover and register architecture adapters from installed packages.
369 Packages can declare adapters in their ``pyproject.toml``:
370 ```toml
371 [project.entry-points."transformer_lens.architectures"]
372 "MyModelForCausalLM" = "my_package.adapters:MyArchitectureAdapter"
373 ```
374 """
375 if cls._entry_points_discovered:
376 return
377 try:
378 eps = entry_points(group="transformer_lens.architectures")
379 except Exception as e:
380 warnings.warn(
381 f"Failed to discover entry points: {e}. " f"External adapters may not be available."
382 )
383 else:
384 for ep in eps:
385 try:
386 if ep.name in cls._adapters:
387 dist_name = (
388 getattr(ep.dist, "name", "unknown")
389 if ep.dist is not None
390 else "unknown"
391 )
392 warnings.warn(
393 f"Custom architecture adapter {ep.name} provided by {dist_name} "
394 f"attempted to override a native adapter. If you'd like to use this "
395 f"custom adapter, register it explicitly with register_adapter"
396 )
397 continue
398 cls._adapters[ep.name] = ep.load()
399 except Exception as e:
400 warnings.warn(
401 f"Failed to load entry point '{ep.name}': {e}. " f"Skipping this adapter."
402 )
403 cls._entry_points_discovered = True
405 @classmethod
406 def select_architecture_adapter(cls, cfg: TransformerBridgeConfig) -> ArchitectureAdapter:
407 """Select the appropriate architecture adapter for the given config.
409 Args:
410 cfg: The TransformerBridgeConfig to select the adapter for.
412 Returns:
413 The selected architecture adapter.
415 Raises:
416 ValueError: If no adapter is found for the given config.
417 """
418 cls.discover_entry_points()
419 if cfg.architecture is not None:
420 if cfg.architecture in cls._adapters:
421 return cls._adapters[cfg.architecture](cfg)
422 else:
423 raise ValueError(f"Unsupported architecture: {cfg.architecture}")
425 raise ValueError(f"TransformerBridgeConfig must have architecture set, got: {cfg}")