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