diff --git a/tests/unit/model_bridge/generalized_components/test_moe_dense_dispatch.py b/tests/unit/model_bridge/generalized_components/test_moe_dense_dispatch.py index 37123f632..55843478b 100644 --- a/tests/unit/model_bridge/generalized_components/test_moe_dense_dispatch.py +++ b/tests/unit/model_bridge/generalized_components/test_moe_dense_dispatch.py @@ -27,19 +27,29 @@ D_MODEL, D_MLP = 8, 16 +# Some architectures name gate/up/down differently +DEFAULT_DENSE_NAMES: tuple[str, str, str] = ("gate_proj", "up_proj", "down_proj") +DENSE_PROJECTION_NAMES: dict[str, tuple[str, str, str]] = { + "Lfm2MoeForCausalLM": ("w1", "w3", "w2"), +} + class _DenseMLP(nn.Module): """Standard SwiGLU gated MLP — the dense-prefix layer shape.""" - def __init__(self) -> None: + def __init__(self, names: tuple[str, str, str] = DEFAULT_DENSE_NAMES) -> None: super().__init__() - self.gate_proj = nn.Linear(D_MODEL, D_MLP, bias=False) - self.up_proj = nn.Linear(D_MODEL, D_MLP, bias=False) - self.down_proj = nn.Linear(D_MLP, D_MODEL, bias=False) + self.names = names + gate, up, down = names + setattr(self, gate, nn.Linear(D_MODEL, D_MLP, bias=False)) + setattr(self, up, nn.Linear(D_MODEL, D_MLP, bias=False)) + setattr(self, down, nn.Linear(D_MLP, D_MODEL, bias=False)) def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: - return self.down_proj( - torch.nn.functional.silu(self.gate_proj(hidden_states)) * self.up_proj(hidden_states) + gate, up, down = self.names + return getattr(self, down)( + torch.nn.functional.silu(getattr(self, gate)(hidden_states)) + * getattr(self, up)(hidden_states) ) @@ -229,6 +239,7 @@ def test_undeclared_dense_projections_leave_moe_mapping(self) -> None: "Qwen3VLMoeForConditionalGeneration", "LLaDA2MoeModelLM", "LagunaForCausalLM", + "Lfm2MoeForCausalLM", "Llama4ForConditionalGeneration", ] @@ -270,7 +281,8 @@ def test_adapter_templates_bind_dense_layers_as_gated_mlps(architecture: str) -> blocks = adapter.component_mapping["blocks"] template = blocks.submodules.get("mlp") or blocks.submodules["feed_forward"] instance = copy.deepcopy(template) - instance.set_original_component(_DenseMLP()) + names = DENSE_PROJECTION_NAMES.get(architecture, DEFAULT_DENSE_NAMES) + instance.set_original_component(_DenseMLP(names)) assert instance._bound_dense is True, f"{architecture} did not declare dense projections" assert instance.hook_aliases["hook_pre"] == "dense_gate.hook_out" @@ -286,7 +298,9 @@ def test_dense_keys_read_the_projection_they_name(architecture: str) -> None: blocks = adapter.component_mapping["blocks"] template = blocks.submodules.get("mlp") or blocks.submodules["feed_forward"] bridge = copy.deepcopy(template) - module = _DenseMLP() + names = DENSE_PROJECTION_NAMES.get(architecture, DEFAULT_DENSE_NAMES) + gate_name, in_name, _ = names + module = _DenseMLP(names) bridge.set_original_component(module) setup_submodules(bridge, adapter, module) @@ -301,8 +315,8 @@ def test_dense_keys_read_the_projection_they_name(architecture: str) -> None: ) with torch.no_grad(): bridge(x) - expected_gate = module.gate_proj(x) - expected_in = module.up_proj(x) + expected_gate = getattr(module, gate_name)(x) + expected_in = getattr(module, in_name)(x) torch.testing.assert_close( captured["dense_gate"], diff --git a/tests/unit/model_bridge/supported_architectures/test_lfm2_moe_adapter.py b/tests/unit/model_bridge/supported_architectures/test_lfm2_moe_adapter.py index 3b06d9d1f..95b4476c3 100644 --- a/tests/unit/model_bridge/supported_architectures/test_lfm2_moe_adapter.py +++ b/tests/unit/model_bridge/supported_architectures/test_lfm2_moe_adapter.py @@ -1,85 +1,469 @@ -"""Unit tests for the Lfm2MoeArchitectureAdapter — no model downloads.""" +"""Unit tests for Lfm2MoeArchitectureAdapter. + +Tests cover: +- Config attribute validation +- Component mapping structure +- Weight conversion keys and rearrange patterns +- Architecture guards +- Setup component tests +""" + +from types import SimpleNamespace import pytest -from tests.unit.model_bridge.supported_architectures.helpers import make_bridge_cfg from transformer_lens.config import TransformerBridgeConfig +from transformer_lens.conversion_utils.conversion_steps import RearrangeTensorConversion +from transformer_lens.conversion_utils.param_processing_conversion import ( + ParamProcessingConversion, +) from transformer_lens.model_bridge.generalized_components import ( + BlockBridge, + DepthwiseConv1DBridge, EmbeddingBridge, + Lfm2ShortConvBridge, + LinearBridge, + MoEBridge, + PositionEmbeddingsAttentionBridge, RMSNormalizationBridge, + RotaryEmbeddingBridge, UnembeddingBridge, ) from transformer_lens.model_bridge.supported_architectures.lfm2_moe import ( Lfm2MoeArchitectureAdapter, - Lfm2MoeBlockBridge, + Lfm2MoeGateBridge, ) +# --------------------------------------------------------------------------- +# Fixtures & Helpers +# --------------------------------------------------------------------------- -@pytest.fixture(scope="class") -def cfg() -> TransformerBridgeConfig: - bridge_cfg = make_bridge_cfg( - "Lfm2MoeForCausalLM", - d_model=64, - d_head=16, - n_layers=4, - n_ctx=128, - n_heads=4, - n_key_value_heads=2, - d_vocab=256, - d_mlp=224, - default_prepend_bos=True, + +def _make_cfg( + n_heads: int = 32, + n_key_value_heads: int = 4, + d_model: int = 128, + n_layers: int = 4, + d_mlp: int = 256, + d_vocab: int = 1000, + n_ctx: int = 512, + layer_types: list[str] = ["conv", "full_attention", "conv", "conv"], + moe_intermediate_size=512, + num_experts=8, + experts_per_token=2, + rope_parameters={"rope_theta": 5_000_000, "rope_type": "default"}, +) -> TransformerBridgeConfig: + """Return a minimal TransformerBridgeConfig for Lfm2Moe adapter tests.""" + cfg = TransformerBridgeConfig( + d_model=d_model, + d_head=d_model // n_heads, + n_layers=n_layers, + n_ctx=n_ctx, + n_heads=n_heads, + n_key_value_heads=n_key_value_heads, + d_vocab=d_vocab, + d_mlp=d_mlp, + architecture="Lfm2MoeForCausalLM", ) - bridge_cfg.layer_types = ["conv", "conv", "full_attention", "conv"] - bridge_cfg.moe_intermediate_size = 56 - bridge_cfg.num_experts = 8 - bridge_cfg.experts_per_token = 2 - bridge_cfg.norm_eps = 1e-5 - bridge_cfg.rope_parameters = {"rope_theta": 5_000_000, "rope_type": "default"} - return bridge_cfg + + cfg.experts_per_token = (experts_per_token,) + cfg.layer_types = layer_types + cfg.moe_intermediate_size = moe_intermediate_size + cfg.num_experts = (num_experts,) + cfg.rope_parameters = rope_parameters + + return cfg + + +@pytest.fixture +def cfg() -> TransformerBridgeConfig: + return _make_cfg() -@pytest.fixture(scope="class") +@pytest.fixture def adapter(cfg: TransformerBridgeConfig) -> Lfm2MoeArchitectureAdapter: return Lfm2MoeArchitectureAdapter(cfg) +# For rotary embedding and attention implementation check in setup component testing + +layer_types = ["conv", "full_attention", "conv", "conv"] + + +def _fake_attn(layer_idx: int) -> SimpleNamespace: + """Per-layer self_attn with a mutable .config so the eager flip is observable.""" + return SimpleNamespace( + config=SimpleNamespace(_attn_implementation="sdpa"), + layer_idx=layer_idx, + ) + + +def _fake_hf_model(pos_emb: object, n_layers: int = 2) -> SimpleNamespace: + """Stub hf_model exposing everything setup_component_testing walks: + - .model.pos_emb -> rotary wiring (lfm2MoE names rotary "pos_emb") + - .config._attn_implementation -> top-level eager flip + - .model.layers[*].self_attn.config._attn_implementation -> per-layer eager flip + """ + return SimpleNamespace( + config=SimpleNamespace(_attn_implementation="sdpa"), + model=SimpleNamespace( + pos_emb=pos_emb, + layers=[ + SimpleNamespace(self_attn=_fake_attn(i)) + for i in range(n_layers) + if layer_types[i] == "full_attention" + ], + ), + ) + + +class DummyAttention: + def __init__(self) -> None: + self.pos_emb = None + + def set_rotary_emb(self, pos_emb: object) -> None: + self.pos_emb = pos_emb + + +class DummyBlock: + def __init__(self, has_attention: bool = True) -> None: + if has_attention: + self.attn = DummyAttention() + + +class DummyBridgeModel: + def __init__(self, blocks: list[DummyBlock]) -> None: + self.blocks = blocks + + +# --------------------------------------------------------------------------- +# Config attribute tests +# --------------------------------------------------------------------------- + + class TestLfm2MoeAdapterConfig: - def test_norm_and_rope_config(self, adapter: Lfm2MoeArchitectureAdapter) -> None: - assert adapter.cfg.normalization_type == "RMS" - assert adapter.cfg.positional_embedding_type == "rotary" - assert adapter.cfg.eps == 1e-5 + """Adapter must set all required config flags to the values Lfm2Moe expects.""" + + def test_attn_implementation_is_eager(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + assert adapter.cfg.attn_implementation == "eager" + + def test_act_fn_is_silu(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + assert adapter.cfg.act_fn == "silu" + + def test_rotary_base_value(self, adapter: Lfm2MoeArchitectureAdapter) -> None: assert adapter.cfg.rotary_base == 5_000_000 - def test_default_prepend_bos_is_false(self, adapter: Lfm2MoeArchitectureAdapter) -> None: - assert adapter.cfg.default_prepend_bos is False +# --------------------------------------------------------------------------- +# Component mapping structure tests +# --------------------------------------------------------------------------- + + +class TestLfm2MoeAdapterComponentMapping: + """Component mapping must have the correct bridge types and HF module names.""" + + # -- Top-level keys -- + + def test_embed_is_embedding_bridge(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + assert isinstance(adapter.component_mapping["embed"], EmbeddingBridge) + + def test_embed_name(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + assert adapter.component_mapping["embed"].name == "model.embed_tokens" -class TestLfm2MoeComponentMapping: - def test_has_residual_only_top_level_mapping(self, adapter: Lfm2MoeArchitectureAdapter) -> None: - mapping = adapter.component_mapping - assert mapping is not None - assert set(mapping) == {"embed", "blocks", "ln_final", "unembed"} + def test_rotary_emb_is_rotary_embedding_bridge( + self, adapter: Lfm2MoeArchitectureAdapter + ) -> None: + assert isinstance(adapter.component_mapping["rotary_emb"], RotaryEmbeddingBridge) + + def test_rotary_emb_name(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + assert adapter.component_mapping["rotary_emb"].name == "model.pos_emb" + + def test_blocks_is_block_bridge(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + """Sequential attn/conv & MLP requires BlockBridge, not ParallelBlockBridge.""" + assert isinstance(adapter.component_mapping["blocks"], BlockBridge) + + def test_blocks_name(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + assert adapter.component_mapping["blocks"].name == "model.layers" + + def test_ln_final_is_rms_normalization_bridge( + self, adapter: Lfm2MoeArchitectureAdapter + ) -> None: + assert isinstance(adapter.component_mapping["ln_final"], RMSNormalizationBridge) + + def test_ln_final_name(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + assert adapter.component_mapping["ln_final"].name == "model.embedding_norm" + + def test_unembed_is_unembedding_bridge(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + assert isinstance(adapter.component_mapping["unembed"], UnembeddingBridge) + + def test_unembed_name(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + assert adapter.component_mapping["unembed"].name == "lm_head" + + # -- Block submodules -- + + def test_blocks_ln1_is_rms_normalization_bridge( + self, adapter: Lfm2MoeArchitectureAdapter + ) -> None: + blocks = adapter.component_mapping["blocks"] + assert isinstance(blocks.submodules["ln1"], RMSNormalizationBridge) + + def test_blocks_ln1_name(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + blocks = adapter.component_mapping["blocks"] + assert blocks.submodules["ln1"].name == "operator_norm" - def test_component_types(self, adapter: Lfm2MoeArchitectureAdapter) -> None: - mapping = adapter.component_mapping - assert isinstance(mapping["embed"], EmbeddingBridge) - assert isinstance(mapping["blocks"], Lfm2MoeBlockBridge) - assert isinstance(mapping["ln_final"], RMSNormalizationBridge) - assert isinstance(mapping["unembed"], UnembeddingBridge) + def test_blocks_ln2_is_rms_normalization_bridge( + self, adapter: Lfm2MoeArchitectureAdapter + ) -> None: + blocks = adapter.component_mapping["blocks"] + assert isinstance(blocks.submodules["ln2"], RMSNormalizationBridge) + + def test_blocks_ln2_name(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + blocks = adapter.component_mapping["blocks"] + assert blocks.submodules["ln2"].name == "ffn_norm" + + def test_attn_is_position_embeddings_attention_bridge( + self, adapter: Lfm2MoeArchitectureAdapter + ) -> None: + blocks = adapter.component_mapping["blocks"] + assert isinstance(blocks.submodules["attn"], PositionEmbeddingsAttentionBridge) - def test_hf_module_paths(self, adapter: Lfm2MoeArchitectureAdapter) -> None: - mapping = adapter.component_mapping - assert mapping["embed"].name == "model.embed_tokens" - assert mapping["blocks"].name == "model.layers" - assert mapping["ln_final"].name == "model.embedding_norm" - assert mapping["unembed"].name == "lm_head" + def test_attn_name(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + blocks = adapter.component_mapping["blocks"] + assert blocks.submodules["attn"].name == "self_attn" + + def test_attn_requires_attention_mask_is_true( + self, adapter: Lfm2MoeArchitectureAdapter + ) -> None: + blocks = adapter.component_mapping["blocks"] + assert blocks.submodules["attn"].requires_attention_mask is True - def test_blocks_only_advertise_supported_residual_aliases( + def test_attn_requires_position_embeddings_is_true( self, adapter: Lfm2MoeArchitectureAdapter ) -> None: blocks = adapter.component_mapping["blocks"] - assert blocks.hook_aliases == { - "hook_resid_pre": "hook_in", - "hook_resid_post": "hook_out", + assert blocks.submodules["attn"].requires_position_embeddings is True + + def test_conv_is_lfm2shortconv_bridge(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + blocks = adapter.component_mapping["blocks"] + assert isinstance(blocks.submodules["conv"], Lfm2ShortConvBridge) + + def test_conv_name(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + blocks = adapter.component_mapping["blocks"] + assert blocks.submodules["conv"].name == "conv" + + def test_mlp_is_moe_bridge(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + blocks = adapter.component_mapping["blocks"] + assert isinstance(blocks.submodules["mlp"], MoEBridge) + + def test_mlp_name(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + blocks = adapter.component_mapping["blocks"] + assert blocks.submodules["mlp"].name == "feed_forward" + + # -- Attention submodules -- + + @pytest.mark.parametrize("slot", ["q", "k", "v", "o"]) + def test_attn_submodule_is_linear_bridge( + self, adapter: Lfm2MoeArchitectureAdapter, slot: str + ) -> None: + attn = adapter.component_mapping["blocks"].submodules["attn"] + assert isinstance(attn.submodules[slot], LinearBridge) + + @pytest.mark.parametrize("slot", ["q_norm", "k_norm"]) + def test_attn_submodule_is_rms_normalization_bridge( + self, adapter: Lfm2MoeArchitectureAdapter, slot: str + ) -> None: + attn = adapter.component_mapping["blocks"].submodules["attn"] + assert isinstance(attn.submodules[slot], RMSNormalizationBridge) + + @pytest.mark.parametrize( + "slot, hf_name", + [ + ("q", "q_proj"), + ("k", "k_proj"), + ("v", "v_proj"), + ("o", "out_proj"), + ("q_norm", "q_layernorm"), + ("k_norm", "k_layernorm"), + ], + ) + def test_attn_submodule_name( + self, adapter: Lfm2MoeArchitectureAdapter, slot: str, hf_name: str + ) -> None: + attn = adapter.component_mapping["blocks"].submodules["attn"] + assert attn.submodules[slot].name == hf_name + + # -- Conv submodules -- + + def test_conv_in_is_linear_bridge(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + conv = adapter.component_mapping["blocks"].submodules["conv"] + assert isinstance(conv.submodules["in"], LinearBridge) + + def test_conv_in_name(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + conv = adapter.component_mapping["blocks"].submodules["conv"] + assert conv.submodules["in"].name == "in_proj" + + def test_conv_conv_is_depthwise_conv1d_bridge( + self, adapter: Lfm2MoeArchitectureAdapter + ) -> None: + conv = adapter.component_mapping["blocks"].submodules["conv"] + assert isinstance(conv.submodules["conv"], DepthwiseConv1DBridge) + + def test_conv_conv_name(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + conv = adapter.component_mapping["blocks"].submodules["conv"] + assert conv.submodules["conv"].name == "conv" + + def test_conv_out_is_linear_bridge(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + conv = adapter.component_mapping["blocks"].submodules["conv"] + assert isinstance(conv.submodules["out"], LinearBridge) + + def test_conv_out_name(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + conv = adapter.component_mapping["blocks"].submodules["conv"] + assert conv.submodules["out"].name == "out_proj" + + # -- MLP submodules -- + + def test_mlp_gate_is_lfm2_moe_gate_bridge(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + mlp = adapter.component_mapping["blocks"].submodules["mlp"] + assert isinstance(mlp.submodules["gate"], Lfm2MoeGateBridge) + + def test_mlp_gate_name(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + mlp = adapter.component_mapping["blocks"].submodules["mlp"] + assert mlp.submodules["gate"].name == "gate" + + def test_mlp_dense_gate_is_linear_bridge(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + mlp = adapter.component_mapping["blocks"].submodules["mlp"] + assert isinstance(mlp.submodules["dense_gate"], LinearBridge) + + def test_mlp_dense_gate_name(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + mlp = adapter.component_mapping["blocks"].submodules["mlp"] + assert mlp.submodules["dense_gate"].name == "w1" + + def test_mlp_dense_in_is_linear_bridge(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + mlp = adapter.component_mapping["blocks"].submodules["mlp"] + assert isinstance(mlp.submodules["dense_in"], LinearBridge) + + def test_mlp_dense_in_name(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + mlp = adapter.component_mapping["blocks"].submodules["mlp"] + assert mlp.submodules["dense_in"].name == "w3" + + def test_mlp_dense_out_is_linear_bridge(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + mlp = adapter.component_mapping["blocks"].submodules["mlp"] + assert isinstance(mlp.submodules["dense_out"], LinearBridge) + + def test_mlp_dense_out_name(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + mlp = adapter.component_mapping["blocks"].submodules["mlp"] + assert mlp.submodules["dense_out"].name == "w2" + + +# --------------------------------------------------------------------------- +# Weight processing conversion tests +# --------------------------------------------------------------------------- + + +class TestLfm2AdapterWeightConversions: + """Adapter must define exactly the four QKVO weight conversions.""" + + def test_conversion_keys_present(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + """Lfm2 has 4 weight matrices (QKVO) per attention layer""" + assert 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", } - assert blocks.submodules == {} + + @pytest.mark.parametrize("slot", ["q", "k", "v"]) + def test_qkv_weight_uses_split_heads_pattern( + self, adapter: Lfm2MoeArchitectureAdapter, slot: str + ) -> None: + conv = adapter.weight_processing_conversions[f"blocks.{{i}}.attn.{slot}.weight"] + expected = adapter.cfg.n_key_value_heads if slot in ["k", "v"] else adapter.cfg.n_heads + assert isinstance(conv, ParamProcessingConversion) + assert isinstance(conv.tensor_conversion, RearrangeTensorConversion) + assert conv.tensor_conversion.pattern == "(n h) m -> n m h" + assert conv.tensor_conversion.axes_lengths["n"] == expected + + def test_o_uses_merge_heads_pattern(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + conv = adapter.weight_processing_conversions["blocks.{i}.attn.o.weight"] + assert isinstance(conv, ParamProcessingConversion) + assert isinstance(conv.tensor_conversion, RearrangeTensorConversion) + assert conv.tensor_conversion.pattern == "m (n h) -> n h m" + assert conv.tensor_conversion.axes_lengths["n"] == adapter.cfg.n_heads + + +# --------------------------------------------------------------------------- +# Architecture guards +# --------------------------------------------------------------------------- + + +class TestLfm2ArchitectureGuards: + """Guard against accidental introduction of features Lfm2 does not have.""" + + def test_no_pos_embed_component(self, adapter: Lfm2MoeArchitectureAdapter) -> None: + """Lfm2 uses rotary embeddings, so there is no learned positional embedding.""" + assert "pos_embed" not in adapter.component_mapping + + +# --------------------------------------------------------------------------- +# Setup component testing tests +# --------------------------------------------------------------------------- + + +class TestLfm2SetupComponentTesting: + """setup_component_testing must wire Lfm2's shared rotary embedding into attention bridges.""" + + def test_setup_flips_top_level_attn_implementation_to_eager( + self, adapter: Lfm2MoeArchitectureAdapter + ) -> None: + """HF reference defaults to sdpa; setup must flip the top-level config to eager.""" + hf = _fake_hf_model(object()) + assert hf.config._attn_implementation == "sdpa" + + adapter.setup_component_testing(hf) + + assert hf.config._attn_implementation == "eager" + + def test_setup_flips_per_layer_attn_implementation_to_eager( + self, adapter: Lfm2MoeArchitectureAdapter + ) -> None: + """Each already-built attn layer caches its own config; setup must flip all of them.""" + hf = _fake_hf_model(object(), n_layers=2) + assert all(l.self_attn.config._attn_implementation == "sdpa" for l in hf.model.layers) + + adapter.setup_component_testing(hf) + + for layer in hf.model.layers: + assert layer.self_attn.config._attn_implementation == "eager" + + def test_sets_rotary_emb_on_template_attention( + self, adapter: Lfm2MoeArchitectureAdapter + ) -> None: + rotary_emb = object() + attn_template = adapter.get_generalized_component("blocks.0.attn") + assert isinstance(attn_template, PositionEmbeddingsAttentionBridge) + assert attn_template._rotary_emb is None + + adapter.setup_component_testing(_fake_hf_model(rotary_emb)) + + assert attn_template._rotary_emb is rotary_emb + + def test_sets_rotary_emb_on_each_bridge_model_attention( + self, adapter: Lfm2MoeArchitectureAdapter + ) -> None: + pos_emb = object() + bridge_model = DummyBridgeModel([DummyBlock(), DummyBlock(), DummyBlock()]) + + adapter.setup_component_testing(_fake_hf_model(pos_emb), bridge_model=bridge_model) + + for block in bridge_model.blocks: + assert block.attn.pos_emb is pos_emb + + def test_skips_bridge_blocks_without_attention( + self, adapter: Lfm2MoeArchitectureAdapter + ) -> None: + pos_emb = object() + bridge_model = DummyBridgeModel([DummyBlock(), DummyBlock(has_attention=False)]) + + adapter.setup_component_testing(_fake_hf_model(pos_emb), bridge_model=bridge_model) + + assert bridge_model.blocks[0].attn.pos_emb is pos_emb diff --git a/transformer_lens/model_bridge/supported_architectures/lfm2_moe.py b/transformer_lens/model_bridge/supported_architectures/lfm2_moe.py index 496b1527a..2751f140f 100644 --- a/transformer_lens/model_bridge/supported_architectures/lfm2_moe.py +++ b/transformer_lens/model_bridge/supported_architectures/lfm2_moe.py @@ -1,76 +1,163 @@ """LiquidAI LFM2 MoE architecture adapter.""" -from typing import Any +from typing import Any, Dict, Optional + +import torch from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter from transformer_lens.model_bridge.generalized_components import ( BlockBridge, + DepthwiseConv1DBridge, EmbeddingBridge, + Lfm2ShortConvBridge, + LinearBridge, + MoEBridge, + MoERouterBridge, + PositionEmbeddingsAttentionBridge, RMSNormalizationBridge, + RotaryEmbeddingBridge, UnembeddingBridge, ) -class Lfm2MoeBlockBridge(BlockBridge): - """Whole-layer LFM2 bridge exposing only residual stream hooks. - - LFM2 MoE interleaves short-convolution and full-attention operator layers. - Wrapping the HF layer as a whole preserves correct execution while avoiding - unresolved standard attention/MLP aliases on layers that do not have them. - """ - - hook_aliases = { - "hook_resid_pre": "hook_in", - "hook_resid_post": "hook_out", - } +class Lfm2MoeGateBridge(MoERouterBridge): + def get_random_inputs( + self, + batch_size: int = 2, + seq_len: int = 8, + device: Optional[torch.device] = None, + dtype: Optional[torch.dtype] = None, + ) -> Dict[str, Any]: + """Random inputs for router component testing. + + The router runs on the reshaped [N, d_model] hidden states and takes a + second `expert_bias` arg (use_expert_bias=True); its top-k gather is + hardcoded to dim=1, so the input must be 2D or the gather indexes the + sequence axis out of bounds. + + Args: + batch_size: Batch size for generated inputs + seq_len: Sequence length for generated inputs + device: Device to place tensors on + dtype: Dtype for generated tensors (defaults to float32) + + Returns: + Dictionary of input tensors matching the component's expected input signature + """ + if device is None: + device = torch.device("cpu") + if dtype is None: + dtype = torch.float32 + d_model = self.config.d_model if self.config and hasattr(self.config, "d_model") else 768 + num_experts = ( + self.config.num_experts if self.config and hasattr(self.config, "num_experts") else 0 + ) + hidden_states = torch.randn(batch_size * seq_len, d_model, device=device, dtype=dtype) + expert_bias = torch.zeros(num_experts, device=device) + return {"args": (hidden_states, expert_bias)} class Lfm2MoeArchitectureAdapter(ArchitectureAdapter): - """Architecture adapter for LiquidAI LFM2 MoE models. - - LFM2 MoE is a hybrid decoder with both short-convolution and full-attention - layers. The adapter delegates each decoder layer to HF and exposes residual - hooks around the whole layer rather than pretending every layer has a - homogeneous attention/MLP substructure. - """ - - # Phases 1-3 compare standard attention/MLP components, which this hybrid - # adapter intentionally doesn't expose (whole-layer residual hooks only). - # Phase 4 (generation + text-quality) needs no component comparison, so it applies. - applicable_phases: list[int] = [4] + """Architecture adapter for LiquidAI Lfm2 MoE models.""" def __init__(self, cfg: Any) -> None: - """Initialize the LFM2 MoE architecture adapter.""" + """Initialize the Lfm2 MoE architecture adapter.""" super().__init__(cfg) self._set_rms_rotary_defaults() - # Hookable attention needs eager; the base prepare hooks force it through - # from_pretrained and onto the loaded config. - self.cfg.attn_implementation = "eager" - self.cfg.default_prepend_bos = False - if hasattr(cfg, "num_experts"): - self.cfg.num_experts = cfg.num_experts - if hasattr(cfg, "experts_per_token"): - self.cfg.experts_per_token = cfg.experts_per_token - if hasattr(cfg, "moe_intermediate_size"): - setattr(self.cfg, "moe_intermediate_size", cfg.moe_intermediate_size) - if hasattr(cfg, "layer_types"): - setattr(self.cfg, "layer_types", cfg.layer_types) - - norm_eps = getattr(cfg, "norm_eps", None) - if norm_eps is not None: - self.cfg.eps = norm_eps + self.cfg.act_fn = "silu" + self.cfg.attn_implementation = "eager" rope_parameters = getattr(cfg, "rope_parameters", None) or {} rope_theta = rope_parameters.get("rope_theta") or getattr(cfg, "rope_theta", None) if rope_theta is not None: self.cfg.rotary_base = rope_theta + self.weight_processing_conversions = { + **self._qkvo_weight_conversions(), + } + self.component_mapping = { "embed": EmbeddingBridge(name="model.embed_tokens"), - "blocks": Lfm2MoeBlockBridge(name="model.layers", config=self.cfg), - # LFM2 stores the decoder-final norm at embedding_norm, not model.norm. + "rotary_emb": RotaryEmbeddingBridge(name="model.pos_emb"), + "blocks": BlockBridge( + name="model.layers", + config=self.cfg, + submodules={ + "ln1": RMSNormalizationBridge( + name="operator_norm", + config=self.cfg, + ), + "ln2": RMSNormalizationBridge( + name="ffn_norm", + config=self.cfg, + ), + "attn": PositionEmbeddingsAttentionBridge( + name="self_attn", + config=self.cfg, + optional=True, + submodules={ + "q": LinearBridge(name="q_proj"), + "k": LinearBridge(name="k_proj"), + "v": LinearBridge(name="v_proj"), + "o": LinearBridge(name="out_proj"), + "q_norm": RMSNormalizationBridge(name="q_layernorm", config=self.cfg), + "k_norm": RMSNormalizationBridge(name="k_layernorm", config=self.cfg), + }, + requires_attention_mask=True, + requires_position_embeddings=True, + ), + "conv": Lfm2ShortConvBridge( + name="conv", + config=self.cfg, + optional=True, + submodules={ + "in": LinearBridge(name="in_proj"), + "conv": DepthwiseConv1DBridge(name="conv"), + "out": LinearBridge(name="out_proj"), + }, + ), + "mlp": MoEBridge( + name="feed_forward", + config=self.cfg, + sparse_required=("gate",), + submodules={ + "gate": Lfm2MoeGateBridge(name="gate", config=self.cfg, optional=True), + "dense_gate": LinearBridge(name="w1", optional=True), + "dense_in": LinearBridge(name="w3", optional=True), + "dense_out": LinearBridge(name="w2", optional=True), + }, + ), + }, + ), "ln_final": RMSNormalizationBridge(name="model.embedding_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 model-specific references for component testing.""" + rotary_emb = hf_model.model.pos_emb + + # Set attention implementation on HF model to eager (vs sdpa default) + if hasattr(hf_model, "config") and hasattr(hf_model.config, "_attn_implementation"): + hf_model.config._attn_implementation = "eager" + + if hasattr(hf_model, "model") and hasattr(hf_model.model, "layers"): + for layer in hf_model.model.layers: + if hasattr(layer, "self_attn") and hasattr(layer.self_attn, "config"): + layer.self_attn.config._attn_implementation = "eager" + + # Set rotary_emb on actual bridge instances + 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) + + # Set on template for get_generalized_component() calls + # Find the first attention layer (LFM2 layer 0 is conv, not attn) + layer_types = getattr(self.cfg, "layer_types", None) + if layer_types is not None and "full_attention" in layer_types: + first_attn_idx = layer_types.index("full_attention") + attn_bridge = self.get_generalized_component(f"blocks.{first_attn_idx}.attn") + attn_bridge.set_rotary_emb(rotary_emb) diff --git a/transformer_lens/tools/model_registry/data/supported_models.json b/transformer_lens/tools/model_registry/data/supported_models.json index fabaf12e6..f0fa44a24 100644 --- a/transformer_lens/tools/model_registry/data/supported_models.json +++ b/transformer_lens/tools/model_registry/data/supported_models.json @@ -9,7 +9,7 @@ "total_architectures": 143, "total_models": 15670, "total_provisional": 7, - "total_verified": 1205, + "total_verified": 1204, "models": [ { "architecture_id": "FalconH1ForCausalLM", @@ -30400,16 +30400,17 @@ { "architecture_id": "Lfm2MoeForCausalLM", "model_id": "LiquidAI/LFM2.5-8B-A1B", - "status": 1, - "verified_date": "2026-06-26", + "status": 3, + "verified_date": "2026-08-15", "metadata": null, - "note": "Full verification completed with issues, low text quality", - "phase1_score": null, - "phase2_score": null, - "phase3_score": null, - "phase4_score": 23.6, + "note": "Below threshold: P3=85.0% but required tests failed: logits_equivalence, loss_equivalence — Text quality score: 24.7/100 (avg perplexity: 19.6) — generated text may be incoherent", + "phase1_score": 100.0, + "phase2_score": 100.0, + "phase3_score": 85.0, + "phase4_score": 24.7, "phase7_score": null, - "phase8_score": null + "phase8_score": null, + "phase9_score": null }, { "architecture_id": "LlamaForCausalLM", diff --git a/transformer_lens/tools/model_registry/data/verification_history.json b/transformer_lens/tools/model_registry/data/verification_history.json index 24dd53444..4c16b1be3 100644 --- a/transformer_lens/tools/model_registry/data/verification_history.json +++ b/transformer_lens/tools/model_registry/data/verification_history.json @@ -23040,6 +23040,26 @@ "notes": "Full verification completed", "invalidated": false, "invalidation_reason": null + }, + { + "model_id": "LiquidAI/LFM2.5-8B-A1B", + "architecture_id": "Lfm2MoeForCausalLM", + "verified_date": "2026-08-14", + "verified_by": "verify_models", + "transformerlens_version": null, + "notes": "Below threshold: P1=50.0% < 100.0% (failed: all_components); P3=38.9% < 75.0% (failed: layer_norm_fo — 22/152 components failed (22 critical)", + "invalidated": false, + "invalidation_reason": null + }, + { + "model_id": "LiquidAI/LFM2.5-8B-A1B", + "architecture_id": "Lfm2MoeForCausalLM", + "verified_date": "2026-08-15", + "verified_by": "verify_models", + "transformerlens_version": null, + "notes": "Below threshold: P3=85.0% but required tests failed: logits_equivalence, loss_equivalence — Text quality score: 24.7/100 (avg perplexity: 19.6) — generated text may be incoherent", + "invalidated": false, + "invalidation_reason": null } ] }