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

1"""StableLM architecture adapter.""" 

2 

3from typing import Any 

4 

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) 

23 

24 

25class StableLmArchitectureAdapter(ArchitectureAdapter): 

26 """Architecture adapter for StableLM models. 

27 

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). 

31 

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) 

37 

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 """ 

48 

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

50 """Initialize the StableLM architecture adapter.""" 

51 super().__init__(cfg) 

52 

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" 

63 

64 n_kv_heads = getattr(self.cfg, "n_key_value_heads", None) or self.cfg.n_heads 

65 

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 } 

79 

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) 

88 

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() 

119 

120 # StableLM has both parallel (use_parallel_residual=True) and sequential variants. 

121 block_cls = ParallelBlockBridge if use_parallel_residual else BlockBridge 

122 

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 }