Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
Original file line number Diff line number Diff line change
Expand Up @@ -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,
Expand All @@ -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,
)

Expand All @@ -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):
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -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
Expand Down
1 change: 1 addition & 0 deletions transformer_lens/tools/model_registry/generate_report.py
Original file line number Diff line number Diff line change
Expand Up @@ -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",
Expand Down
Loading