Repository navigation
Add Muse Glimmer architecture adapter #1822
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from 1 commit
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| 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]) |
| 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"), | ||
| "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: | ||
|
Collaborator
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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 |
||
| multiplier = float(getattr(self.cfg, "output_multiplier", 1.0)) | ||
| return super().apply_output_logits_transform(logits * multiplier) | ||
There was a problem hiding this comment.
Choose a reason for hiding this comment
The reason will be displayed to describe this comment to others. Learn more.
HF runs
vision_adapterbefore this Linear andperception_emb_normafter it. Sovision_projector.hook_outis 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.pyis 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?