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