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
1 change: 1 addition & 0 deletions tests/QUARANTINES.md
Original file line number Diff line number Diff line change
Expand Up @@ -158,6 +158,7 @@ now bound solely by the surviving
| [`unit/model_bridge/test_hook_alias_resolution.py`:90](unit/model_bridge/test_hook_alias_resolution.py) | `xfail(strict=True)` per-arch | Hook-alias gaps |
| [`unit/model_bridge/supported_architectures/test_qwen3_5_adapter.py`:448,464,494,514,605,700,805,947,1133](unit/model_bridge/supported_architectures/test_qwen3_5_adapter.py) | `skipif` ×9 | Qwen3_5 classes absent from installed transformers |
| [`unit/model_bridge/supported_architectures/test_qwen3_next_adapter.py`:397](unit/model_bridge/supported_architectures/test_qwen3_next_adapter.py) | `skipif` | Qwen3NextForCausalLM absent from installed transformers |
| [`unit/model_bridge/supported_architectures/test_muse_glimmer_adapter.py`:8](unit/model_bridge/supported_architectures/test_muse_glimmer_adapter.py) (whole file) | `importorskip` | MuseGlimmer classes absent from installed transformers (needs >= 5.15; lock is 5.13) |
| [`integration/test_weight_processing_integration.py`:279](integration/test_weight_processing_integration.py) | `skip` | Weight-processing edge case |
| [`acceptance/model_bridge/compatibility/test_backward_hooks.py`:11](acceptance/model_bridge/compatibility/test_backward_hooks.py) | `skip` | Backward-hook compatibility |

Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,96 @@
"""Tests for MuseGlimmerArchitectureAdapter on a tiny random config."""

import copy

import pytest
import torch

pytest.importorskip("transformers.models.muse_glimmer")

from transformers import AutoModelForImageTextToText
from transformers.models.muse_glimmer import MuseGlimmerConfig

from transformer_lens.model_bridge.bridge import TransformerBridge
from transformer_lens.model_bridge.sources._bridge_builder import (
build_bridge_config_from_hf,
)
from transformer_lens.model_bridge.supported_architectures.muse_glimmer import (
MuseGlimmerArchitectureAdapter,
)

ARCH = "MuseGlimmerForConditionalGeneration"
N_LAYERS = 4
TOKENS = torch.tensor([[5, 17, 29, 3, 11, 42, 7, 23], [8, 9, 10, 11, 12, 13, 14, 15]])


class _Tok:
pass


@pytest.fixture(scope="module")
def models():
torch.manual_seed(0)
cfg = MuseGlimmerConfig(
text_config=dict(
vocab_size=64,
hidden_size=64,
intermediate_size=96,
num_hidden_layers=N_LAYERS,
num_attention_heads=4,
num_key_value_heads=2,
head_dim=16,
max_position_embeddings=64,
sliding_window=4,
),
vision_config=dict(
hidden_size=32,
intermediate_size=48,
num_hidden_layers=1,
num_attention_heads=2,
pos_emb_height=4,
pos_emb_width=4,
),
out_hidden_size=128,
projector_hidden_size=32,
)
cfg.architectures = [ARCH]
hf = AutoModelForImageTextToText.from_config(cfg, attn_implementation="eager")
hf = hf.to(torch.float32).eval()
reference = copy.deepcopy(hf)
bridge_cfg = build_bridge_config_from_hf(hf.config, ARCH, "muse-glimmer-tiny", torch.float32)
bridge = TransformerBridge(hf, MuseGlimmerArchitectureAdapter(bridge_cfg), tokenizer=_Tok())
return bridge, reference


def test_forward_matches_hf(models):
bridge, reference = models
with torch.no_grad():
bridge_logits = bridge(TOKENS)
hf_logits = reference(input_ids=TOKENS).logits
assert torch.allclose(bridge_logits, hf_logits, atol=1e-4, rtol=0)


def test_run_with_cache_hooks(models):
bridge, reference = models
with torch.no_grad():
_, cache = bridge.run_with_cache(TOKENS)
hf_attn = reference(input_ids=TOKENS, output_attentions=True).attentions
batch, seq = TOKENS.shape
for i in range(N_LAYERS):
for name in ("hook_resid_pre", "hook_attn_out", "hook_mlp_out", "hook_resid_post"):
assert cache[f"blocks.{i}.{name}"].shape == (batch, seq, 64)
assert cache[f"blocks.{i}.attn.hook_z"].shape == (batch, seq, 4, 16)
torch.testing.assert_close(cache[f"blocks.{i}.attn.hook_pattern"], hf_attn[i])


def test_vision_projector_hook_out_is_the_merged_embedding(models):
"""HF scatters ``perception_emb_norm``'s output into the text embeddings.

The projector bridge therefore wraps that module rather than the raw
``vision_projection`` Linear, so ``vision_projector.hook_out`` is the value
HF actually merges (matching the other multimodal adapters).
"""
bridge, _ = models
projector = bridge.vision_projector
assert projector.name == "model.perception_emb_norm"
assert type(projector.original_component).__name__ == "MuseGlimmerRMSNorm"
42 changes: 42 additions & 0 deletions tests/unit/model_bridge/test_output_logits_contract.py
Original file line number Diff line number Diff line change
Expand Up @@ -12,6 +12,9 @@
from transformer_lens.model_bridge.supported_architectures.granite import (
GraniteArchitectureAdapter,
)
from transformer_lens.model_bridge.supported_architectures.muse_glimmer import (
MuseGlimmerArchitectureAdapter,
)


def _text_config(**overrides: object) -> SimpleNamespace:
Expand Down Expand Up @@ -93,3 +96,42 @@ def test_granite_and_falcon_apply_declared_output_scalars() -> None:
FalconH1ArchitectureAdapter(falcon_cfg).apply_output_logits_transform(logits),
torch.tensor([[-2.0, 1.0]]),
)


def test_nested_text_config_muse_multiplier_is_used() -> None:
wrapper = SimpleNamespace(
text_config=_text_config(final_logit_softcapping=20.0, output_multiplier=0.5)
)
cfg = build_bridge_config_from_hf(
wrapper,
"MuseGlimmerForConditionalGeneration",
"tiny-muse-output-contract",
torch.float32,
)

assert cfg.output_logits_soft_cap == 20.0
assert cfg.output_multiplier == 0.5


def test_muse_glimmer_applies_multiplier_before_the_softcap() -> None:
"""Muse Glimmer scales logits by ``output_multiplier`` and then tanh-softcaps them.

The tiny model's logits are too small for the cap to change anything, so the
forward-parity tests cannot catch a missing factor; this pins the contract with
logits large enough for both the multiplier and the cap to matter.
"""
wrapper = SimpleNamespace(
text_config=_text_config(final_logit_softcapping=20.0, output_multiplier=0.5)
)
cfg = build_bridge_config_from_hf(
wrapper,
"MuseGlimmerForConditionalGeneration",
"tiny-muse-output-contract",
torch.float32,
)
logits = torch.tensor([[-40.0, -6.0, 6.0, 40.0]])

torch.testing.assert_close(
MuseGlimmerArchitectureAdapter(cfg).apply_output_logits_transform(logits),
20.0 * torch.tanh(logits * 0.5 / 20.0),
)
2 changes: 2 additions & 0 deletions transformer_lens/factories/architecture_adapter_factory.py
Original file line number Diff line number Diff line change
Expand Up @@ -98,6 +98,7 @@
MixtralArchitectureAdapter,
ModernBertDecoderArchitectureAdapter,
MPTArchitectureAdapter,
MuseGlimmerArchitectureAdapter,
MusicFlamingoArchitectureAdapter,
NanoChatArchitectureAdapter,
NanogptArchitectureAdapter,
Expand Down Expand Up @@ -267,6 +268,7 @@
"Mistral3ForConditionalGeneration": Mistral3ArchitectureAdapter,
"MPTForCausalLM": MPTArchitectureAdapter,
"MptForCausalLM": MPTArchitectureAdapter,
"MuseGlimmerForConditionalGeneration": MuseGlimmerArchitectureAdapter,
"NeoForCausalLM": NeoArchitectureAdapter,
"NeoXForCausalLM": NeoxArchitectureAdapter,
"NeelSoluOldForCausalLM": NeelSoluOldArchitectureAdapter,
Expand Down
2 changes: 2 additions & 0 deletions transformer_lens/model_bridge/sources/_bridge_builder.py
Original file line number Diff line number Diff line change
Expand Up @@ -154,6 +154,8 @@
"n_layers_in_coda",
"injection_type",
"qk_bias",
# Muse Glimmer
"output_multiplier",
# RWKV-7 (attention-free recurrent, generalized delta-rule time-mixing).
# head_dim is intentionally omitted: it is a read-only alias of d_head on
# TransformerBridgeConfig, so a passthrough setattr would raise.
Expand Down
1 change: 1 addition & 0 deletions transformer_lens/model_bridge/sources/_hf_format.py
Original file line number Diff line number Diff line change
Expand Up @@ -350,6 +350,7 @@ def determine_architecture_from_hf_config(hf_config):
"mistral3": "Mistral3ForConditionalGeneration",
"mixtral": "MixtralForCausalLM",
"mpt": "MptForCausalLM",
"muse_glimmer": "MuseGlimmerForConditionalGeneration",
"gemma": "GemmaForCausalLM",
"gemma2": "Gemma2ForCausalLM",
"gemma3": "Gemma3ForCausalLM",
Expand Down
Original file line number Diff line number Diff line change
Expand Up @@ -166,6 +166,9 @@
ModernBertDecoderArchitectureAdapter,
)
from transformer_lens.model_bridge.supported_architectures.mpt import MPTArchitectureAdapter
from transformer_lens.model_bridge.supported_architectures.muse_glimmer import (
MuseGlimmerArchitectureAdapter,
)
from transformer_lens.model_bridge.supported_architectures.music_flamingo import (
MusicFlamingoArchitectureAdapter,
)
Expand Down Expand Up @@ -405,6 +408,7 @@
"MixtralArchitectureAdapter",
"ModernBertDecoderArchitectureAdapter",
"MPTArchitectureAdapter",
"MuseGlimmerArchitectureAdapter",
"MusicFlamingoArchitectureAdapter",
"NanogptArchitectureAdapter",
"NanoChatArchitectureAdapter",
Expand Down
Original file line number Diff line number Diff line change
@@ -0,0 +1,80 @@
"""Muse Glimmer (MuseGlimmerForConditionalGeneration) architecture adapter."""

from typing import Any

import torch

from transformer_lens.model_bridge.architecture_adapter import ArchitectureAdapter
from transformer_lens.model_bridge.generalized_components import (
AttentionBridge,
BlockBridge,
EmbeddingBridge,
LinearBridge,
RotaryEmbeddingBridge,
UnembeddingBridge,
VisionProjectionBridge,
)
from transformer_lens.model_bridge.generalized_components.base import (
GeneralizedComponent,
)


class MuseGlimmerArchitectureAdapter(ArchitectureAdapter):
"""Architecture adapter for MuseGlimmerForConditionalGeneration models."""

_testing_lm_attr = "model.language_model"
_testing_wire_rotary = False

# Sandwich norms rescale sublayer outputs, so folding them is not function-preserving.
supports_fold_ln = False

def __init__(self, cfg: Any) -> None:
super().__init__(cfg)

self.cfg.is_multimodal = True
self._extract_vision_dims(cfg)

self._set_rms_rotary_defaults()
self.cfg.attn_implementation = "eager"
self.weight_processing_conversions: dict = {}

self.component_mapping = {
"vision_encoder": GeneralizedComponent(name="model.vision_tower"),
# HF runs ``vision_adapter -> vision_projection -> perception_emb_norm`` and
# scatters only the last value into the text embeddings, so the projector
# bridge wraps the norm: ``hook_out`` is the merged embedding, as in the
# other VLM adapters (Gemma 3's norm lives inside its projector module).
# ``hook_in`` is then the projection output rather than the vision-tower one.
"vision_projector": VisionProjectionBridge(name="model.perception_emb_norm"),
"embed": EmbeddingBridge(name="model.language_model.embed_tokens"),
"rotary_emb": RotaryEmbeddingBridge(name="model.language_model.rotary_emb"),
"blocks": BlockBridge(
name="model.language_model.layers",
submodules={
"ln1": GeneralizedComponent(name="input_layernorm"),
"ln1_post": GeneralizedComponent(name="post_attention_layernorm"),
"ln2": GeneralizedComponent(name="pre_feedforward_layernorm"),
"ln2_post": GeneralizedComponent(name="post_feedforward_layernorm"),
"attn": AttentionBridge(
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"),
"gate": LinearBridge(name="gate_proj"),
},
maintain_native_attention=True,
requires_attention_mask=True,
),
"mlp": self._gated_mlp(),
},
),
"ln_final": GeneralizedComponent(name="model.language_model.norm"),
"unembed": UnembeddingBridge(name="lm_head", config=self.cfg),
}

def apply_output_logits_transform(self, logits: torch.Tensor) -> torch.Tensor:

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Bridge forward returns HF's own logits, so the parity test still passes with the multiplier or the softcap deleted from this method. Only JacobianLens calls it. The tiny model's logits are also too small for the cap to change anything at 1e-4. The Granite and Falcon-H1 cases in test_output_logits_contract.py cover this kind of override, and they run under the locked transformers too. Could you add a Muse Glimmer case there?

multiplier = float(getattr(self.cfg, "output_multiplier", 1.0))
return super().apply_output_logits_transform(logits * multiplier)
2 changes: 2 additions & 0 deletions transformer_lens/tools/model_registry/__init__.py
Original file line number Diff line number Diff line change
Expand Up @@ -150,6 +150,7 @@
"NemotronHForCausalLM",
"MPTForCausalLM",
"MptForCausalLM",
"MuseGlimmerForConditionalGeneration",
"MistralForCausalLM",
"Mistral3ForConditionalGeneration",
"MixtralForCausalLM",
Expand Down Expand Up @@ -315,6 +316,7 @@
"MixtralForCausalLM": ["mistralai"],
"MPTForCausalLM": ["mosaicml"],
"MptForCausalLM": ["mosaicml"],
"MuseGlimmerForConditionalGeneration": ["meta-models"],
"MT5ForConditionalGeneration": ["google", "bigscience", "csebuetnlp"],
"Olmo2ForCausalLM": ["allenai", "HPLT"],
"Olmo3ForCausalLM": ["allenai"],
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 @@ -45,6 +45,7 @@
"MistralForCausalLM": "Mistral AI's efficient 7B parameter model with sliding window attention",
"Mistral3ForConditionalGeneration": "Mistral AI's Mistral-Small VLM (Pixtral tower + Mistral decoder)",
"MixtralForCausalLM": "Mistral AI's Mixture of Experts model",
"MuseGlimmerForConditionalGeneration": "Meta's Muse Glimmer VLM (gated QK-norm attention, interleaved NoPE layers)",
"GemmaForCausalLM": "Google's Gemma lightweight open model family",
"Gemma2ForCausalLM": "Google's Gemma 2 with improved architecture",
"Gemma3ForCausalLM": "Google's Gemma 3 latest generation",
Expand Down
1 change: 1 addition & 0 deletions transformer_lens/utilities/architectures.py
Original file line number Diff line number Diff line change
Expand Up @@ -57,6 +57,7 @@
"Florence2ForConditionalGeneration",
"Mistral3ForConditionalGeneration",
"Llama4ForConditionalGeneration",
"MuseGlimmerForConditionalGeneration",
"Qwen2_5_VLForConditionalGeneration",
"Qwen3VLForConditionalGeneration",
"Qwen3VLMoeForConditionalGeneration",
Expand Down
Loading