Coverage for transformer_lens/model_bridge/supported_architectures/qwen3.py: 100%
39 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"""Qwen3 architecture adapter.
3Base adapter for the Qwen3 model family. Provides shared config setup,
4attention bridge construction, and setup_component_testing used by
5Qwen3, Qwen3.5, and Qwen3Next variants.
6"""
8from typing import Any
10import torch
12from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
13from transformer_lens.model_bridge.generalized_components import (
14 BlockBridge,
15 EmbeddingBridge,
16 LinearBridge,
17 RMSNormalizationBridge,
18 RotaryEmbeddingBridge,
19 UnembeddingBridge,
20)
21from transformer_lens.model_bridge.generalized_components.gated_delta_net import (
22 GatedDeltaNetBridge,
23)
24from transformer_lens.model_bridge.generalized_components.position_embeddings_attention import (
25 PositionEmbeddingsAttentionBridge,
26)
29class Qwen3ArchitectureAdapter(ArchitectureAdapter):
30 """Architecture adapter for Qwen3 dense models.
32 RMSNorm, RoPE, GQA, Q/K head norms, gated MLP. No biases.
33 Serves as base class for Qwen3.5 and Qwen3Next hybrid variants.
34 """
36 _testing_hybrid = True
38 def __init__(self, cfg: Any, *, hybrid: bool = False, lm_prefix: str = "model") -> None:
39 super().__init__(cfg)
40 self._setup_qwen3_config(cfg)
41 if hybrid:
42 self.supports_fold_ln = False
43 self.weight_processing_conversions: dict = {}
44 else:
45 self.weight_processing_conversions = {**self._qkvo_weight_conversions()}
46 self.component_mapping = self._build_component_mapping(hybrid=hybrid, lm_prefix=lm_prefix)
48 def _setup_qwen3_config(self, cfg: Any) -> None:
49 """Config shared across all Qwen3 variants (dense, hybrid, MoE)."""
50 self._set_rms_rotary_defaults()
51 self.cfg.default_prepend_bos = False
52 self.cfg.attn_implementation = "eager"
54 def _build_attention_bridge(self, optional: bool = False) -> PositionEmbeddingsAttentionBridge:
55 """Standard Qwen3 attention bridge with Q/K norms."""
56 return PositionEmbeddingsAttentionBridge(
57 name="self_attn",
58 config=self.cfg,
59 optional=optional,
60 submodules={
61 "q": LinearBridge(name="q_proj"),
62 "k": LinearBridge(name="k_proj"),
63 "v": LinearBridge(name="v_proj"),
64 "o": LinearBridge(name="o_proj"),
65 "q_norm": RMSNormalizationBridge(name="q_norm", config=self.cfg),
66 "k_norm": RMSNormalizationBridge(name="k_norm", config=self.cfg),
67 },
68 )
70 def _build_mlp_bridge(self):
71 """Dense gated MLP (gate_proj + up_proj -> down_proj). Override for MoE."""
72 return self._gated_mlp()
74 def _build_linear_attn_bridge(self, optional: bool = False) -> GatedDeltaNetBridge:
75 """GatedDeltaNet linear-attention bridge for hybrid variants."""
76 return GatedDeltaNetBridge(
77 name="linear_attn",
78 config=self.cfg,
79 optional=optional,
80 )
82 def _build_component_mapping(self, *, hybrid: bool = False, lm_prefix: str = "model") -> dict:
83 """Parametric component mapping. hybrid=True adds optional linear_attn; lm_prefix
84 nests the text model (``model``, or ``model.language_model`` for multimodal). lm_head
85 stays top-level.
86 """
87 block_submodules: dict = {
88 "ln1": RMSNormalizationBridge(name="input_layernorm", config=self.cfg),
89 "ln2": RMSNormalizationBridge(name="post_attention_layernorm", config=self.cfg),
90 "attn": self._build_attention_bridge(optional=hybrid),
91 "mlp": self._build_mlp_bridge(),
92 }
93 if hybrid:
94 block_submodules["linear_attn"] = self._build_linear_attn_bridge(optional=True)
95 return {
96 "embed": EmbeddingBridge(name=f"{lm_prefix}.embed_tokens"),
97 "rotary_emb": RotaryEmbeddingBridge(name=f"{lm_prefix}.rotary_emb", config=self.cfg),
98 "blocks": BlockBridge(name=f"{lm_prefix}.layers", submodules=block_submodules),
99 "ln_final": RMSNormalizationBridge(name=f"{lm_prefix}.norm", config=self.cfg),
100 "unembed": UnembeddingBridge(name="lm_head"),
101 }
103 @staticmethod
104 def _preprocess_gated_q_proj(
105 state_dict: dict[str, torch.Tensor], n_heads: int, d_head: int
106 ) -> dict[str, torch.Tensor]:
107 """Slice query half from gated q_proj.weight (interleaved per-head layout).
109 q_proj.weight has shape (n_heads * d_head * 2, hidden_size) with
110 interleaved [query, gate] rows per head. Extracts query-only half.
111 """
112 keys_to_update = [k for k in state_dict if k.endswith(".self_attn.q_proj.weight")]
113 for key in keys_to_update:
114 w = state_dict[key]
115 w = w.view(n_heads, d_head * 2, -1)
116 state_dict[key] = w[:, :d_head, :].reshape(n_heads * d_head, -1)
117 return state_dict