Coverage for transformer_lens/model_bridge/supported_architectures/qwen2.py: 100%
11 statements
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-21 19:27 +0000
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-21 19:27 +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 RMSNormalizationBridge,
10 RotaryEmbeddingBridge,
11 UnembeddingBridge,
12)
15class Qwen2ArchitectureAdapter(ArchitectureAdapter):
16 """Architecture adapter for Qwen2 models.
18 Qwen2 hardcodes q/k/v biases (o_proj, MLP, and norms are bias-free); the
19 include_biases conversions keep GQA K/V biases in the per-head
20 (n_kv_heads, d_head) layout weight processing expects.
21 """
23 _testing_eager: Optional[str] = None
25 def __init__(self, cfg: Any) -> None:
26 """Initialize the Qwen2 architecture adapter."""
27 super().__init__(cfg)
29 self._set_rms_rotary_defaults()
31 self.cfg.default_prepend_bos = False
33 self.weight_processing_conversions = {
34 **self._qkvo_weight_conversions(include_biases=True),
35 }
36 self.component_mapping = {
37 "embed": EmbeddingBridge(name="model.embed_tokens"),
38 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb"),
39 "blocks": BlockBridge(
40 name="model.layers",
41 config=self.cfg,
42 submodules={
43 "ln1": RMSNormalizationBridge(name="input_layernorm", config=self.cfg),
44 "ln2": RMSNormalizationBridge(name="post_attention_layernorm", config=self.cfg),
45 "attn": self._build_attention_bridge(),
46 # GatedMLPBridge: hook_pre = gate pre-activation (HT GatedMLP
47 # semantics) + compat reconstruction; plain MLPBridge pointed
48 # hook_pre at the up-projection.
49 "mlp": self._gated_mlp(),
50 },
51 ),
52 "ln_final": RMSNormalizationBridge(name="model.norm", config=self.cfg),
53 "unembed": UnembeddingBridge(name="lm_head", config=self.cfg),
54 }