Coverage for transformer_lens/model_bridge/supported_architectures/qwen3_moe.py: 100%
11 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"""Qwen3MoE (Mixture of Experts) 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 Qwen3MoeArchitectureAdapter(ArchitectureAdapter):
20 """Architecture adapter for Qwen3MoE (Mixture of Experts) models.
22 Qwen3MoE is a sparse MoE decoder-only Transformer, structurally close to OLMoE.
23 Key features:
25 - Pre-norm: RMSNorm applied BEFORE attention and BEFORE MLP.
26 - Q/K normalization: RMSNorm applied to queries and keys after projection.
27 - Sparse MoE: 128 experts with top-8 routing (public 30B-A3B checkpoints).
28 - Batched expert parameters: gate_up_proj and down_proj as single 3D tensors,
29 not a ModuleList.
30 - final_rms=True (Qwen3-style; OLMoE uses False).
31 - No biases on any projections.
32 - GQA: n_key_value_heads < n_heads in all public checkpoints.
34 Only the all-MoE configuration is supported (decoder_sparse_step=1,
35 mlp_only_layers=[]). Models with dense fallback layers cannot be wrapped
36 because MoEBridge does not handle the dense Qwen3MoeMLP path.
38 Optional Parameters (may not exist in state_dict):
39 -------------------------------------------------
40 - blocks.{i}.attn.b_Q - No bias on query projection
41 - blocks.{i}.attn.b_K - No bias on key projection
42 - blocks.{i}.attn.b_V - No bias on value projection
43 - blocks.{i}.attn.b_O - No bias on output projection
44 - blocks.{i}.ln1.b - RMSNorm has no bias
45 - blocks.{i}.ln2.b - RMSNorm has no bias
46 - ln_final.b - RMSNorm has no bias
47 """
49 def __init__(self, cfg: Any) -> None:
50 """Initialize the Qwen3MoE architecture adapter."""
51 super().__init__(cfg)
53 self._set_rms_rotary_defaults()
54 # Force eager attention for output_attentions hook support
55 self.cfg.attn_implementation = "eager"
56 self.cfg.default_prepend_bos = False # Qwen3 family convention
58 # QKVO rearrangements; MoE expert and gate weights pass through unchanged
59 self.weight_processing_conversions = {
60 **self._qkvo_weight_conversions(),
61 }
63 # Component mapping — PRE-NORM architecture:
64 # ln1 = input_layernorm (applied BEFORE attention)
65 # ln2 = post_attention_layernorm (applied BEFORE MLP)
66 self.component_mapping = {
67 "embed": EmbeddingBridge(name="model.embed_tokens"),
68 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb", config=self.cfg),
69 "blocks": BlockBridge(
70 name="model.layers",
71 submodules={
72 "ln1": RMSNormalizationBridge(name="input_layernorm", config=self.cfg),
73 "ln2": RMSNormalizationBridge(name="post_attention_layernorm", config=self.cfg),
74 "attn": PositionEmbeddingsAttentionBridge(
75 name="self_attn",
76 config=self.cfg,
77 submodules={
78 "q": LinearBridge(name="q_proj"),
79 "k": LinearBridge(name="k_proj"),
80 "v": LinearBridge(name="v_proj"),
81 "o": LinearBridge(name="o_proj"),
82 "q_norm": RMSNormalizationBridge(name="q_norm", config=self.cfg),
83 "k_norm": RMSNormalizationBridge(name="k_norm", config=self.cfg),
84 },
85 requires_attention_mask=True,
86 requires_position_embeddings=True,
87 ),
88 # Qwen3MoeSparseMoeBlock stores experts as batched 3D tensors
89 # rather than a ModuleList. MoEBridge wraps the entire block and
90 # delegates to HF's native forward — same pattern as OLMoE.
91 "mlp": MoEBridge(
92 name="mlp",
93 config=self.cfg,
94 sparse_required=("gate",),
95 submodules={
96 # Dense fallback layers (mlp_only_layers /
97 # decoder_sparse_step) have no router; their
98 # projections bind gated-MLP neuron hooks (#1645).
99 "gate": MoERouterBridge(name="gate", optional=True),
100 "dense_gate": LinearBridge(name="gate_proj", optional=True),
101 "dense_in": LinearBridge(name="up_proj", optional=True),
102 "dense_out": LinearBridge(name="down_proj", optional=True),
103 },
104 ),
105 },
106 ),
107 "ln_final": RMSNormalizationBridge(name="model.norm", config=self.cfg),
108 "unembed": UnembeddingBridge(name="lm_head", config=self.cfg),
109 }