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

1"""Qwen2 architecture adapter.""" 

2 

3from typing import Any, Optional 

4 

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) 

13 

14 

15class Qwen2ArchitectureAdapter(ArchitectureAdapter): 

16 """Architecture adapter for Qwen2 models. 

17 

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

22 

23 _testing_eager: Optional[str] = None 

24 

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

26 """Initialize the Qwen2 architecture adapter.""" 

27 super().__init__(cfg) 

28 

29 self._set_rms_rotary_defaults() 

30 

31 self.cfg.default_prepend_bos = False 

32 

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 }