Coverage for transformer_lens/model_bridge/supported_architectures/lfm2.py: 94%
26 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"""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 Lfm2ShortConvBridge,
11 LinearBridge,
12 PositionEmbeddingsAttentionBridge,
13 RMSNormalizationBridge,
14 RotaryEmbeddingBridge,
15 UnembeddingBridge,
16)
17from transformer_lens.utilities.attn_implementation import force_eager_attention
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._set_rms_rotary_defaults()
28 self.cfg.act_fn = "silu"
30 self.cfg.attn_implementation = "eager"
32 if hasattr(cfg, "n_key_value_heads") and cfg.n_key_value_heads is not None: 32 ↛ 35line 32 didn't jump to line 35 because the condition on line 32 was always true
33 self.cfg.n_key_value_heads = cfg.n_key_value_heads
35 self.weight_processing_conversions = {
36 **self._qkvo_weight_conversions(),
37 }
39 self.component_mapping = {
40 "embed": EmbeddingBridge(name="model.embed_tokens"),
41 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb"),
42 "blocks": BlockBridge(
43 name="model.layers",
44 submodules={
45 "ln1": RMSNormalizationBridge(
46 name="operator_norm",
47 config=self.cfg,
48 ),
49 "ln2": RMSNormalizationBridge(
50 name="ffn_norm",
51 config=self.cfg,
52 ),
53 "attn": PositionEmbeddingsAttentionBridge(
54 name="self_attn",
55 config=self.cfg,
56 optional=True,
57 submodules={
58 "q": LinearBridge(name="q_proj"),
59 "k": LinearBridge(name="k_proj"),
60 "v": LinearBridge(name="v_proj"),
61 "o": LinearBridge(name="out_proj"),
62 "q_norm": RMSNormalizationBridge(name="q_layernorm", config=self.cfg),
63 "k_norm": RMSNormalizationBridge(name="k_layernorm", config=self.cfg),
64 },
65 requires_attention_mask=True,
66 requires_position_embeddings=True,
67 ),
68 "conv": Lfm2ShortConvBridge(
69 name="conv",
70 config=self.cfg,
71 optional=True,
72 submodules={
73 "in": LinearBridge(name="in_proj"),
74 "conv": DepthwiseConv1DBridge(name="conv"),
75 "out": LinearBridge(name="out_proj"),
76 },
77 ),
78 "mlp": self._gated_mlp(name="feed_forward", gate="w1", up="w3", down="w2"),
79 },
80 ),
81 "ln_final": RMSNormalizationBridge(name="model.embedding_norm", config=self.cfg),
82 "unembed": UnembeddingBridge(name="lm_head", config=self.cfg),
83 }
85 def setup_component_testing(self, hf_model: Any, bridge_model: Any = None) -> None:
86 """Set up model-specific references for component testing."""
87 rotary_emb = hf_model.model.rotary_emb
89 # Set attention implementation on HF model to eager (vs sdpa default)
90 force_eager_attention(hf_model, per_layer=True)
92 # Set rotary_emb on actual bridge instances
93 if bridge_model is not None and hasattr(bridge_model, "blocks"):
94 for block in bridge_model.blocks:
95 if hasattr(block, "attn"):
96 block.attn.set_rotary_emb(rotary_emb)
98 # Set on template for get_generalized_component() calls
99 # Find the first attention layer (LFM2 layer 0 is conv, not attn)
100 layer_types = getattr(self.cfg, "layer_types", None)
101 if layer_types is not None and "full_attention" in layer_types: 101 ↛ exitline 101 didn't return from function 'setup_component_testing' because the condition on line 101 was always true
102 first_attn_idx = layer_types.index("full_attention")
103 attn_bridge = self.get_generalized_component(f"blocks.{first_attn_idx}.attn")
104 attn_bridge.set_rotary_emb(rotary_emb)