Coverage for transformer_lens/model_bridge/supported_architectures/lfm2.py: 91%
35 statements
« prev ^ index » next coverage.py v7.10.1, created at 2026-08-11 18:50 +0000
« prev ^ index » next coverage.py v7.10.1, created at 2026-08-11 18:50 +0000
1"""Lfm2 architecture adapter."""
3from typing import Any
5from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
6from transformer_lens.model_bridge.generalized_components import (
7 BlockBridge,
8 DepthwiseConv1DBridge,
9 EmbeddingBridge,
10 GatedMLPBridge,
11 Lfm2ShortConvBridge,
12 LinearBridge,
13 PositionEmbeddingsAttentionBridge,
14 RMSNormalizationBridge,
15 RotaryEmbeddingBridge,
16 UnembeddingBridge,
17)
20class Lfm2ArchitectureAdapter(ArchitectureAdapter):
21 """Architecture adapter for Lfm2 models."""
23 def __init__(self, cfg: Any) -> None:
24 """Initialize the Lfm2 architecture adapter."""
25 super().__init__(cfg)
27 self.cfg.normalization_type = "RMS"
28 self.cfg.positional_embedding_type = "rotary"
29 self.cfg.final_rms = True
30 self.cfg.gated_mlp = True
31 self.cfg.attn_only = False
32 self.cfg.uses_rms_norm = True
33 self.cfg.act_fn = "silu"
35 self.cfg.attn_implementation = "eager"
37 if hasattr(cfg, "n_key_value_heads") and cfg.n_key_value_heads is not None: 37 ↛ 40line 37 didn't jump to line 40 because the condition on line 37 was always true
38 self.cfg.n_key_value_heads = cfg.n_key_value_heads
40 self.weight_processing_conversions = {
41 **self._qkvo_weight_conversions(),
42 }
44 self.component_mapping = {
45 "embed": EmbeddingBridge(name="model.embed_tokens"),
46 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb"),
47 "blocks": BlockBridge(
48 name="model.layers",
49 submodules={
50 "ln1": RMSNormalizationBridge(
51 name="operator_norm",
52 config=self.cfg,
53 ),
54 "ln2": RMSNormalizationBridge(
55 name="ffn_norm",
56 config=self.cfg,
57 ),
58 "attn": PositionEmbeddingsAttentionBridge(
59 name="self_attn",
60 config=self.cfg,
61 optional=True,
62 submodules={
63 "q": LinearBridge(name="q_proj"),
64 "k": LinearBridge(name="k_proj"),
65 "v": LinearBridge(name="v_proj"),
66 "o": LinearBridge(name="out_proj"),
67 "q_norm": RMSNormalizationBridge(name="q_layernorm", config=self.cfg),
68 "k_norm": RMSNormalizationBridge(name="k_layernorm", config=self.cfg),
69 },
70 requires_attention_mask=True,
71 requires_position_embeddings=True,
72 ),
73 "conv": Lfm2ShortConvBridge(
74 name="conv",
75 config=self.cfg,
76 optional=True,
77 submodules={
78 "in": LinearBridge(name="in_proj"),
79 "conv": DepthwiseConv1DBridge(name="conv"),
80 "out": LinearBridge(name="out_proj"),
81 },
82 ),
83 "mlp": GatedMLPBridge(
84 name="feed_forward",
85 config=self.cfg,
86 submodules={
87 "gate": LinearBridge(name="w1"),
88 "in": LinearBridge(name="w3"),
89 "out": LinearBridge(name="w2"),
90 },
91 ),
92 },
93 ),
94 "ln_final": RMSNormalizationBridge(name="model.embedding_norm", config=self.cfg),
95 "unembed": UnembeddingBridge(name="lm_head", config=self.cfg),
96 }
98 def setup_component_testing(self, hf_model: Any, bridge_model: Any = None) -> None:
99 """Set up model-specific references for component testing."""
100 rotary_emb = hf_model.model.rotary_emb
102 # Set attention implementation on HF model to eager (vs sdpa default)
103 if hasattr(hf_model, "config") and hasattr(hf_model.config, "_attn_implementation"): 103 ↛ 106line 103 didn't jump to line 106 because the condition on line 103 was always true
104 hf_model.config._attn_implementation = "eager"
106 if hasattr(hf_model, "model") and hasattr(hf_model.model, "layers"): 106 ↛ 112line 106 didn't jump to line 112 because the condition on line 106 was always true
107 for layer in hf_model.model.layers:
108 if hasattr(layer, "self_attn") and hasattr(layer.self_attn, "config"): 108 ↛ 107line 108 didn't jump to line 107 because the condition on line 108 was always true
109 layer.self_attn.config._attn_implementation = "eager"
111 # Set rotary_emb on actual bridge instances
112 if bridge_model is not None and hasattr(bridge_model, "blocks"):
113 for block in bridge_model.blocks:
114 if hasattr(block, "attn"):
115 block.attn.set_rotary_emb(rotary_emb)
117 # Set on template for get_generalized_component() calls
118 # Find the first attention layer (LFM2 layer 0 is conv, not attn)
119 layer_types = getattr(self.cfg, "layer_types", None)
120 if layer_types is not None and "full_attention" in layer_types: 120 ↛ exitline 120 didn't return from function 'setup_component_testing' because the condition on line 120 was always true
121 first_attn_idx = layer_types.index("full_attention")
122 attn_bridge = self.get_generalized_component(f"blocks.{first_attn_idx}.attn")
123 attn_bridge.set_rotary_emb(rotary_emb)