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

1"""Lfm2 architecture adapter.""" 

2 

3from typing import Any 

4 

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 

18 

19 

20class Lfm2ArchitectureAdapter(ArchitectureAdapter): 

21 """Architecture adapter for Lfm2 models.""" 

22 

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

24 """Initialize the Lfm2 architecture adapter.""" 

25 super().__init__(cfg) 

26 

27 self._set_rms_rotary_defaults() 

28 self.cfg.act_fn = "silu" 

29 

30 self.cfg.attn_implementation = "eager" 

31 

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 

34 

35 self.weight_processing_conversions = { 

36 **self._qkvo_weight_conversions(), 

37 } 

38 

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 } 

84 

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 

88 

89 # Set attention implementation on HF model to eager (vs sdpa default) 

90 force_eager_attention(hf_model, per_layer=True) 

91 

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) 

97 

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)