Coverage for transformer_lens/model_bridge/supported_architectures/mistral.py: 100%

9 statements  

« prev     ^ index     » next       coverage.py v7.10.1, created at 2026-09-01 16:23 +0000

1"""Mistral 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 EmbeddingBridge, 

9 LinearBridge, 

10 PositionEmbeddingsAttentionBridge, 

11 RMSNormalizationBridge, 

12 RotaryEmbeddingBridge, 

13 UnembeddingBridge, 

14) 

15 

16 

17class MistralArchitectureAdapter(ArchitectureAdapter): 

18 """Architecture adapter for Mistral models.""" 

19 

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

21 """Initialize the Mistral architecture adapter.""" 

22 super().__init__(cfg) 

23 

24 self._set_rms_rotary_defaults(final_rms=False) 

25 

26 self.weight_processing_conversions = { 

27 **self._qkvo_weight_conversions(), 

28 } 

29 

30 self.component_mapping = { 

31 "embed": EmbeddingBridge(name="model.embed_tokens"), 

32 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb", config=self.cfg), 

33 "blocks": BlockBridge( 

34 name="model.layers", 

35 config=self.cfg, 

36 submodules={ 

37 "ln1": RMSNormalizationBridge(name="input_layernorm", config=self.cfg), 

38 "ln2": RMSNormalizationBridge(name="post_attention_layernorm", config=self.cfg), 

39 "attn": PositionEmbeddingsAttentionBridge( 

40 name="self_attn", 

41 config=self.cfg, 

42 requires_position_embeddings=True, 

43 requires_attention_mask=True, 

44 submodules={ 

45 "q": LinearBridge(name="q_proj"), 

46 "k": LinearBridge(name="k_proj"), 

47 "v": LinearBridge(name="v_proj"), 

48 "o": LinearBridge(name="o_proj"), 

49 }, 

50 ), 

51 "mlp": self._gated_mlp(), 

52 }, 

53 ), 

54 "ln_final": RMSNormalizationBridge(name="model.norm", config=self.cfg), 

55 "unembed": UnembeddingBridge(name="lm_head"), 

56 }