Coverage for transformer_lens/model_bridge/supported_architectures/qwen.py: 96%

48 statements  

« prev     ^ index     » next       coverage.py v7.10.1, created at 2026-09-01 16:23 +0000

1"""Qwen architecture adapter.""" 

2 

3from typing import Any 

4 

5import torch 

6 

7from transformer_lens.conversion_utils.conversion_steps import RearrangeTensorConversion 

8from transformer_lens.conversion_utils.param_processing_conversion import ( 

9 ParamProcessingConversion, 

10) 

11from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter 

12from transformer_lens.model_bridge.generalized_components import ( 

13 BlockBridge, 

14 EmbeddingBridge, 

15 JointQKVAttentionBridge, 

16 LinearBridge, 

17 NormalizationBridge, 

18 UnembeddingBridge, 

19) 

20 

21 

22class QwenArchitectureAdapter(ArchitectureAdapter): 

23 """Architecture adapter for Qwen models.""" 

24 

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

26 """Initialize the Qwen architecture adapter.""" 

27 super().__init__(cfg) 

28 

29 # Set config variables for weight processing 

30 self.cfg.normalization_type = "RMS" 

31 self.cfg.positional_embedding_type = "rotary" 

32 self.cfg.final_rms = True 

33 self.cfg.gated_mlp = True 

34 self.cfg.attn_only = False 

35 

36 self.weight_processing_conversions = { 

37 "blocks.{i}.attn.q": ParamProcessingConversion( 

38 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=self.cfg.n_heads), 

39 source_key="transformer.h.{i}.attn.c_attn.weight", 

40 ), 

41 "blocks.{i}.attn.k": ParamProcessingConversion( 

42 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=self.cfg.n_heads), 

43 source_key="transformer.h.{i}.attn.c_attn.weight", 

44 ), 

45 "blocks.{i}.attn.v": ParamProcessingConversion( 

46 tensor_conversion=RearrangeTensorConversion("(n h) m -> n m h", n=self.cfg.n_heads), 

47 source_key="transformer.h.{i}.attn.c_attn.weight", 

48 ), 

49 "blocks.{i}.attn.o": ParamProcessingConversion( 

50 tensor_conversion=RearrangeTensorConversion("m (n h) -> n h m", n=self.cfg.n_heads), 

51 source_key="transformer.h.{i}.attn.c_proj.weight", 

52 ), 

53 } 

54 

55 self.component_mapping = { 

56 "embed": EmbeddingBridge(name="transformer.wte"), 

57 "blocks": BlockBridge( 

58 name="transformer.h", 

59 submodules={ 

60 "ln1": NormalizationBridge(name="ln_1", config=self.cfg), 

61 "attn": JointQKVAttentionBridge( 

62 name="attn", 

63 config=self.cfg, 

64 split_qkv_matrix=self._split_qkv_matrix, 

65 submodules={ 

66 "qkv": LinearBridge(name="c_attn"), 

67 "o": LinearBridge(name="c_proj"), 

68 }, 

69 ), 

70 "ln2": NormalizationBridge(name="ln_2", config=self.cfg), 

71 "mlp": self._gated_mlp(gate="w1", up="w2", down="c_proj"), 

72 }, 

73 ), 

74 "ln_final": NormalizationBridge(name="transformer.ln_f", config=self.cfg), 

75 "unembed": UnembeddingBridge(name="lm_head"), 

76 } 

77 

78 def _split_qkv_matrix( 

79 self, original_attention_component: Any 

80 ) -> tuple[torch.nn.Linear, torch.nn.Linear, torch.nn.Linear]: 

81 """Split Qwen's fused c_attn linear layer into q, k, v projections.""" 

82 

83 assert original_attention_component is not None 

84 assert hasattr(original_attention_component, "c_attn") 

85 

86 c_attn = original_attention_component.c_attn 

87 assert isinstance(c_attn, torch.nn.Linear) 

88 

89 d_model = self.cfg.d_model 

90 qkv_weights = c_attn.weight.detach().clone() 

91 

92 if qkv_weights.shape == (d_model, 3 * d_model): 

93 # Weight stored as [in_features, 3*out_features] (Conv1D style) 

94 W_Q, W_K, W_V = torch.tensor_split(qkv_weights, 3, dim=1) 

95 W_Q, W_K, W_V = W_Q.T.contiguous(), W_K.T.contiguous(), W_V.T.contiguous() 

96 elif qkv_weights.shape == (3 * d_model, d_model): 

97 # Standard Linear layout [3*out_features, in_features] 

98 W_Q, W_K, W_V = torch.tensor_split(qkv_weights, 3, dim=0) 

99 else: 

100 raise ValueError( 

101 f"Unexpected c_attn weight shape {qkv_weights.shape} for Qwen attention " 

102 f"(expected ({d_model}, {3*d_model}) or ({3*d_model}, {d_model}))" 

103 ) 

104 

105 if c_attn.bias is not None: 

106 qkv_bias = c_attn.bias.detach().clone() 

107 if qkv_bias.shape[0] != 3 * d_model: 107 ↛ 108line 107 didn't jump to line 108 because the condition on line 107 was never true

108 raise ValueError( 

109 f"Unexpected c_attn bias shape {qkv_bias.shape} for Qwen attention " 

110 f"(expected ({3*d_model},))" 

111 ) 

112 b_Q, b_K, b_V = torch.tensor_split(qkv_bias, 3, dim=0) 

113 else: 

114 device = qkv_weights.device 

115 dtype = qkv_weights.dtype 

116 b_Q = torch.zeros(d_model, device=device, dtype=dtype) 

117 b_K = torch.zeros_like(b_Q) 

118 b_V = torch.zeros_like(b_Q) 

119 

120 def build_linear(weight: torch.Tensor, bias: torch.Tensor) -> torch.nn.Linear: 

121 linear = torch.nn.Linear( 

122 d_model, d_model, bias=True, device=weight.device, dtype=weight.dtype 

123 ) 

124 linear.weight = torch.nn.Parameter(weight.contiguous()) 

125 linear.bias = torch.nn.Parameter(bias.contiguous()) 

126 return linear 

127 

128 q_proj = build_linear(W_Q, b_Q) 

129 k_proj = build_linear(W_K, b_K) 

130 v_proj = build_linear(W_V, b_V) 

131 

132 return q_proj, k_proj, v_proj