Coverage for transformer_lens/model_bridge/supported_architectures/stablelm.py: 100%
26 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"""StableLM architecture adapter."""
3from typing import Any
5from transformer_lens.conversion_utils.conversion_steps import RearrangeTensorConversion
6from transformer_lens.conversion_utils.param_processing_conversion import (
7 ParamProcessingConversion,
8)
9from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
10from transformer_lens.model_bridge.generalized_components import (
11 BlockBridge,
12 EmbeddingBridge,
13 LinearBridge,
14 NormalizationBridge,
15 ParallelBlockBridge,
16 PositionEmbeddingsAttentionBridge,
17 RotaryEmbeddingBridge,
18 UnembeddingBridge,
19)
20from transformer_lens.model_bridge.generalized_components.base import (
21 GeneralizedComponent,
22)
25class StableLmArchitectureAdapter(ArchitectureAdapter):
26 """Architecture adapter for StableLM models.
28 StableLM uses a Llama-like architecture with separate Q/K/V projections and
29 gated MLP, but differs in using standard LayerNorm (not RMSNorm) and partial
30 rotary embeddings (25% of head dimensions by default).
32 Supports optional features:
33 - Grouped Query Attention (num_key_value_heads != num_attention_heads)
34 - QKV bias (use_qkv_bias=True on some models like stable-code-3b)
35 - Parallel residual connections (use_parallel_residual=True)
36 - Per-head QK LayerNorm (qk_layernorm=True)
38 Optional Parameters (may not exist in state_dict):
39 -------------------------------------------------
40 - blocks.{i}.attn.b_Q - Only present when use_qkv_bias=True
41 - blocks.{i}.attn.b_K - Only present when use_qkv_bias=True
42 - blocks.{i}.attn.b_V - Only present when use_qkv_bias=True
43 - blocks.{i}.attn.b_O - No bias on output projection
44 - blocks.{i}.mlp.b_in - No bias on MLP up_proj
45 - blocks.{i}.mlp.b_gate - No bias on MLP gate_proj
46 - blocks.{i}.mlp.b_out - No bias on MLP down_proj
47 """
49 def __init__(self, cfg: Any) -> None:
50 """Initialize the StableLM architecture adapter."""
51 super().__init__(cfg)
53 # Set config variables for weight processing
54 self.cfg.normalization_type = "LN"
55 self.cfg.positional_embedding_type = "rotary"
56 self.cfg.final_rms = False
57 self.cfg.gated_mlp = True
58 self.cfg.attn_only = False
59 self.cfg.uses_rms_norm = False
60 # The bridge reimplements attention; the HF reference must run the
61 # matching eager math.
62 self.cfg.attn_implementation = "eager"
64 n_kv_heads = getattr(self.cfg, "n_key_value_heads", None) or self.cfg.n_heads
66 self.weight_processing_conversions = {
67 **self._qkvo_weight_conversions(),
68 # Bias conversions for models with use_qkv_bias=True
69 "blocks.{i}.attn.q.bias": ParamProcessingConversion(
70 tensor_conversion=RearrangeTensorConversion("(n h) -> n h", n=self.cfg.n_heads),
71 ),
72 "blocks.{i}.attn.k.bias": ParamProcessingConversion(
73 tensor_conversion=RearrangeTensorConversion("(n h) -> n h", n=n_kv_heads),
74 ),
75 "blocks.{i}.attn.v.bias": ParamProcessingConversion(
76 tensor_conversion=RearrangeTensorConversion("(n h) -> n h", n=n_kv_heads),
77 ),
78 }
80 # When parallel_attn_mlp=True (HF: use_parallel_residual=True), both attn
81 # and MLP read from ln1 output:
82 # x = x + attn(ln1(x)) + mlp(ln1(x))
83 # When False, they are sequential with separate norms:
84 # x = x + attn(ln1(x)); x = x + mlp(ln2(x))
85 # HF sets post_attention_layernorm=None when use_parallel_residual=True,
86 # so we must not include ln2 in that case.
87 use_parallel_residual = getattr(cfg, "parallel_attn_mlp", False)
89 block_submodules: dict[str, Any] = {
90 "ln1": NormalizationBridge(
91 name="input_layernorm",
92 config=self.cfg,
93 use_native_layernorm_autograd=True,
94 ),
95 }
96 if not use_parallel_residual:
97 block_submodules["ln2"] = NormalizationBridge(
98 name="post_attention_layernorm",
99 config=self.cfg,
100 use_native_layernorm_autograd=True,
101 )
102 block_submodules["attn"] = PositionEmbeddingsAttentionBridge(
103 name="self_attn",
104 config=self.cfg,
105 submodules={
106 "q": LinearBridge(name="q_proj"),
107 "k": LinearBridge(name="k_proj"),
108 "v": LinearBridge(name="v_proj"),
109 "o": LinearBridge(name="o_proj"),
110 # Per-head LN containers, present only when qk_layernorm=True
111 # (stablelm-2-12b); applied post-reshape like HF.
112 "q_norm": GeneralizedComponent(name="q_layernorm", optional=True),
113 "k_norm": GeneralizedComponent(name="k_layernorm", optional=True),
114 },
115 requires_attention_mask=True,
116 requires_position_embeddings=True,
117 )
118 block_submodules["mlp"] = self._gated_mlp()
120 # StableLM has both parallel (use_parallel_residual=True) and sequential variants.
121 block_cls = ParallelBlockBridge if use_parallel_residual else BlockBridge
123 self.component_mapping = {
124 "embed": EmbeddingBridge(name="model.embed_tokens"),
125 "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb"),
126 "blocks": block_cls(
127 name="model.layers",
128 submodules=block_submodules,
129 ),
130 "ln_final": NormalizationBridge(
131 name="model.norm",
132 config=self.cfg,
133 use_native_layernorm_autograd=True,
134 ),
135 "unembed": UnembeddingBridge(name="lm_head", config=self.cfg),
136 }