Coverage for transformer_lens/model_bridge/supported_architectures/minimax_m2.py: 100%
10 statements
« prev ^ index » next coverage.py v7.10.1, created at 2026-08-11 18:50 +0000
« prev ^ index » next coverage.py v7.10.1, created at 2026-08-11 18:50 +0000
1"""MiniMax-M2 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 MiniMaxM2ArchitectureAdapter(ArchitectureAdapter):
20 """Architecture adapter for MiniMaxM2ForCausalLM models -- Qwen3-MoE-like, but with
21 full-width (not per-head) Q/K norm and a sigmoid + e_score_correction_bias router."""
23 def __init__(self, cfg: Any) -> None:
24 """Initialize the MiniMax-M2 architecture adapter."""
25 super().__init__(cfg)
27 self._set_rms_rotary_defaults()
28 # Verified against MiniMaxAI/MiniMax-M2: tokenizer does not prepend BOS.
29 self.cfg.default_prepend_bos = False
31 # QKVO rearrangements; MoE expert and router weights pass through unchanged
32 self.weight_processing_conversions = {
33 **self._qkvo_weight_conversions(),
34 }
36 self.component_mapping = {
37 "embed": EmbeddingBridge(name="model.embed_tokens"),
38 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb", config=self.cfg),
39 "blocks": BlockBridge(
40 name="model.layers",
41 submodules={
42 "ln1": RMSNormalizationBridge(name="input_layernorm", config=self.cfg),
43 "ln2": RMSNormalizationBridge(name="post_attention_layernorm", config=self.cfg),
44 "attn": PositionEmbeddingsAttentionBridge(
45 name="self_attn",
46 config=self.cfg,
47 submodules={
48 "q": LinearBridge(name="q_proj"),
49 "k": LinearBridge(name="k_proj"),
50 "v": LinearBridge(name="v_proj"),
51 "o": LinearBridge(name="o_proj"),
52 # Full-width (all-heads) RMSNorm, applied pre-reshape.
53 "q_norm": RMSNormalizationBridge(name="q_norm", config=self.cfg),
54 "k_norm": RMSNormalizationBridge(name="k_norm", config=self.cfg),
55 },
56 requires_attention_mask=True,
57 requires_position_embeddings=True,
58 ),
59 "mlp": MoEBridge(
60 name="mlp",
61 config=self.cfg,
62 submodules={
63 "gate": MoERouterBridge(name="gate"),
64 },
65 ),
66 },
67 ),
68 "ln_final": RMSNormalizationBridge(name="model.norm", config=self.cfg),
69 "unembed": UnembeddingBridge(name="lm_head", config=self.cfg),
70 }