Skip to content
Merged
Show file tree
Hide file tree
Changes from 1 commit
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,83 @@
"""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])
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,74 @@
"""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,
)
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"),
"vision_projector": GeneralizedComponent(name="model.vision_projection"),

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.

HF runs vision_adapter before this Linear and perception_emb_norm after it. So vision_projector.hook_out is an intermediate value, rescaled before it reaches the text stream. In the other VLM adapters this hook is the embedding that actually gets merged in. gemma3_multimodal.py is a good reference: its projector's norm runs inside the wrapped module, so hook_out is exactly what HF scatters into the text embeddings. Can the mapping expose that point here too?

"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