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
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
1"""Mistral 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 EmbeddingBridge,
9 LinearBridge,
10 PositionEmbeddingsAttentionBridge,
11 RMSNormalizationBridge,
12 RotaryEmbeddingBridge,
13 UnembeddingBridge,
14)
17class MistralArchitectureAdapter(ArchitectureAdapter):
18 """Architecture adapter for Mistral models."""
20 def __init__(self, cfg: Any) -> None:
21 """Initialize the Mistral architecture adapter."""
22 super().__init__(cfg)
24 self._set_rms_rotary_defaults(final_rms=False)
26 self.weight_processing_conversions = {
27 **self._qkvo_weight_conversions(),
28 }
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 }