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

10 statements  

« prev     ^ index     » next       coverage.py v7.10.1, created at 2026-09-21 19:27 +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 RMSNormalizationBridge, 

10 RotaryEmbeddingBridge, 

11 UnembeddingBridge, 

12) 

13 

14 

15class LlamaArchitectureAdapter(ArchitectureAdapter): 

16 """Architecture adapter for Llama models. 

17 

18 Optional Parameters (may not exist in state_dict): 

19 ------------------------------------------------- 

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

21 

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

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

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

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

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

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

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

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

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

31 - ln_final.b - RMSNorm has no bias 

32 

33 Weight processing must handle these missing biases gracefully using 

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

35 """ 

36 

37 _testing_eager: Optional[str] = None 

38 

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

40 """Initialize the Llama architecture adapter.""" 

41 super().__init__(cfg) 

42 

43 self._set_rms_rotary_defaults() 

44 

45 self.weight_processing_conversions = { 

46 **self._qkvo_weight_conversions(), 

47 } 

48 

49 self.component_mapping = { 

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

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

52 "blocks": BlockBridge( 

53 name="model.layers", 

54 submodules={ 

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

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

57 "attn": self._build_attention_bridge(), 

58 "mlp": self._gated_mlp(), 

59 }, 

60 ), 

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

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

63 }