Coverage for transformer_lens/model_bridge/supported_architectures/qwen2.py: 100%
14 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"""Qwen2 architecture adapter."""
3from typing import Any, Optional
5from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
6from transformer_lens.model_bridge.generalized_components import (
7 BlockBridge,
8 EmbeddingBridge,
9 LinearBridge,
10 PositionEmbeddingsAttentionBridge,
11 RMSNormalizationBridge,
12 RotaryEmbeddingBridge,
13 UnembeddingBridge,
14)
17class Qwen2ArchitectureAdapter(ArchitectureAdapter):
18 """Architecture adapter for Qwen2 models.
20 Qwen2 hardcodes q/k/v biases (o_proj, MLP, and norms are bias-free); the
21 include_biases conversions keep GQA K/V biases in the per-head
22 (n_kv_heads, d_head) layout weight processing expects.
23 """
25 _testing_eager: Optional[str] = None
27 _attention_bridge_cls = PositionEmbeddingsAttentionBridge
29 def __init__(self, cfg: Any) -> None:
30 """Initialize the Qwen2 architecture adapter."""
31 super().__init__(cfg)
33 self._set_rms_rotary_defaults()
35 self.cfg.default_prepend_bos = False
37 self.weight_processing_conversions = {
38 **self._qkvo_weight_conversions(include_biases=True),
39 }
40 self.component_mapping = {
41 "embed": EmbeddingBridge(name="model.embed_tokens"),
42 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb"),
43 "blocks": BlockBridge(
44 name="model.layers",
45 config=self.cfg,
46 submodules={
47 "ln1": RMSNormalizationBridge(name="input_layernorm", config=self.cfg),
48 "ln2": RMSNormalizationBridge(name="post_attention_layernorm", config=self.cfg),
49 "attn": self._build_attention_bridge(),
50 # GatedMLPBridge: hook_pre = gate pre-activation (HT GatedMLP
51 # semantics) + compat reconstruction; plain MLPBridge pointed
52 # hook_pre at the up-projection.
53 "mlp": self._gated_mlp(),
54 },
55 ),
56 "ln_final": RMSNormalizationBridge(name="model.norm", config=self.cfg),
57 "unembed": UnembeddingBridge(name="lm_head", config=self.cfg),
58 }
60 def _build_attention_bridge(self):
61 """Attention bridge seam; subclasses swap the class or the construction."""
62 return self._attention_bridge_cls(
63 name="self_attn",
64 config=self.cfg,
65 submodules={
66 "q": LinearBridge(name="q_proj"),
67 "k": LinearBridge(name="k_proj"),
68 "v": LinearBridge(name="v_proj"),
69 "o": LinearBridge(name="o_proj"),
70 },
71 requires_attention_mask=True,
72 requires_position_embeddings=True,
73 )