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
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
1"""Qwen architecture adapter."""
3from typing import Any
5import torch
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)
22class QwenArchitectureAdapter(ArchitectureAdapter):
23 """Architecture adapter for Qwen models."""
25 def __init__(self, cfg: Any) -> None:
26 """Initialize the Qwen architecture adapter."""
27 super().__init__(cfg)
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
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 }
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 }
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."""
83 assert original_attention_component is not None
84 assert hasattr(original_attention_component, "c_attn")
86 c_attn = original_attention_component.c_attn
87 assert isinstance(c_attn, torch.nn.Linear)
89 d_model = self.cfg.d_model
90 qkv_weights = c_attn.weight.detach().clone()
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 )
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)
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
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)
132 return q_proj, k_proj, v_proj