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

1"""Qwen3 architecture adapter. 

2 

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

7 

8from typing import Any 

9 

10import torch 

11 

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) 

27 

28 

29class Qwen3ArchitectureAdapter(ArchitectureAdapter): 

30 """Architecture adapter for Qwen3 dense models. 

31 

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

35 

36 _testing_hybrid = True 

37 

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) 

47 

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" 

53 

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 ) 

69 

70 def _build_mlp_bridge(self): 

71 """Dense gated MLP (gate_proj + up_proj -> down_proj). Override for MoE.""" 

72 return self._gated_mlp() 

73 

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 ) 

81 

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 } 

102 

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

108 

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