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-09-21 19:27 +0000

1"""MiniMax-M2 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 MoEBridge, 

11 MoERouterBridge, 

12 PositionEmbeddingsAttentionBridge, 

13 RMSNormalizationBridge, 

14 RotaryEmbeddingBridge, 

15 UnembeddingBridge, 

16) 

17 

18 

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

22 

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

24 """Initialize the MiniMax-M2 architecture adapter.""" 

25 super().__init__(cfg) 

26 

27 self._set_rms_rotary_defaults() 

28 # Verified against MiniMaxAI/MiniMax-M2: tokenizer does not prepend BOS. 

29 self.cfg.default_prepend_bos = False 

30 

31 # QKVO rearrangements; MoE expert and router weights pass through unchanged 

32 self.weight_processing_conversions = { 

33 **self._qkvo_weight_conversions(), 

34 } 

35 

36 # Deliberate mirror of olmoe.py / qwen3_moe.py: same wiring by structural 

37 # coincidence, not lineage (norm/router semantics differ per vendor), 

38 # so each file keeps its mapping inline and readable. 

39 self.component_mapping = { 

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

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

42 "blocks": BlockBridge( 

43 name="model.layers", 

44 submodules={ 

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

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

47 "attn": PositionEmbeddingsAttentionBridge( 

48 name="self_attn", 

49 config=self.cfg, 

50 submodules={ 

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

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

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

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

55 # Full-width (all-heads) RMSNorm, applied pre-reshape. 

56 "q_norm": RMSNormalizationBridge(name="q_norm", config=self.cfg), 

57 "k_norm": RMSNormalizationBridge(name="k_norm", config=self.cfg), 

58 }, 

59 requires_attention_mask=True, 

60 requires_position_embeddings=True, 

61 ), 

62 "mlp": MoEBridge( 

63 name="mlp", 

64 config=self.cfg, 

65 submodules={ 

66 "gate": MoERouterBridge(name="gate"), 

67 }, 

68 ), 

69 }, 

70 ), 

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

72 "unembed": UnembeddingBridge(name="lm_head", config=self.cfg), 

73 }