Coverage for transformer_lens/model_bridge/supported_architectures/mixtral.py: 100%
9 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"""Mixtral 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 MoEBridge,
11 MoERouterBridge,
12 PositionEmbeddingsAttentionBridge,
13 RMSNormalizationBridge,
14 RotaryEmbeddingBridge,
15 UnembeddingBridge,
16)
19class MixtralArchitectureAdapter(ArchitectureAdapter):
20 """Architecture adapter for Mixtral models.
22 Mixtral uses a pre-norm architecture with RMSNorm, rotary position embeddings
23 (RoPE), and a Sparse Mixture of Experts MLP. Key features:
25 - Pre-norm: RMSNorm applied BEFORE attention and BEFORE MLP.
26 - Rotary embeddings: stored at model.rotary_emb and passed per-forward-call.
27 - Sparse MoE: batched expert parameters (gate_up_proj, down_proj as 3D tensors).
28 - MixtralAttention.forward() requires position_embeddings and attention_mask args.
29 - Optional GQA (n_key_value_heads may differ from n_heads).
30 """
32 def __init__(self, cfg: Any) -> None:
33 """Initialize the Mixtral architecture adapter."""
34 super().__init__(cfg)
36 self._set_rms_rotary_defaults(final_rms=False)
38 self.weight_processing_conversions = {
39 **self._qkvo_weight_conversions(),
40 }
42 # Set up component mapping
43 self.component_mapping = {
44 "embed": EmbeddingBridge(name="model.embed_tokens"),
45 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb", config=self.cfg),
46 "blocks": BlockBridge(
47 name="model.layers",
48 submodules={
49 "ln1": RMSNormalizationBridge(name="input_layernorm", config=self.cfg),
50 "ln2": RMSNormalizationBridge(name="post_attention_layernorm", config=self.cfg),
51 # MixtralAttention.forward() requires position_embeddings and
52 # attention_mask as positional arguments (not optional kwargs).
53 "attn": PositionEmbeddingsAttentionBridge(
54 name="self_attn",
55 config=self.cfg,
56 submodules={
57 "q": LinearBridge(name="q_proj"),
58 "k": LinearBridge(name="k_proj"),
59 "v": LinearBridge(name="v_proj"),
60 "o": LinearBridge(name="o_proj"),
61 },
62 requires_attention_mask=True,
63 requires_position_embeddings=True,
64 ),
65 # Mixtral uses batched expert parameters (gate_up_proj, down_proj
66 # as 3D tensors) rather than a ModuleList of individual experts.
67 # MoEBridge wraps the entire MLP module and delegates to HF's
68 # native forward pass. 5.13 renamed the decoder-layer attr
69 # block_sparse_moe -> mlp.
70 "mlp": MoEBridge(
71 name="mlp",
72 config=self.cfg,
73 submodules={
74 "gate": MoERouterBridge(name="gate"),
75 },
76 ),
77 },
78 ),
79 "ln_final": RMSNormalizationBridge(name="model.norm", config=self.cfg),
80 "unembed": UnembeddingBridge(name="lm_head"),
81 }