From 3d382ee071e39f6d65cf41e70eb5a80d2c75716a Mon Sep 17 00:00:00 2001 From: Sanjid Mzi Date: Wed, 22 Jul 2026 20:15:16 +0000 Subject: [PATCH] Add Starcoder2 architecture adapter --- .../model_bridge/test_starcoder2_adapter.py | 105 ++++++++ .../test_starcoder2_adapter.py | 230 ++++++++++++++++++ .../factories/architecture_adapter_factory.py | 2 + .../supported_architectures/__init__.py | 4 + .../supported_architectures/starcoder2.py | 131 ++++++++++ .../tools/model_registry/__init__.py | 2 + .../tools/model_registry/generate_report.py | 1 + 7 files changed, 475 insertions(+) create mode 100644 tests/integration/model_bridge/test_starcoder2_adapter.py create mode 100644 tests/unit/model_bridge/supported_architectures/test_starcoder2_adapter.py create mode 100644 transformer_lens/model_bridge/supported_architectures/starcoder2.py diff --git a/tests/integration/model_bridge/test_starcoder2_adapter.py b/tests/integration/model_bridge/test_starcoder2_adapter.py new file mode 100644 index 000000000..f5b724ee0 --- /dev/null +++ b/tests/integration/model_bridge/test_starcoder2_adapter.py @@ -0,0 +1,105 @@ +"""Integration tests for the Starcoder2 architecture adapter. + +Builds a tiny ``Starcoder2ForCausalLM`` from config (no download) and checks that +the TransformerBridge reproduces its structure and — most importantly — its +forward-pass logits. Starcoder2 exercises a LayerNorm + non-gated-MLP + biased ++ GQA decoder, so a logit match is strong evidence the adapter's weight +conversions and component mapping are correct. +""" + +import tempfile + +import pytest +import torch +from transformers import AutoTokenizer, Starcoder2Config, Starcoder2ForCausalLM + +from transformer_lens.model_bridge.bridge import TransformerBridge + +D_MODEL = 64 +N_HEADS = 8 +N_KV_HEADS = 2 +N_LAYERS = 2 + + +def _make_hf_model() -> Starcoder2ForCausalLM: + cfg = Starcoder2Config( + hidden_size=D_MODEL, + intermediate_size=128, + num_hidden_layers=N_LAYERS, + num_attention_heads=N_HEADS, + num_key_value_heads=N_KV_HEADS, + max_position_embeddings=128, + vocab_size=1000, + use_bias=True, + ) + torch.manual_seed(0) + model = Starcoder2ForCausalLM(cfg) + model.eval() + return model + + +@pytest.fixture(scope="module") +def hf_and_bridge(): + """Return a matched (hf_model, bridge) pair loaded from the same weights.""" + hf_model = _make_hf_model() + with tempfile.TemporaryDirectory() as tmpdir: + hf_model.save_pretrained(tmpdir) + tok = AutoTokenizer.from_pretrained("gpt2") + tok.save_pretrained(tmpdir) + bridge = TransformerBridge.boot_transformers(tmpdir, device="cpu") + bridge.eval() + return hf_model, bridge + + +def _tokens() -> torch.Tensor: + return torch.tensor([[5, 12, 7, 99, 3, 42, 8, 1]]) + + +# --------------------------------------------------------------------------- +# Structure +# --------------------------------------------------------------------------- + + +class TestStarcoder2BridgeStructure: + def test_block_count(self, hf_and_bridge) -> None: + _, bridge = hf_and_bridge + assert len(bridge.blocks) == N_LAYERS + + def test_attention_is_separate_qkv(self, hf_and_bridge) -> None: + _, bridge = hf_and_bridge + attn = bridge.blocks[0].attn + for proj in ("q", "k", "v", "o"): + assert hasattr(attn, proj) + + def test_mlp_is_non_gated(self, hf_and_bridge) -> None: + _, bridge = hf_and_bridge + mlp = bridge.blocks[0].mlp + assert hasattr(mlp, "in") + assert hasattr(mlp, "out") + assert not hasattr(mlp, "gate") + + +# --------------------------------------------------------------------------- +# Forward pass +# --------------------------------------------------------------------------- + + +class TestStarcoder2ForwardPass: + def test_forward_returns_correct_shape(self, hf_and_bridge) -> None: + _, bridge = hf_and_bridge + logits = bridge(_tokens()) + assert logits.shape == (1, 8, 1000) + + def test_forward_matches_hf(self, hf_and_bridge) -> None: + """The decisive test: bridge logits must equal HuggingFace's up to float noise.""" + hf_model, bridge = hf_and_bridge + tokens = _tokens() + with torch.no_grad(): + hf_logits = hf_model(tokens).logits + bridge_logits = bridge(tokens) + assert torch.allclose(hf_logits, bridge_logits, atol=1e-4, rtol=1e-4) + + def test_forward_produces_no_nans(self, hf_and_bridge) -> None: + _, bridge = hf_and_bridge + logits = bridge(_tokens()) + assert not torch.isnan(logits).any() diff --git a/tests/unit/model_bridge/supported_architectures/test_starcoder2_adapter.py b/tests/unit/model_bridge/supported_architectures/test_starcoder2_adapter.py new file mode 100644 index 000000000..380e03c5d --- /dev/null +++ b/tests/unit/model_bridge/supported_architectures/test_starcoder2_adapter.py @@ -0,0 +1,230 @@ +"""Unit tests for Starcoder2ArchitectureAdapter. + +Tests cover: +- Config flags set by the adapter (LayerNorm, non-gated MLP, rotary) +- Component mapping structure (bridge types and HF module names) +- The GQA-aware weight-conversion key set, including the per-head bias + rearrangements that distinguish Starcoder2 from the bias-free Llama family + +Behavioural coverage (forward pass vs HuggingFace) lives in +``tests/integration/model_bridge/test_starcoder2_adapter.py``. +""" + +import pytest + +from transformer_lens.config import TransformerBridgeConfig +from transformer_lens.model_bridge.generalized_components import ( + BlockBridge, + EmbeddingBridge, + LinearBridge, + MLPBridge, + NormalizationBridge, + PositionEmbeddingsAttentionBridge, + RotaryEmbeddingBridge, + UnembeddingBridge, +) +from transformer_lens.model_bridge.supported_architectures.starcoder2 import ( + Starcoder2ArchitectureAdapter, +) + +# --------------------------------------------------------------------------- +# Helpers / fixtures +# --------------------------------------------------------------------------- + +N_HEADS = 8 +N_KV_HEADS = 2 +D_MODEL = 64 +D_MLP = 256 +N_LAYERS = 2 +N_CTX = 256 +D_VOCAB = 1000 + + +def _make_cfg(n_kv_heads: int | None = N_KV_HEADS) -> TransformerBridgeConfig: + """Return a minimal TransformerBridgeConfig for Starcoder2 adapter tests.""" + return TransformerBridgeConfig( + d_model=D_MODEL, + d_head=D_MODEL // N_HEADS, + n_layers=N_LAYERS, + n_ctx=N_CTX, + n_heads=N_HEADS, + d_vocab=D_VOCAB, + d_mlp=D_MLP, + n_key_value_heads=n_kv_heads, + architecture="Starcoder2ForCausalLM", + ) + + +@pytest.fixture +def cfg() -> TransformerBridgeConfig: + return _make_cfg() + + +@pytest.fixture +def adapter(cfg: TransformerBridgeConfig) -> Starcoder2ArchitectureAdapter: + return Starcoder2ArchitectureAdapter(cfg) + + +# --------------------------------------------------------------------------- +# Config flag tests +# --------------------------------------------------------------------------- + + +class TestStarcoder2AdapterConfig: + """Tests that the adapter sets the correct config flags.""" + + def test_normalization_type_is_layernorm(self, adapter: Starcoder2ArchitectureAdapter) -> None: + """Starcoder2 uses LayerNorm (with bias), not RMSNorm.""" + assert adapter.cfg.normalization_type == "LN" + + def test_positional_embedding_type(self, adapter: Starcoder2ArchitectureAdapter) -> None: + assert adapter.cfg.positional_embedding_type == "rotary" + + def test_not_final_rms(self, adapter: Starcoder2ArchitectureAdapter) -> None: + assert adapter.cfg.final_rms is False + + def test_mlp_is_not_gated(self, adapter: Starcoder2ArchitectureAdapter) -> None: + """Starcoder2 uses a plain c_fc -> c_proj MLP, not a gated MLP.""" + assert adapter.cfg.gated_mlp is False + + def test_not_attn_only(self, adapter: Starcoder2ArchitectureAdapter) -> None: + assert adapter.cfg.attn_only is False + + def test_n_key_value_heads_propagated(self, adapter: Starcoder2ArchitectureAdapter) -> None: + assert adapter.cfg.n_key_value_heads == N_KV_HEADS + + +# --------------------------------------------------------------------------- +# Component mapping tests +# --------------------------------------------------------------------------- + + +class TestStarcoder2AdapterComponentMapping: + """Tests that component_mapping has the correct bridge types and HF module names.""" + + def test_top_level_keys(self, adapter: Starcoder2ArchitectureAdapter) -> None: + assert set(adapter.component_mapping.keys()) == { + "embed", + "rotary_emb", + "blocks", + "ln_final", + "unembed", + } + + def test_top_level_bridge_types(self, adapter: Starcoder2ArchitectureAdapter) -> None: + mapping = adapter.component_mapping + assert isinstance(mapping["embed"], EmbeddingBridge) + assert isinstance(mapping["rotary_emb"], RotaryEmbeddingBridge) + assert isinstance(mapping["blocks"], BlockBridge) + assert isinstance(mapping["ln_final"], NormalizationBridge) + assert isinstance(mapping["unembed"], UnembeddingBridge) + + def test_top_level_hf_paths(self, adapter: Starcoder2ArchitectureAdapter) -> None: + mapping = adapter.component_mapping + assert mapping["embed"].name == "model.embed_tokens" + assert mapping["rotary_emb"].name == "model.rotary_emb" + assert mapping["blocks"].name == "model.layers" + assert mapping["ln_final"].name == "model.norm" + assert mapping["unembed"].name == "lm_head" + + def test_block_submodule_keys(self, adapter: Starcoder2ArchitectureAdapter) -> None: + """Sequential pre-norm block: two LayerNorms, not the single norm of a parallel block.""" + blocks = adapter.component_mapping["blocks"] + assert set(blocks.submodules.keys()) == {"ln1", "ln2", "attn", "mlp"} + + def test_block_bridge_types(self, adapter: Starcoder2ArchitectureAdapter) -> None: + blocks = adapter.component_mapping["blocks"] + assert isinstance(blocks.submodules["ln1"], NormalizationBridge) + assert isinstance(blocks.submodules["ln2"], NormalizationBridge) + assert isinstance(blocks.submodules["attn"], PositionEmbeddingsAttentionBridge) + assert isinstance(blocks.submodules["mlp"], MLPBridge) + + def test_block_hf_paths(self, adapter: Starcoder2ArchitectureAdapter) -> None: + blocks = adapter.component_mapping["blocks"] + assert blocks.submodules["ln1"].name == "input_layernorm" + assert blocks.submodules["ln2"].name == "post_attention_layernorm" + assert blocks.submodules["attn"].name == "self_attn" + assert blocks.submodules["mlp"].name == "mlp" + + +# --------------------------------------------------------------------------- +# Attention mapping tests +# --------------------------------------------------------------------------- + + +class TestStarcoder2AdapterAttention: + """Tests the separate-QKV attention mapping.""" + + def test_attention_submodule_keys(self, adapter: Starcoder2ArchitectureAdapter) -> None: + """Unlike GPTBigCode's combined c_attn, Starcoder2 uses separate q/k/v projections.""" + attn = adapter.component_mapping["blocks"].submodules["attn"] + assert set(attn.submodules.keys()) == {"q", "k", "v", "o"} + + def test_attention_hf_paths(self, adapter: Starcoder2ArchitectureAdapter) -> None: + attn = adapter.component_mapping["blocks"].submodules["attn"] + assert attn.submodules["q"].name == "q_proj" + assert attn.submodules["k"].name == "k_proj" + assert attn.submodules["v"].name == "v_proj" + assert attn.submodules["o"].name == "o_proj" + + def test_attention_linear_bridge_types(self, adapter: Starcoder2ArchitectureAdapter) -> None: + attn = adapter.component_mapping["blocks"].submodules["attn"] + for submodule in attn.submodules.values(): + assert isinstance(submodule, LinearBridge) + + +# --------------------------------------------------------------------------- +# MLP mapping tests +# --------------------------------------------------------------------------- + + +class TestStarcoder2AdapterMLP: + """Tests the non-gated c_fc -> c_proj MLP mapping.""" + + def test_mlp_submodule_keys(self, adapter: Starcoder2ArchitectureAdapter) -> None: + """Non-gated MLP has only in/out, no gate.""" + mlp = adapter.component_mapping["blocks"].submodules["mlp"] + assert set(mlp.submodules.keys()) == {"in", "out"} + + def test_mlp_hf_paths(self, adapter: Starcoder2ArchitectureAdapter) -> None: + mlp = adapter.component_mapping["blocks"].submodules["mlp"] + assert mlp.submodules["in"].name == "c_fc" + assert mlp.submodules["out"].name == "c_proj" + + def test_mlp_linear_bridge_types(self, adapter: Starcoder2ArchitectureAdapter) -> None: + mlp = adapter.component_mapping["blocks"].submodules["mlp"] + for submodule in mlp.submodules.values(): + assert isinstance(submodule, LinearBridge) + + +# --------------------------------------------------------------------------- +# Weight conversion key tests +# --------------------------------------------------------------------------- + + +class TestStarcoder2AdapterWeightConversions: + """Tests the GQA-aware Q/K/V/O weight and bias conversions.""" + + def test_conversion_key_set(self, adapter: Starcoder2ArchitectureAdapter) -> None: + """Starcoder2 has biases, so q/k/v get both a weight and a per-head bias conversion. + + The output projection ``o`` has no per-head bias (its bias stays ``[d_model]``), + so there is no ``blocks.{i}.attn.o.bias`` conversion. + """ + assert set(adapter.weight_processing_conversions.keys()) == { + "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_no_output_bias_conversion(self, adapter: Starcoder2ArchitectureAdapter) -> None: + assert "blocks.{i}.attn.o.bias" not in adapter.weight_processing_conversions + + def test_missing_kv_heads_falls_back_to_n_heads(self) -> None: + """Without n_key_value_heads the adapter still builds a full conversion set (MHA).""" + adapter = Starcoder2ArchitectureAdapter(_make_cfg(n_kv_heads=None)) + assert "blocks.{i}.attn.k.bias" in adapter.weight_processing_conversions diff --git a/transformer_lens/factories/architecture_adapter_factory.py b/transformer_lens/factories/architecture_adapter_factory.py index 1e1f96bd0..a0bd9a4eb 100644 --- a/transformer_lens/factories/architecture_adapter_factory.py +++ b/transformer_lens/factories/architecture_adapter_factory.py @@ -88,6 +88,7 @@ RWKV7ArchitectureAdapter, SmolLM3ArchitectureAdapter, StableLmArchitectureAdapter, + Starcoder2ArchitectureAdapter, T5ArchitectureAdapter, T5Gemma2ArchitectureAdapter, T5GemmaArchitectureAdapter, @@ -178,6 +179,7 @@ "RWKV7ForCausalLM": RWKV7ArchitectureAdapter, "SmolLM3ForCausalLM": SmolLM3ArchitectureAdapter, "StableLmForCausalLM": StableLmArchitectureAdapter, + "Starcoder2ForCausalLM": Starcoder2ArchitectureAdapter, "T5ForConditionalGeneration": T5ArchitectureAdapter, "MT5ForConditionalGeneration": T5ArchitectureAdapter, "T5GemmaForConditionalGeneration": T5GemmaArchitectureAdapter, diff --git a/transformer_lens/model_bridge/supported_architectures/__init__.py b/transformer_lens/model_bridge/supported_architectures/__init__.py index dcec4fdc6..e1937bd6c 100644 --- a/transformer_lens/model_bridge/supported_architectures/__init__.py +++ b/transformer_lens/model_bridge/supported_architectures/__init__.py @@ -235,6 +235,9 @@ from transformer_lens.model_bridge.supported_architectures.stablelm import ( StableLmArchitectureAdapter, ) +from transformer_lens.model_bridge.supported_architectures.starcoder2 import ( + Starcoder2ArchitectureAdapter, +) from transformer_lens.model_bridge.supported_architectures.t5 import ( T5ArchitectureAdapter, ) @@ -330,6 +333,7 @@ "RWKV7ArchitectureAdapter", "SmolLM3ArchitectureAdapter", "StableLmArchitectureAdapter", + "Starcoder2ArchitectureAdapter", "T5ArchitectureAdapter", "T5GemmaArchitectureAdapter", "T5Gemma2ArchitectureAdapter", diff --git a/transformer_lens/model_bridge/supported_architectures/starcoder2.py b/transformer_lens/model_bridge/supported_architectures/starcoder2.py new file mode 100644 index 000000000..8be7b133d --- /dev/null +++ b/transformer_lens/model_bridge/supported_architectures/starcoder2.py @@ -0,0 +1,131 @@ +"""Starcoder2 architecture adapter.""" + +from typing import Any + +from transformer_lens.conversion_utils.conversion_steps import RearrangeTensorConversion +from transformer_lens.conversion_utils.param_processing_conversion import ( + ParamProcessingConversion, +) +from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter +from transformer_lens.model_bridge.generalized_components import ( + BlockBridge, + EmbeddingBridge, + LinearBridge, + MLPBridge, + NormalizationBridge, + PositionEmbeddingsAttentionBridge, + RotaryEmbeddingBridge, + UnembeddingBridge, +) + + +class Starcoder2ArchitectureAdapter(ArchitectureAdapter): + """Architecture adapter for Starcoder2 models (``Starcoder2ForCausalLM``). + + Starcoder2 is a Llama-shaped decoder — sequential pre-norm blocks, GQA, and + rotary position embeddings — but differs from Llama in three ways it shares + with its GPTBigCode (StarCoder) predecessor: + + - **LayerNorm (with bias)** rather than RMSNorm (``normalization_type="LN"``). + - **Non-gated GELU MLP** (``c_fc`` -> ``c_proj``) rather than a gated MLP. + - **Biases on every attention and MLP projection** (``use_bias=True``). + + Unlike GPTBigCode it uses separate ``q_proj``/``k_proj``/``v_proj`` projections + and rotary embeddings instead of a combined ``c_attn`` and learned positions. + """ + + def __init__(self, cfg: Any) -> None: + """Initialize the Starcoder2 architecture adapter.""" + super().__init__(cfg) + + # Config variables for weight processing + self.cfg.normalization_type = "LN" + self.cfg.positional_embedding_type = "rotary" + self.cfg.final_rms = False + self.cfg.gated_mlp = False + self.cfg.attn_only = False + + self.default_config = { + "d_model": cfg.d_model, + "d_head": cfg.d_model // cfg.n_heads, + "n_heads": cfg.n_heads, + "n_layers": cfg.n_layers, + "d_vocab": cfg.d_vocab, + } + + # GQA: Starcoder2 uses num_key_value_heads (< n_heads) for K/V projections. + n_kv_heads = None + if hasattr(cfg, "n_key_value_heads") and cfg.n_key_value_heads is not None: + self.default_config["n_key_value_heads"] = cfg.n_key_value_heads + self.cfg.n_key_value_heads = cfg.n_key_value_heads + n_kv_heads = cfg.n_key_value_heads + + # Standard Q/K/V/O weight rearrangement, plus per-head bias rearrangement + # (Starcoder2 has biases on the attention projections, which Llama does not). + # K/V use the GQA head count; O has no per-head bias (its bias stays [d_model]). + n_kv = n_kv_heads if n_kv_heads is not None else self.cfg.n_heads + self.weight_processing_conversions = { + **self._qkvo_weight_conversions(n_kv_heads=n_kv_heads), + "blocks.{i}.attn.q.bias": ParamProcessingConversion( + tensor_conversion=RearrangeTensorConversion("(n h) -> n h", n=self.cfg.n_heads), + ), + "blocks.{i}.attn.k.bias": ParamProcessingConversion( + tensor_conversion=RearrangeTensorConversion("(n h) -> n h", n=n_kv), + ), + "blocks.{i}.attn.v.bias": ParamProcessingConversion( + tensor_conversion=RearrangeTensorConversion("(n h) -> n h", n=n_kv), + ), + } + + self.component_mapping = { + "embed": EmbeddingBridge(name="model.embed_tokens"), + "rotary_emb": RotaryEmbeddingBridge(name="model.rotary_emb"), + "blocks": BlockBridge( + name="model.layers", + submodules={ + "ln1": NormalizationBridge(name="input_layernorm", config=self.cfg), + "ln2": NormalizationBridge(name="post_attention_layernorm", config=self.cfg), + "attn": PositionEmbeddingsAttentionBridge( + name="self_attn", + config=self.cfg, + submodules={ + "q": LinearBridge(name="q_proj"), + "k": LinearBridge(name="k_proj"), + "v": LinearBridge(name="v_proj"), + "o": LinearBridge(name="o_proj"), + }, + requires_attention_mask=True, + requires_position_embeddings=True, + ), + "mlp": MLPBridge( + name="mlp", + submodules={ + "in": LinearBridge(name="c_fc"), + "out": LinearBridge(name="c_proj"), + }, + ), + }, + ), + "ln_final": NormalizationBridge(name="model.norm", config=self.cfg), + "unembed": UnembeddingBridge(name="lm_head", config=self.cfg), + } + + def setup_component_testing(self, hf_model: Any, bridge_model: Any = None) -> None: + """Set up rotary embedding references for Starcoder2 component testing. + + Starcoder2 uses RoPE. Mirror the Llama-family setup: set the shared + ``model.rotary_emb`` reference on each attention bridge instance. + + Args: + hf_model: The HuggingFace Starcoder2 model instance. + bridge_model: The TransformerBridge model, if available. + """ + rotary_emb = hf_model.model.rotary_emb + + if bridge_model is not None and hasattr(bridge_model, "blocks"): + for block in bridge_model.blocks: + if hasattr(block, "attn"): + block.attn.set_rotary_emb(rotary_emb) + + attn_bridge = self.get_generalized_component("blocks.0.attn") + attn_bridge.set_rotary_emb(rotary_emb) diff --git a/transformer_lens/tools/model_registry/__init__.py b/transformer_lens/tools/model_registry/__init__.py index e606aa4c3..3a508f7b6 100644 --- a/transformer_lens/tools/model_registry/__init__.py +++ b/transformer_lens/tools/model_registry/__init__.py @@ -120,6 +120,7 @@ "RWKV7ForCausalLM", "SmolLM3ForCausalLM", "StableLmForCausalLM", + "Starcoder2ForCausalLM", "T5ForConditionalGeneration", "MT5ForConditionalGeneration", "T5GemmaForConditionalGeneration", @@ -207,6 +208,7 @@ "RWKV7ForCausalLM": ["fla-hub"], "SmolLM3ForCausalLM": ["HuggingFaceTB"], "StableLmForCausalLM": ["stabilityai"], + "Starcoder2ForCausalLM": ["bigcode"], "T5ForConditionalGeneration": ["google-t5", "google", "Salesforce", "MBZUAI"], "T5GemmaForConditionalGeneration": ["google"], "T5Gemma2ForConditionalGeneration": ["google"], diff --git a/transformer_lens/tools/model_registry/generate_report.py b/transformer_lens/tools/model_registry/generate_report.py index 0a4f8cb05..85707de20 100644 --- a/transformer_lens/tools/model_registry/generate_report.py +++ b/transformer_lens/tools/model_registry/generate_report.py @@ -63,6 +63,7 @@ "OlmoeForCausalLM": "Allen AI's OLMoE Mixture of Experts model", "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)", "T5GemmaForConditionalGeneration": "Google's T5Gemma encoder-decoder model with Gemma-style RoPE, GQA, and gated MLP", "T5Gemma2ForConditionalGeneration": "Google's T5Gemma2 multimodal encoder-decoder model with merged self+cross decoder attention, QK-norm, and dual RoPE (text-only bridge support)",