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

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 LinearBridge, 

10 PositionEmbeddingsAttentionBridge, 

11 RMSNormalizationBridge, 

12 RotaryEmbeddingBridge, 

13 UnembeddingBridge, 

14) 

15 

16 

17class Qwen2ArchitectureAdapter(ArchitectureAdapter): 

18 """Architecture adapter for Qwen2 models. 

19 

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

24 

25 _testing_eager: Optional[str] = None 

26 

27 _attention_bridge_cls = PositionEmbeddingsAttentionBridge 

28 

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

30 """Initialize the Qwen2 architecture adapter.""" 

31 super().__init__(cfg) 

32 

33 self._set_rms_rotary_defaults() 

34 

35 self.cfg.default_prepend_bos = False 

36 

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 } 

59 

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 )