Coverage for transformer_lens/model_bridge/supported_architectures/llama.py: 100%

13 statements  

« prev     ^ index     » next       coverage.py v7.10.1, created at 2026-09-01 16:23 +0000

1"""Llama 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 LlamaArchitectureAdapter(ArchitectureAdapter): 

18 """Architecture adapter for Llama models. 

19 

20 Optional Parameters (may not exist in state_dict): 

21 ------------------------------------------------- 

22 LLaMA models do NOT have biases on attention and MLP projections: 

23 

24 - blocks.{i}.attn.b_Q - No bias on query projection 

25 - blocks.{i}.attn.b_K - No bias on key projection 

26 - blocks.{i}.attn.b_V - No bias on value projection 

27 - blocks.{i}.attn.b_O - No bias on output projection 

28 - blocks.{i}.mlp.b_in - No bias on MLP input (up_proj) 

29 - blocks.{i}.mlp.b_gate - No bias on MLP gate projection 

30 - blocks.{i}.mlp.b_out - No bias on MLP output (down_proj) 

31 - blocks.{i}.ln1.b - RMSNorm has no bias 

32 - blocks.{i}.ln2.b - RMSNorm has no bias 

33 - ln_final.b - RMSNorm has no bias 

34 

35 Weight processing must handle these missing biases gracefully using 

36 ProcessWeights._safe_get_tensor() or by checking for None values. 

37 """ 

38 

39 _testing_eager: Optional[str] = None 

40 

41 _attention_bridge_cls = PositionEmbeddingsAttentionBridge 

42 

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

44 """Initialize the Llama architecture adapter.""" 

45 super().__init__(cfg) 

46 

47 self._set_rms_rotary_defaults() 

48 

49 self.weight_processing_conversions = { 

50 **self._qkvo_weight_conversions(), 

51 } 

52 

53 self.component_mapping = { 

54 "embed": EmbeddingBridge(name="model.embed_tokens"), 

55 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb"), 

56 "blocks": BlockBridge( 

57 name="model.layers", 

58 submodules={ 

59 "ln1": RMSNormalizationBridge(name="input_layernorm", config=self.cfg), 

60 "ln2": RMSNormalizationBridge(name="post_attention_layernorm", config=self.cfg), 

61 "attn": self._build_attention_bridge(), 

62 "mlp": self._gated_mlp(), 

63 }, 

64 ), 

65 "ln_final": RMSNormalizationBridge(name="model.norm", config=self.cfg), 

66 "unembed": UnembeddingBridge(name="lm_head", config=self.cfg), 

67 } 

68 

69 def _build_attention_bridge(self): 

70 """Attention bridge seam; subclasses swap the class or the construction.""" 

71 return self._attention_bridge_cls( 

72 name="self_attn", 

73 config=self.cfg, 

74 submodules={ 

75 "q": LinearBridge(name="q_proj"), 

76 "k": LinearBridge(name="k_proj"), 

77 "v": LinearBridge(name="v_proj"), 

78 "o": LinearBridge(name="o_proj"), 

79 }, 

80 requires_attention_mask=True, 

81 requires_position_embeddings=True, 

82 )