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

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 GatedMLPBridge, 

11 Lfm2ShortConvBridge, 

12 LinearBridge, 

13 PositionEmbeddingsAttentionBridge, 

14 RMSNormalizationBridge, 

15 RotaryEmbeddingBridge, 

16 UnembeddingBridge, 

17) 

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.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" 

34 

35 self.cfg.attn_implementation = "eager" 

36 

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 

39 

40 self.weight_processing_conversions = { 

41 **self._qkvo_weight_conversions(), 

42 } 

43 

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 } 

97 

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 

101 

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" 

105 

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" 

110 

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) 

116 

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)