diff --git a/tests/unit/model_bridge/supported_architectures/test_gemma4_adapter.py b/tests/unit/model_bridge/supported_architectures/test_gemma4_adapter.py index afff05df4..722910764 100644 --- a/tests/unit/model_bridge/supported_architectures/test_gemma4_adapter.py +++ b/tests/unit/model_bridge/supported_architectures/test_gemma4_adapter.py @@ -6,6 +6,7 @@ from transformer_lens.model_bridge.generalized_components import ( DelegatedAttentionBlockBridge, EmbeddingBridge, + GatedMLPBridge, LinearBridge, RotaryEmbeddingBridge, UnembeddingBridge, @@ -161,6 +162,12 @@ def test_moe_submodules_are_optional(): def test_gated_mlp_decomposition(): mlp = _adapter().component_mapping["blocks"].submodules["mlp"] + assert isinstance(mlp, GatedMLPBridge) assert mlp.submodules["gate"].name == "gate_proj" assert mlp.submodules["in"].name == "up_proj" assert mlp.submodules["out"].name == "down_proj" + assert mlp.hook_aliases == { + "hook_pre": "gate.hook_out", + "hook_pre_linear": "in.hook_out", + "hook_post": "out.hook_in", + } diff --git a/transformer_lens/model_bridge/supported_architectures/gemma4.py b/transformer_lens/model_bridge/supported_architectures/gemma4.py index bceb4ec73..0ddd39079 100644 --- a/transformer_lens/model_bridge/supported_architectures/gemma4.py +++ b/transformer_lens/model_bridge/supported_architectures/gemma4.py @@ -145,14 +145,7 @@ def __init__(self, cfg: Any) -> None: "v_norm": GeneralizedComponent(name="v_norm", optional=True), }, ), - "mlp": GeneralizedComponent( - name="mlp", - submodules={ - "gate": LinearBridge(name="gate_proj"), - "in": LinearBridge(name="up_proj"), - "out": LinearBridge(name="down_proj"), - }, - ), + "mlp": self._gated_mlp(), }, ), "ln_final": GeneralizedComponent(name="model.language_model.norm"),