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
« prev ^ index » next coverage.py v7.10.1, created at 2026-09-01 16:23 +0000
1"""Llama 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 LinearBridge,
10 PositionEmbeddingsAttentionBridge,
11 RMSNormalizationBridge,
12 RotaryEmbeddingBridge,
13 UnembeddingBridge,
14)
17class LlamaArchitectureAdapter(ArchitectureAdapter):
18 """Architecture adapter for Llama models.
20 Optional Parameters (may not exist in state_dict):
21 -------------------------------------------------
22 LLaMA models do NOT have biases on attention and MLP projections:
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
35 Weight processing must handle these missing biases gracefully using
36 ProcessWeights._safe_get_tensor() or by checking for None values.
37 """
39 _testing_eager: Optional[str] = None
41 _attention_bridge_cls = PositionEmbeddingsAttentionBridge
43 def __init__(self, cfg: Any) -> None:
44 """Initialize the Llama architecture adapter."""
45 super().__init__(cfg)
47 self._set_rms_rotary_defaults()
49 self.weight_processing_conversions = {
50 **self._qkvo_weight_conversions(),
51 }
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 }
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 )