diff --git a/tests/unit/model_bridge/supported_architectures/test_starcoder2_adapter.py b/tests/unit/model_bridge/supported_architectures/test_starcoder2_adapter.py index 7718cb1c8..6e021fb11 100644 --- a/tests/unit/model_bridge/supported_architectures/test_starcoder2_adapter.py +++ b/tests/unit/model_bridge/supported_architectures/test_starcoder2_adapter.py @@ -13,17 +13,23 @@ SUPPORTED_ARCHITECTURES, ) from transformer_lens.model_bridge.generalized_components import ( + BlockBridge, + EmbeddingBridge, GatedMLPBridge, + LinearBridge, MLPBridge, NormalizationBridge, + PositionEmbeddingsAttentionBridge, RMSNormalizationBridge, + RotaryEmbeddingBridge, + UnembeddingBridge, ) from transformer_lens.model_bridge.supported_architectures.starcoder2 import ( Starcoder2ArchitectureAdapter, ) -def _make_cfg() -> TransformerBridgeConfig: +def _make_cfg(n_key_value_heads: int | None = 2) -> TransformerBridgeConfig: return make_bridge_cfg( "Starcoder2ForCausalLM", d_model=64, @@ -33,7 +39,7 @@ def _make_cfg() -> TransformerBridgeConfig: n_heads=4, d_mlp=256, d_vocab=512, - n_key_value_heads=2, + n_key_value_heads=n_key_value_heads, default_prepend_bos=True, ) @@ -59,14 +65,56 @@ def test_qkv_bias_conversions_use_kv_head_count(self, adapter): assert key in conv, f"missing {key}" assert conv[key].tensor_conversion.axes_lengths["h"] == adapter.cfg.n_key_value_heads + def test_conversion_key_set(self, adapter): + """q/k/v get both a weight and a per-head bias conversion; o's bias stays + [d_model], so it must not gain a per-head reshape.""" + assert set(adapter.weight_processing_conversions) == { + "blocks.{i}.attn.q.weight", + "blocks.{i}.attn.k.weight", + "blocks.{i}.attn.v.weight", + "blocks.{i}.attn.o.weight", + "blocks.{i}.attn.q.bias", + "blocks.{i}.attn.k.bias", + "blocks.{i}.attn.v.bias", + } + + def test_missing_kv_heads_falls_back_to_n_heads(self): + """MHA checkpoints omit n_key_value_heads; the bias reshapes must still + be emitted, split by n_heads.""" + adapter = Starcoder2ArchitectureAdapter(_make_cfg(n_key_value_heads=None)) + conv = adapter.weight_processing_conversions + assert conv["blocks.{i}.attn.k.bias"].tensor_conversion.axes_lengths["h"] == 4 + class TestStarcoder2ComponentMapping: + def test_top_level_mapping(self, adapter): + mapping = adapter.component_mapping + assert isinstance(mapping["embed"], EmbeddingBridge) + assert isinstance(mapping["rotary_emb"], RotaryEmbeddingBridge) + assert isinstance(mapping["blocks"], BlockBridge) + assert isinstance(mapping["unembed"], UnembeddingBridge) + assert mapping["embed"].name == "model.embed_tokens" + assert mapping["blocks"].name == "model.layers" + assert mapping["unembed"].name == "lm_head" + + def test_attention_is_separate_qkvo(self, adapter): + """Separate q/k/v/o projections — unlike GPTBigCode's fused c_attn.""" + attn = adapter.component_mapping["blocks"].submodules["attn"] + assert isinstance(attn, PositionEmbeddingsAttentionBridge) + assert attn.name == "self_attn" + expected = {"q": "q_proj", "k": "k_proj", "v": "v_proj", "o": "o_proj"} + assert set(attn.submodules) == set(expected) + for key, hf_name in expected.items(): + assert isinstance(attn.submodules[key], LinearBridge) + assert attn.submodules[key].name == hf_name + def test_norms_are_plain_layernorm(self, adapter): """StarCoder2 uses nn.LayerNorm despite its llama-like shape.""" submodules = adapter.component_mapping["blocks"].submodules - assert isinstance(submodules["ln1"], NormalizationBridge) - assert not isinstance(submodules["ln1"], RMSNormalizationBridge) - assert submodules["ln1"].name == "input_layernorm" + for key, hf_name in (("ln1", "input_layernorm"), ("ln2", "post_attention_layernorm")): + assert isinstance(submodules[key], NormalizationBridge) + assert not isinstance(submodules[key], RMSNormalizationBridge) + assert submodules[key].name == hf_name assert adapter.component_mapping["ln_final"].name == "model.norm" def test_mlp_is_plain_c_fc_c_proj(self, adapter): diff --git a/transformer_lens/model_bridge/supported_architectures/starcoder2.py b/transformer_lens/model_bridge/supported_architectures/starcoder2.py index fa4acf077..87600d8a3 100644 --- a/transformer_lens/model_bridge/supported_architectures/starcoder2.py +++ b/transformer_lens/model_bridge/supported_architectures/starcoder2.py @@ -3,6 +3,9 @@ BigCode's StarCoder2 (``Starcoder2ForCausalLM``): pre-norm decoder with plain LayerNorm (not RMS), separate biased q/k/v/o projections, GQA, RoPE, and a non-gated ``c_fc``/``c_proj`` MLP. + +Not a drop-in for its GPTBigCode predecessor: that one fuses q/k/v into a +single ``c_attn`` and uses learned positions, so it needs a different bridge. """ from typing import Any diff --git a/transformer_lens/tools/model_registry/generate_report.py b/transformer_lens/tools/model_registry/generate_report.py index 559425cf2..2f39db8d7 100644 --- a/transformer_lens/tools/model_registry/generate_report.py +++ b/transformer_lens/tools/model_registry/generate_report.py @@ -79,6 +79,7 @@ "OlmoHybridForCausalLM": "Allen AI's OLMo Hybrid (GatedDeltaNet linear attention + full attention)", "StableLmForCausalLM": "Stability AI's StableLM model", "SmolLM3ForCausalLM": "Hugging Face's SmolLM3 compact open model with NoPE layers", + "Starcoder2ForCausalLM": "BigCode's StarCoder2 code generation model", "T5ForConditionalGeneration": "Google's T5 encoder-decoder model (partial support)", "T5WithLMHeadModel": "Legacy T5 class name on old google-t5 checkpoints (t5-3b, t5-11b)", "T5GemmaForConditionalGeneration": "Google's T5Gemma encoder-decoder model with Gemma-style RoPE, GQA, and gated MLP",