From 58cf2f6688502f565526fbe5277dac98e093c204 Mon Sep 17 00:00:00 2001 From: buffett0323 Date: Sun, 16 Aug 2026 14:42:27 -0700 Subject: [PATCH] Add Liger Kernel support for Muse Glimmer --- README.md | 1 + src/liger_kernel/transformers/__init__.py | 9 + .../transformers/model/muse_glimmer.py | 148 +++++++++ .../transformers/model/output_classes.py | 15 + src/liger_kernel/transformers/monkey_patch.py | 143 +++++++++ src/liger_kernel/transformers/rms_norm.py | 49 +++ src/liger_kernel/transformers/swiglu.py | 20 ++ test/convergence/bf16/test_mini_models.py | 88 +++++ .../bf16/test_mini_models_multimodal.py | 120 +++++++ .../bf16/test_mini_models_with_logits.py | 89 ++++++ test/convergence/fp32/test_mini_models.py | 86 +++++ .../fp32/test_mini_models_multimodal.py | 119 +++++++ .../fp32/test_mini_models_with_logits.py | 86 +++++ .../Muse-Glimmer-30B/tokenizer_config.json | 132 ++++++++ test/transformers/test_monkey_patch.py | 132 ++++++++ test/transformers/test_muse_glimmer.py | 300 ++++++++++++++++++ test/utils.py | 12 + 17 files changed, 1549 insertions(+) create mode 100644 src/liger_kernel/transformers/model/muse_glimmer.py create mode 100644 test/resources/fake_configs/meta-models/Muse-Glimmer-30B/tokenizer_config.json create mode 100644 test/transformers/test_muse_glimmer.py diff --git a/README.md b/README.md index 491644892..d2171637d 100644 --- a/README.md +++ b/README.md @@ -377,6 +377,7 @@ loss.backward() | Ministral | `liger_kernel.transformers.apply_liger_kernel_to_ministral` | RoPE, RMSNorm, SwiGLU, CrossEntropyLoss, FusedLinearCrossEntropy | | Mistral | `liger_kernel.transformers.apply_liger_kernel_to_mistral` | RoPE, RMSNorm, SwiGLU, CrossEntropyLoss, FusedLinearCrossEntropy | | Mixtral | `liger_kernel.transformers.apply_liger_kernel_to_mixtral` | RoPE, RMSNorm, SwiGLU, CrossEntropyLoss, FusedLinearCrossEntropy | +| Muse Glimmer | `liger_kernel.transformers.apply_liger_kernel_to_muse_glimmer` | LayerNorm, RoPE, RMSNorm, SwiGLU, CrossEntropyLoss, FusedLinearCrossEntropy | | Nemotron | `liger_kernel.transformers.apply_liger_kernel_to_nemotron` | ReLUSquared, CrossEntropyLoss, FusedLinearCrossEntropy | | Pixtral | `liger_kernel.transformers.apply_liger_kernel_to_pixtral` | RoPE, RMSNorm, SwiGLU| | Gemma1 | `liger_kernel.transformers.apply_liger_kernel_to_gemma` | RoPE, RMSNorm, GeGLU, CrossEntropyLoss, FusedLinearCrossEntropy | diff --git a/src/liger_kernel/transformers/__init__.py b/src/liger_kernel/transformers/__init__.py index 26bdef91b..9f44ff7a8 100644 --- a/src/liger_kernel/transformers/__init__.py +++ b/src/liger_kernel/transformers/__init__.py @@ -20,6 +20,8 @@ from liger_kernel.transformers.poly_norm import LigerPolyNorm # noqa: F401 from liger_kernel.transformers.relu_squared import LigerReLUSquared # noqa: F401 from liger_kernel.transformers.rms_norm import LigerRMSNorm # noqa: F401 +from liger_kernel.transformers.rms_norm import LigerRMSNormForMuseGlimmer # noqa: F401 +from liger_kernel.transformers.rms_norm import LigerRMSNormForMuseGlimmerTextCentered # noqa: F401 from liger_kernel.transformers.rope import liger_rotary_pos_emb # noqa: F401 from liger_kernel.transformers.softmax import LigerSoftmax # noqa: F401 from liger_kernel.transformers.sparsemax import LigerSparsemax # noqa: F401 @@ -29,6 +31,7 @@ from liger_kernel.transformers.swiglu import LigerPhi3SwiGLUMLP # noqa: F401 from liger_kernel.transformers.swiglu import LigerQwen3MoeSwiGLUMLP # noqa: F401 from liger_kernel.transformers.swiglu import LigerSwiGLUMLP # noqa: F401 +from liger_kernel.transformers.swiglu import LigerSwiGLUMLPForMuseGlimmer # noqa: F401 from liger_kernel.transformers.tiled_mlp import LigerTiledGEGLUMLP # noqa: F401 from liger_kernel.transformers.tiled_mlp import LigerTiledSwiGLUMLP # noqa: F401 from liger_kernel.transformers.tvd import LigerTVDLoss # noqa: F401 @@ -63,6 +66,7 @@ from liger_kernel.transformers.monkey_patch import apply_liger_kernel_to_mistral # noqa: F401 from liger_kernel.transformers.monkey_patch import apply_liger_kernel_to_mixtral # noqa: F401 from liger_kernel.transformers.monkey_patch import apply_liger_kernel_to_mllama # noqa: F401 + from liger_kernel.transformers.monkey_patch import apply_liger_kernel_to_muse_glimmer # noqa: F401 from liger_kernel.transformers.monkey_patch import apply_liger_kernel_to_nemotron # noqa: F401 from liger_kernel.transformers.monkey_patch import apply_liger_kernel_to_olmo2 # noqa: F401 from liger_kernel.transformers.monkey_patch import apply_liger_kernel_to_olmo3 # noqa: F401 @@ -137,6 +141,7 @@ def __getattr__(name: str): "apply_liger_kernel_to_ministral", "apply_liger_kernel_to_mistral", "apply_liger_kernel_to_mixtral", + "apply_liger_kernel_to_muse_glimmer", "apply_liger_kernel_to_nemotron", "apply_liger_kernel_to_mllama", "apply_liger_kernel_to_olmo2", @@ -184,6 +189,8 @@ def __getattr__(name: str): "LigerPolyNorm", "LigerReLUSquared", "LigerRMSNorm", + "LigerRMSNormForMuseGlimmer", + "LigerRMSNormForMuseGlimmerTextCentered", "liger_rotary_pos_emb", "liger_llama4_text_rotary_pos_emb", "liger_llama4_vision_rotary_pos_emb", @@ -192,6 +199,7 @@ def __getattr__(name: str): "LigerPhi3SwiGLUMLP", "LigerQwen3MoeSwiGLUMLP", "LigerSwiGLUMLP", + "LigerSwiGLUMLPForMuseGlimmer", "LigerTiledGEGLUMLP", "LigerTiledSwiGLUMLP", "LigerTVDLoss", @@ -228,6 +236,7 @@ def __getattr__(name: str): "apply_liger_kernel_to_ministral", "apply_liger_kernel_to_mistral", "apply_liger_kernel_to_mixtral", + "apply_liger_kernel_to_muse_glimmer", "apply_liger_kernel_to_nemotron", "apply_liger_kernel_to_mllama", "apply_liger_kernel_to_olmo2", diff --git a/src/liger_kernel/transformers/model/muse_glimmer.py b/src/liger_kernel/transformers/model/muse_glimmer.py new file mode 100644 index 000000000..7779091d1 --- /dev/null +++ b/src/liger_kernel/transformers/model/muse_glimmer.py @@ -0,0 +1,148 @@ +from typing import Optional +from typing import Tuple +from typing import Union + +import torch + +from transformers.cache_utils import Cache +from transformers.utils import can_return_tuple + +from liger_kernel.transformers.model.loss_utils import LigerForCausalLMLoss +from liger_kernel.transformers.model.loss_utils import unpack_cross_entropy_result +from liger_kernel.transformers.model.output_classes import LigerMuseGlimmerCausalLMOutputWithPast + + +@can_return_tuple +def lce_forward( + self, + input_ids: Optional[torch.LongTensor] = None, + pixel_values: Optional[torch.FloatTensor] = None, + image_grid_thw: Optional[torch.LongTensor] = None, + pixel_values_videos: Optional[torch.FloatTensor] = None, + video_grid_thw: Optional[torch.LongTensor] = None, + attention_mask: Optional[torch.Tensor] = None, + position_ids: Optional[torch.LongTensor] = None, + past_key_values: Optional[Cache] = None, + inputs_embeds: Optional[torch.FloatTensor] = None, + labels: Optional[torch.LongTensor] = None, + use_cache: Optional[bool] = None, + logits_to_keep: Union[int, torch.Tensor] = 0, + skip_logits: Optional[bool] = None, + **kwargs, +) -> Union[Tuple, LigerMuseGlimmerCausalLMOutputWithPast]: + r""" + labels (`torch.LongTensor` of shape `(batch_size, sequence_length)`, *optional*): + Labels for computing the masked language modeling loss. Indices should either be in `[0, ..., + config.text_config.vocab_size]` or -100 (see `input_ids` docstring). Tokens with indices set to `-100` are + ignored (masked), the loss is only computed for the tokens with labels in + `[0, ..., config.text_config.vocab_size]`. + skip_logits (`bool`, *optional*): + Whether to skip materializing the logits and use Liger's fused linear cross entropy instead. Defaults to + `True` during training when labels are provided. + + Example: + + ```python + >>> from transformers import AutoProcessor, MuseGlimmerForConditionalGeneration + + >>> model = MuseGlimmerForConditionalGeneration.from_pretrained(MUSE_GLIMMER_CHECKPOINT) + >>> processor = AutoProcessor.from_pretrained(MUSE_GLIMMER_CHECKPOINT) + + >>> messages = [ + { + "role": "user", + "content": [ + {"type": "image", "image": "https://example.com/image.jpeg"}, + {"type": "text", "text": "Describe the image."}, + ], + } + ] + + >>> inputs = processor.apply_chat_template( + messages, tokenize=True, add_generation_prompt=True, return_dict=True, return_tensors="pt" + ) + + >>> generated_ids = model.generate(**inputs, max_new_tokens=1024) + >>> output_text = processor.batch_decode(generated_ids, skip_special_tokens=True)[0] + ``` + """ + outputs = self.model( + input_ids=input_ids, + pixel_values=pixel_values, + image_grid_thw=image_grid_thw, + pixel_values_videos=pixel_values_videos, + video_grid_thw=video_grid_thw, + attention_mask=attention_mask, + position_ids=position_ids, + past_key_values=past_key_values, + inputs_embeds=inputs_embeds, + use_cache=use_cache, + **kwargs, + ) + + hidden_states = outputs.last_hidden_state + slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep + kept_hidden_states = hidden_states[:, slice_indices, :] + + text_config = self.config.text_config + output_multiplier = text_config.output_multiplier + final_logit_softcapping = text_config.final_logit_softcapping + + shift_labels = kwargs.pop("shift_labels", None) + loss = None + logits = None + token_accuracy = None + predicted_tokens = None + + if skip_logits and labels is None and shift_labels is None: + raise ValueError("skip_logits is True, but labels and shift_labels are None") + + if skip_logits is None: + skip_logits = self.training and (labels is not None or shift_labels is not None) + + if skip_logits: + # MuseGlimmer computes `T * tanh(logits * m / T)` where `m = output_multiplier` and + # `T = final_logit_softcapping`, but Liger's fused linear cross entropy softcap only + # implements `T * tanh(logits / T)`. Fold `m` into the hidden states instead, since + # `(m * h) @ W.T == m * (h @ W.T)`. Scaling the hidden states rather than the lm_head + # weight keeps the extra tensor at `[batch, seq, hidden]` instead of `[vocab, hidden]`. + # Make kept_hidden_states contiguous for LigerForCausalLMLoss after slicing + kept_hidden_states = kept_hidden_states * output_multiplier + + result = LigerForCausalLMLoss( + hidden_states=kept_hidden_states, + lm_head_weight=self.lm_head.weight, + labels=labels, + shift_labels=shift_labels, + hidden_size=text_config.hidden_size, + final_logit_softcapping=final_logit_softcapping, + **kwargs, + ) + loss, _, token_accuracy, predicted_tokens = unpack_cross_entropy_result(result) + else: + logits = self.lm_head(kept_hidden_states) + + logits = logits * output_multiplier + logits = logits / final_logit_softcapping + logits = torch.tanh(logits) + logits = logits * final_logit_softcapping + + if labels is not None or shift_labels is not None: + loss = self.loss_function( + logits=logits, + labels=labels, + shift_labels=shift_labels, + vocab_size=text_config.vocab_size, + **kwargs, + ) + + return LigerMuseGlimmerCausalLMOutputWithPast( + loss=loss, + logits=logits, + past_key_values=outputs.past_key_values, + hidden_states=outputs.hidden_states, + attentions=outputs.attentions, + image_hidden_states=outputs.image_hidden_states, + token_accuracy=token_accuracy, + predicted_tokens=predicted_tokens, + ) diff --git a/src/liger_kernel/transformers/model/output_classes.py b/src/liger_kernel/transformers/model/output_classes.py index c2ef4a7c7..63f74013b 100644 --- a/src/liger_kernel/transformers/model/output_classes.py +++ b/src/liger_kernel/transformers/model/output_classes.py @@ -43,6 +43,13 @@ except Exception: _LlavaCausalLMOutputWithPast = None +try: + from transformers.models.muse_glimmer.modeling_muse_glimmer import ( + MuseGlimmerCausalLMOutputWithPast as _MuseGlimmerCausalLMOutputWithPast, + ) +except Exception: + _MuseGlimmerCausalLMOutputWithPast = None + try: from transformers.models.paligemma.modeling_paligemma import ( PaliGemmaCausalLMOutputWithPast as _PaliGemmaCausalLMOutputWithPast, @@ -145,6 +152,14 @@ class LigerInternVLCausalLMOutputWithPast(_InternVLCausalLMOutputWithPast): predicted_tokens: Optional[torch.LongTensor] = None +if _MuseGlimmerCausalLMOutputWithPast is not None: + + @dataclass + class LigerMuseGlimmerCausalLMOutputWithPast(_MuseGlimmerCausalLMOutputWithPast): + token_accuracy: Optional[torch.FloatTensor] = None + predicted_tokens: Optional[torch.LongTensor] = None + + if _PaliGemmaCausalLMOutputWithPast is not None: @dataclass diff --git a/src/liger_kernel/transformers/monkey_patch.py b/src/liger_kernel/transformers/monkey_patch.py index 4d33d7e41..faa1c0242 100755 --- a/src/liger_kernel/transformers/monkey_patch.py +++ b/src/liger_kernel/transformers/monkey_patch.py @@ -32,12 +32,15 @@ from liger_kernel.transformers.qwen2vl_mrope import liger_multimodal_rotary_pos_emb from liger_kernel.transformers.relu_squared import LigerReLUSquared from liger_kernel.transformers.rms_norm import LigerRMSNorm +from liger_kernel.transformers.rms_norm import LigerRMSNormForMuseGlimmer +from liger_kernel.transformers.rms_norm import LigerRMSNormForMuseGlimmerTextCentered from liger_kernel.transformers.rope import liger_rotary_pos_emb from liger_kernel.transformers.rope import liger_rotary_pos_emb_vision from liger_kernel.transformers.swiglu import LigerBlockSparseTop2MLP from liger_kernel.transformers.swiglu import LigerExperts from liger_kernel.transformers.swiglu import LigerPhi3SwiGLUMLP from liger_kernel.transformers.swiglu import LigerSwiGLUMLP +from liger_kernel.transformers.swiglu import LigerSwiGLUMLPForMuseGlimmer try: import peft @@ -100,6 +103,28 @@ def _patch_rms_norm_module(module, offset=0.0, eps=1e-6, casting_mode="llama", i _bind_method_to_module(module, "_get_name", lambda self: LigerRMSNorm.__name__) +def _patch_muse_glimmer_scale_free_rms_norm_module(module, eps=1e-6): + # MuseGlimmerRMSNorm(with_scale=False) has no weight parameter, so the Liger kernel has + # nothing to scale by. Bind the torch fallback that matches HF exactly. + + assert getattr(module, "weight", None) is None, ( + f"{type(module).__name__} has a weight parameter and cannot use the scale-free " + "fallback -- use _patch_rms_norm_module(offset=0.0, casting_mode='gemma') instead." + ) + module.variance_epsilon = getattr(module, "variance_epsilon", None) or getattr(module, "eps", None) or eps + module.with_scale = False + module.offset = 0.0 + module.casting_mode = "gemma" + module.in_place = False + module.row_mode = None + _bind_method_to_module(module, "forward", LigerRMSNormForMuseGlimmer.forward) + _bind_method_to_module(module, "_get_name", lambda self: LigerRMSNormForMuseGlimmer.__name__) + + +def _patch_muse_glimmer_text_centered_rms_norm_module(module, eps=1e-6): + _patch_rms_norm_module(module, offset=1.0, eps=eps, casting_mode="gemma", in_place=False) + + def _patch_layer_norm_module(module, eps=1e-6): # Check if the module is a PEFT ModulesToSaveWrapper # If it is, we need to patch the modules_to_save.default and original_modules @@ -869,6 +894,123 @@ def apply_liger_kernel_to_mixtral( _patch_rms_norm_module(decoder_layer.post_attention_layernorm) +def apply_liger_kernel_to_muse_glimmer( + rope: bool = True, + cross_entropy: bool = False, + fused_linear_cross_entropy: bool = True, + layer_norm: bool = True, + rms_norm: bool = True, + swiglu: bool = True, + model: PreTrainedModel = None, +) -> None: + """ + Apply Liger kernels to HuggingFace MuseGlimmer models. + + Vision RoPE is left unchanged because its 4D layout is incompatible with + Liger's 3D vision RoPE kernel. + + Args: + rope: Patch text RoPE. Default: True. + cross_entropy: Patch cross entropy. Default: False. + fused_linear_cross_entropy: Patch fused cross entropy. Default: True. + layer_norm: Patch vision LayerNorm when `model` is provided. Default: True. + rms_norm: Patch RMSNorm. Default: True. + swiglu: Patch SwiGLU. Default: True. + model: Existing model instance to patch. Default: None. + """ + assert not (cross_entropy and fused_linear_cross_entropy), ( + "cross_entropy and fused_linear_cross_entropy cannot both be True." + ) + + from transformers.models.muse_glimmer import modeling_muse_glimmer + from transformers.models.muse_glimmer.modeling_muse_glimmer import MuseGlimmerForConditionalGeneration + from transformers.models.muse_glimmer.modeling_muse_glimmer import MuseGlimmerModel + from transformers.models.muse_glimmer.modeling_muse_glimmer import MuseGlimmerTextModel + + from liger_kernel.transformers.model.muse_glimmer import lce_forward as muse_glimmer_lce_forward + + if rope: + modeling_muse_glimmer.apply_rotary_pos_emb = liger_rotary_pos_emb + + if rms_norm: + modeling_muse_glimmer.MuseGlimmerRMSNorm = LigerRMSNormForMuseGlimmer + modeling_muse_glimmer.MuseGlimmerTextCenteredRMSNorm = LigerRMSNormForMuseGlimmerTextCentered + + if swiglu: + modeling_muse_glimmer.MuseGlimmerTextMLP = LigerSwiGLUMLPForMuseGlimmer + + if cross_entropy: + from transformers.loss.loss_utils import nn + + nn.functional.cross_entropy = liger_cross_entropy + + if layer_norm and model is None: + # MuseGlimmer vision LayerNorm uses torch.nn.LayerNorm directly, so patching requires a model instance. + logger.warning_once( + "layer_norm=True is a no-op for MuseGlimmer when `model` is None: the vision " + "tower constructs `nn.LayerNorm` directly, so it can only be patched on an " + "existing instance. Pass `model=` to enable the LayerNorm kernel." + ) + + if fused_linear_cross_entropy: + if model is None: + modeling_muse_glimmer.MuseGlimmerForConditionalGeneration.forward = muse_glimmer_lce_forward + elif isinstance(model, MuseGlimmerForConditionalGeneration): + model.forward = MethodType(muse_glimmer_lce_forward, model) + # MuseGlimmerModel and MuseGlimmerTextModel lack lm_head, so CausalLM forward isn't patched; their layer kernels are still patched below. + + if model is not None: + # The model instance already exists, so we need to additionally patch the + # instance variables that reference already-instantiated modules + + if isinstance(model, MuseGlimmerForConditionalGeneration): + base_model = model.model + elif isinstance(model, MuseGlimmerModel): + base_model = model + elif isinstance(model, MuseGlimmerTextModel): + base_model = None + else: + raise TypeError( + "Unsupported MuseGlimmer model type. `model` must be `MuseGlimmerForConditionalGeneration`, " + f"`MuseGlimmerModel` or `MuseGlimmerTextModel`. Got: {type(model)}" + ) + + text_model = model if base_model is None else base_model.language_model + vision_model = None if base_model is None else getattr(base_model, "vision_tower", None) + + if rms_norm: + if base_model is not None and getattr(base_model, "perception_emb_norm", None) is not None: + _patch_muse_glimmer_scale_free_rms_norm_module(base_model.perception_emb_norm) + + if text_model is not None: + _patch_rms_norm_module(text_model.norm, offset=0.0, casting_mode="gemma", in_place=False) + + embed_tokens = getattr(text_model, "embed_tokens", None) + if embed_tokens is not None and getattr(embed_tokens, "embed_norm", None) is not None: + _patch_muse_glimmer_scale_free_rms_norm_module(embed_tokens.embed_norm) + + if text_model is not None: + for decoder_layer in text_model.layers: + if swiglu: + _patch_swiglu_module(decoder_layer.mlp, LigerSwiGLUMLPForMuseGlimmer) + if rms_norm: + _patch_muse_glimmer_text_centered_rms_norm_module(decoder_layer.input_layernorm) + _patch_muse_glimmer_text_centered_rms_norm_module(decoder_layer.post_attention_layernorm) + _patch_muse_glimmer_text_centered_rms_norm_module(decoder_layer.pre_feedforward_layernorm) + _patch_muse_glimmer_text_centered_rms_norm_module(decoder_layer.post_feedforward_layernorm) + + self_attn = getattr(decoder_layer, "self_attn", None) + if self_attn is not None and getattr(self_attn, "qk_norm", None) is not None: + _patch_muse_glimmer_scale_free_rms_norm_module(self_attn.qk_norm) + + if layer_norm and vision_model is not None: + _patch_layer_norm_module(vision_model.ln_pre) + _patch_layer_norm_module(vision_model.ln_post) + for vision_layer in vision_model.layers: + _patch_layer_norm_module(vision_layer.norm1) + _patch_layer_norm_module(vision_layer.norm2) + + def apply_liger_kernel_to_pixtral( rope: bool = True, rms_norm: bool = True, @@ -3554,6 +3696,7 @@ def __init__(self, hidden_size, eps=1e-6, **kwargs): "ministral": apply_liger_kernel_to_ministral, "mistral": apply_liger_kernel_to_mistral, "mixtral": apply_liger_kernel_to_mixtral, + "muse_glimmer": apply_liger_kernel_to_muse_glimmer, "nemotron": apply_liger_kernel_to_nemotron, "olmo2": apply_liger_kernel_to_olmo2, "pixtral": apply_liger_kernel_to_pixtral, diff --git a/src/liger_kernel/transformers/rms_norm.py b/src/liger_kernel/transformers/rms_norm.py index 03a01a574..ed6123cc6 100644 --- a/src/liger_kernel/transformers/rms_norm.py +++ b/src/liger_kernel/transformers/rms_norm.py @@ -110,6 +110,55 @@ def forward(self, hidden_states): return super().forward(hidden_states) +class LigerRMSNormForMuseGlimmer(LigerRMSNorm): + """MuseGlimmerRMSNorm semantics (see transformers.models.muse_glimmer.modeling_muse_glimmer): + + - weight initialized to ones, applied directly (no ``(1 + w)`` offset) + - fp32 compute, cast back to input dtype (gemma-style casting) + - ``with_scale=False`` variant has NO weight parameter and is used for ``qk_norm``, + the embedding ``embed_norm`` and ``perception_emb_norm``. + + When ``with_scale=False`` the Liger kernel has no weight to multiply by, so we fall back + to a plain torch implementation that matches HF exactly. + """ + + def __init__( + self, + dim=None, + eps=1e-6, + offset=0.0, + casting_mode="gemma", + init_fn="ones", + in_place=False, + with_scale=True, + ): + super().__init__(dim, eps, offset, casting_mode, init_fn, in_place, elementwise_affine=with_scale) + self.with_scale = with_scale + + def forward(self, hidden_states): + if not self.with_scale: + # Mirrors HF's MuseGlimmerRMSNorm forward for the with_scale=False case: + # scale-free RMS normalization with fp32 compute, cast back to input dtype. + input_dtype = hidden_states.dtype + x = hidden_states.float() + mean_sq = x.pow(2).mean(-1, keepdim=True) + self.variance_epsilon + return (x * torch.pow(mean_sq, -0.5)).to(input_dtype) + return super().forward(hidden_states) + + +class LigerRMSNormForMuseGlimmerTextCentered(LigerRMSNorm): + """MuseGlimmerTextCenteredRMSNorm scales by ``(1 + weight)`` with zero-initialized weights, + computing in fp32 and casting back — i.e. Gemma semantics. ``in_place=False`` because each + decoder layer chains ``pre_feedforward_layernorm`` and ``post_feedforward_layernorm`` around + a residual connection. + """ + + def __init__( + self, hidden_size, eps=1e-6, offset=1.0, casting_mode="gemma", init_fn="zeros", in_place=False, row_mode=None + ): + super().__init__(hidden_size, eps, offset, casting_mode, init_fn, in_place, row_mode) + + class LigerRMSNormForOlmo2(LigerRMSNorm): def __init__( self, hidden_size, eps=1e-6, offset=0.0, casting_mode="llama", init_fn="ones", in_place=False, row_mode=None diff --git a/src/liger_kernel/transformers/swiglu.py b/src/liger_kernel/transformers/swiglu.py index 4836478e5..2057d5e86 100644 --- a/src/liger_kernel/transformers/swiglu.py +++ b/src/liger_kernel/transformers/swiglu.py @@ -21,6 +21,26 @@ def forward(self, x): return self.down_proj(LigerSiLUMulFunction.apply(self.gate_proj(x), self.up_proj(x))) +class LigerSwiGLUMLPForMuseGlimmer(LigerSwiGLUMLP): + """SwiGLU MLP wrapper for MuseGlimmerTextMLP. + + MuseGlimmerTextConfig names its activation field ``hidden_activation`` (Gemma-style) + rather than ``hidden_act``, so the base class' validation would raise AttributeError. + """ + + def __init__(self, config): + nn.Module.__init__(self) + self.config = config + self.hidden_size = config.hidden_size + self.intermediate_size = config.intermediate_size + self.gate_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False) + self.up_proj = nn.Linear(self.hidden_size, self.intermediate_size, bias=False) + self.down_proj = nn.Linear(self.intermediate_size, self.hidden_size, bias=False) + hidden_activation = getattr(config, "hidden_activation", None) or getattr(config, "hidden_act", None) + if hidden_activation not in ["silu", "swish"]: + raise ValueError(f"Activation function {hidden_activation} not supported.") + + class LigerBlockSparseTop2MLP(nn.Module): def __init__(self, config): super().__init__() diff --git a/test/convergence/bf16/test_mini_models.py b/test/convergence/bf16/test_mini_models.py index c78a97477..115d859e5 100644 --- a/test/convergence/bf16/test_mini_models.py +++ b/test/convergence/bf16/test_mini_models.py @@ -46,6 +46,7 @@ from liger_kernel.transformers import apply_liger_kernel_to_mistral from liger_kernel.transformers import apply_liger_kernel_to_mixtral from liger_kernel.transformers import apply_liger_kernel_to_mllama +from liger_kernel.transformers import apply_liger_kernel_to_muse_glimmer from liger_kernel.transformers import apply_liger_kernel_to_nemotron from liger_kernel.transformers import apply_liger_kernel_to_olmo2 from liger_kernel.transformers import apply_liger_kernel_to_olmo3 @@ -90,6 +91,7 @@ from test.utils import revert_liger_kernel_to_mistral from test.utils import revert_liger_kernel_to_mixtral from test.utils import revert_liger_kernel_to_mllama +from test.utils import revert_liger_kernel_to_muse_glimmer from test.utils import revert_liger_kernel_to_nemotron from test.utils import revert_liger_kernel_to_olmo2 from test.utils import revert_liger_kernel_to_olmo3 @@ -174,6 +176,21 @@ QWEN3_VL_AVAILABLE = False +try: + # MuseGlimmer is only available in transformers>=5.15.0 + import transformers + + from packaging import version + from transformers.models.muse_glimmer.configuration_muse_glimmer import MuseGlimmerConfig + from transformers.models.muse_glimmer.configuration_muse_glimmer import MuseGlimmerTextConfig + from transformers.models.muse_glimmer.configuration_muse_glimmer import MuseGlimmerVisionConfig + from transformers.models.muse_glimmer.modeling_muse_glimmer import MuseGlimmerForConditionalGeneration + + MUSE_GLIMMER_AVAILABLE = version.parse(transformers.__version__) >= version.parse("5.15.0") +except ImportError: + MUSE_GLIMMER_AVAILABLE = False + + try: # Qwen3-VL-MoE is only available in transformers>=4.57.0 import transformers @@ -965,6 +982,58 @@ ) +if MUSE_GLIMMER_AVAILABLE: + MINI_MODEL_SETUPS["mini_muse_glimmer"] = MiniModelConfig( + liger_kernel_patch_func=apply_liger_kernel_to_muse_glimmer, + liger_kernel_patch_revert_func=revert_liger_kernel_to_muse_glimmer, + model_class=MuseGlimmerForConditionalGeneration, + mini_model_config=MuseGlimmerConfig( + attn_implementation="sdpa", + image_token_id=32768, + video_token_id=32769, + # The vision adapter consumes `vision hidden_size * merge_size ** 2` after pixel shuffle + out_hidden_size=128 * 2**2, + projector_hidden_size=256, + projector_hidden_act="gelu", + text_config=MuseGlimmerTextConfig( + bos_token_id=1, + eos_token_id=2, + pad_token_id=None, + vocab_size=32000, + hidden_size=896, + intermediate_size=2176, + num_hidden_layers=4, + num_attention_heads=8, + num_key_value_heads=2, + head_dim=112, + hidden_activation="silu", + max_position_embeddings=4096, + initializer_range=0.02, + rms_norm_eps=1e-5, + post_norm_eps=1e-8, + sliding_window=128, + attention_dropout=0.0, + attention_bias=False, + tie_word_embeddings=False, + use_cache=True, + ), + vision_config=MuseGlimmerVisionConfig( + hidden_size=128, + intermediate_size=256, + num_hidden_layers=2, + num_attention_heads=4, + hidden_act="gelu", + patch_size=14, + pos_emb_height=4, + pos_emb_width=4, + max_position_embeddings=16, + merge_size=2, + layer_norm_eps=1e-5, + ), + ), + ) + + if QWEN3_VL_AVAILABLE: MINI_MODEL_SETUPS["mini_qwen3_vl"] = MiniModelConfig( liger_kernel_patch_func=apply_liger_kernel_to_qwen3_vl, @@ -2039,6 +2108,25 @@ def run_mini_model( ), ], ), + pytest.param( + "mini_muse_glimmer", + 32, + 1e-5, + torch.bfloat16, + 5e-3, + 1e-2, + 1e-1, + 1e-2, + 1e-2, + 1e-2, + marks=[ + pytest.mark.skipif(not supports_bfloat16(), reason="bfloat16 not supported on this GPU"), + pytest.mark.skipif( + not MUSE_GLIMMER_AVAILABLE, + reason="MuseGlimmer not available in this version of transformers", + ), + ], + ), pytest.param( "mini_qwen3_vl", 32, diff --git a/test/convergence/bf16/test_mini_models_multimodal.py b/test/convergence/bf16/test_mini_models_multimodal.py index 0e8097e68..07ebc6626 100644 --- a/test/convergence/bf16/test_mini_models_multimodal.py +++ b/test/convergence/bf16/test_mini_models_multimodal.py @@ -18,6 +18,7 @@ from liger_kernel.transformers import apply_liger_kernel_to_llama4 from liger_kernel.transformers import apply_liger_kernel_to_llava from liger_kernel.transformers import apply_liger_kernel_to_mllama +from liger_kernel.transformers import apply_liger_kernel_to_muse_glimmer from liger_kernel.transformers import apply_liger_kernel_to_paligemma from liger_kernel.transformers import apply_liger_kernel_to_pixtral from liger_kernel.transformers import apply_liger_kernel_to_qwen2_5_vl @@ -46,6 +47,7 @@ from test.utils import revert_liger_kernel_to_llama4 from test.utils import revert_liger_kernel_to_llava from test.utils import revert_liger_kernel_to_mllama +from test.utils import revert_liger_kernel_to_muse_glimmer from test.utils import revert_liger_kernel_to_Paligemma from test.utils import revert_liger_kernel_to_pixtral from test.utils import revert_liger_kernel_to_qwen2_5_vl @@ -123,6 +125,24 @@ except ImportError: QWEN3_VL_AVAILABLE = False + +try: + # MuseGlimmer is only available in transformers>=5.15.0 + import transformers + + from packaging import version + from transformers.models.muse_glimmer.configuration_muse_glimmer import MuseGlimmerConfig + from transformers.models.muse_glimmer.configuration_muse_glimmer import MuseGlimmerTextConfig + from transformers.models.muse_glimmer.configuration_muse_glimmer import MuseGlimmerVisionConfig + from transformers.models.muse_glimmer.image_processing_muse_glimmer import MuseGlimmerImageProcessor + from transformers.models.muse_glimmer.modeling_muse_glimmer import MuseGlimmerForConditionalGeneration + from transformers.models.muse_glimmer.processing_muse_glimmer import MuseGlimmerProcessor + from transformers.models.muse_glimmer.video_processing_muse_glimmer import MuseGlimmerVideoProcessor + + MUSE_GLIMMER_AVAILABLE = version.parse(transformers.__version__) >= version.parse("5.15.0") +except ImportError: + MUSE_GLIMMER_AVAILABLE = False + try: from transformers.models.qwen3_vl_moe.configuration_qwen3_vl_moe import Qwen3VLMoeConfig from transformers.models.qwen3_vl_moe.configuration_qwen3_vl_moe import Qwen3VLMoeTextConfig @@ -813,6 +833,60 @@ ), ) + +if MUSE_GLIMMER_AVAILABLE: + MINI_MODEL_SETUPS["mini_muse_glimmer"] = MiniModelConfig( + liger_kernel_patch_func=apply_liger_kernel_to_muse_glimmer, + liger_kernel_patch_revert_func=revert_liger_kernel_to_muse_glimmer, + model_class=MuseGlimmerForConditionalGeneration, + mini_model_config=MuseGlimmerConfig( + attn_implementation="sdpa", + # Must match the ids the fake tokenizer assigns to `<|patch|>` / `<|video|>` + image_token_id=11, + video_token_id=10, + # The vision adapter consumes `vision hidden_size * merge_size ** 2` after pixel shuffle + out_hidden_size=128 * 2**2, + projector_hidden_size=256, + projector_hidden_act="gelu", + text_config=MuseGlimmerTextConfig( + bos_token_id=0, + eos_token_id=1, + pad_token_id=2, + vocab_size=32000, + hidden_size=512, + intermediate_size=1024, + num_hidden_layers=4, + num_attention_heads=8, + num_key_value_heads=2, + head_dim=64, + hidden_activation="silu", + max_position_embeddings=4096, + initializer_range=0.02, + rms_norm_eps=1e-5, + post_norm_eps=1e-8, + sliding_window=128, + attention_dropout=0.0, + attention_bias=False, + tie_word_embeddings=False, + use_cache=False, + ), + vision_config=MuseGlimmerVisionConfig( + hidden_size=128, + intermediate_size=256, + num_hidden_layers=2, + num_attention_heads=4, + hidden_act="gelu", + patch_size=14, + pos_emb_height=8, + pos_emb_width=8, + max_position_embeddings=64, + merge_size=2, + layer_norm_eps=1e-5, + ), + ), + ) + + if QWEN3_VL_AVAILABLE: MINI_MODEL_SETUPS["mini_qwen3_vl"] = MiniModelConfig( liger_kernel_patch_func=apply_liger_kernel_to_qwen3_vl, @@ -1106,6 +1180,29 @@ def create_processor(model_name: str): tokenizer=qwen_tokenizer, ) + elif model_name == "mini_muse_glimmer": + tokenizer_config = load_tokenizer_config( + os.path.join(FAKE_CONFIGS_PATH, "meta-models/Muse-Glimmer-30B/tokenizer_config.json") + ) + tokenizer_base = train_bpe_tokenizer( + [ + token.content + for key, token in sorted( + tokenizer_config["added_tokens_decoder"].items(), + key=lambda x: int(x[0]), + ) + ] + ) + muse_tokenizer = PreTrainedTokenizerFast(tokenizer_object=tokenizer_base, **tokenizer_config) + # `max_image_tokens` is capped so the 64x64 procedural test image stays cheap. + image_processor = MuseGlimmerImageProcessor(patch_size=14, merge_size=2, max_image_tokens=256) + video_processor = MuseGlimmerVideoProcessor(patch_size=14, merge_size=2) + return MuseGlimmerProcessor( + image_processor=image_processor, + video_processor=video_processor, + tokenizer=muse_tokenizer, + ) + elif model_name in ("mini_qwen3_vl", "mini_qwen3_vl_moe", "mini_qwen3_5", "mini_qwen3_5_moe"): tokenizer_config = load_tokenizer_config( os.path.join(FAKE_CONFIGS_PATH, "Qwen/Qwen3-VL-4B-Instruct/tokenizer_config.json") @@ -1612,6 +1709,29 @@ def run_mini_model_multimodal( pytest.mark.skipif(not is_torchvision_available(), reason="Qwen2VLVideoProcessor requires torchvision"), ], ), + pytest.param( + "mini_muse_glimmer", + 32, + 1e-5, + torch.bfloat16, + 5e-2, + 5e-2, + 1e-1, + 1e-2, + 1e-2, + 1e-2, + marks=[ + pytest.mark.skipif(not supports_bfloat16(), reason="bfloat16 not supported on this GPU"), + pytest.mark.skipif( + not MUSE_GLIMMER_AVAILABLE, + reason="MuseGlimmer not available in this version of transformers", + ), + pytest.mark.skipif( + not is_torchvision_available(), + reason="MuseGlimmerImageProcessor requires torchvision", + ), + ], + ), pytest.param( "mini_qwen3_vl", 32, diff --git a/test/convergence/bf16/test_mini_models_with_logits.py b/test/convergence/bf16/test_mini_models_with_logits.py index fa7a53ec0..d86a2cfdd 100644 --- a/test/convergence/bf16/test_mini_models_with_logits.py +++ b/test/convergence/bf16/test_mini_models_with_logits.py @@ -45,6 +45,7 @@ from liger_kernel.transformers import apply_liger_kernel_to_mistral from liger_kernel.transformers import apply_liger_kernel_to_mixtral from liger_kernel.transformers import apply_liger_kernel_to_mllama +from liger_kernel.transformers import apply_liger_kernel_to_muse_glimmer from liger_kernel.transformers import apply_liger_kernel_to_nemotron from liger_kernel.transformers import apply_liger_kernel_to_olmo2 from liger_kernel.transformers import apply_liger_kernel_to_olmo3 @@ -87,6 +88,7 @@ from test.utils import revert_liger_kernel_to_mistral from test.utils import revert_liger_kernel_to_mixtral from test.utils import revert_liger_kernel_to_mllama +from test.utils import revert_liger_kernel_to_muse_glimmer from test.utils import revert_liger_kernel_to_nemotron from test.utils import revert_liger_kernel_to_olmo2 from test.utils import revert_liger_kernel_to_olmo3 @@ -174,6 +176,21 @@ except ImportError: QWEN3_VL_AVAILABLE = False + +try: + # MuseGlimmer is only available in transformers>=5.15.0 + import transformers + + from packaging import version + from transformers.models.muse_glimmer.configuration_muse_glimmer import MuseGlimmerConfig + from transformers.models.muse_glimmer.configuration_muse_glimmer import MuseGlimmerTextConfig + from transformers.models.muse_glimmer.configuration_muse_glimmer import MuseGlimmerVisionConfig + from transformers.models.muse_glimmer.modeling_muse_glimmer import MuseGlimmerForConditionalGeneration + + MUSE_GLIMMER_AVAILABLE = version.parse(transformers.__version__) >= version.parse("5.15.0") +except ImportError: + MUSE_GLIMMER_AVAILABLE = False + try: import transformers @@ -687,6 +704,59 @@ ), ) + +if MUSE_GLIMMER_AVAILABLE: + MINI_MODEL_SETUPS["mini_muse_glimmer"] = MiniModelConfig( + liger_kernel_patch_func=apply_liger_kernel_to_muse_glimmer, + liger_kernel_patch_revert_func=revert_liger_kernel_to_muse_glimmer, + model_class=MuseGlimmerForConditionalGeneration, + mini_model_config=MuseGlimmerConfig( + attn_implementation="sdpa", + image_token_id=32768, + video_token_id=32769, + # The vision adapter consumes `vision hidden_size * merge_size ** 2` after pixel shuffle + out_hidden_size=128 * 2**2, + projector_hidden_size=256, + projector_hidden_act="gelu", + text_config=MuseGlimmerTextConfig( + bos_token_id=1, + eos_token_id=2, + pad_token_id=None, + vocab_size=32000, + hidden_size=896, + intermediate_size=2176, + num_hidden_layers=4, + num_attention_heads=8, + num_key_value_heads=2, + head_dim=112, + hidden_activation="silu", + max_position_embeddings=4096, + initializer_range=0.02, + rms_norm_eps=1e-5, + post_norm_eps=1e-8, + sliding_window=128, + attention_dropout=0.0, + attention_bias=False, + tie_word_embeddings=False, + use_cache=True, + ), + vision_config=MuseGlimmerVisionConfig( + hidden_size=128, + intermediate_size=256, + num_hidden_layers=2, + num_attention_heads=4, + hidden_act="gelu", + patch_size=14, + pos_emb_height=4, + pos_emb_width=4, + max_position_embeddings=16, + merge_size=2, + layer_norm_eps=1e-5, + ), + ), + ) + + if QWEN3_VL_AVAILABLE: MINI_MODEL_SETUPS["mini_qwen3_vl"] = MiniModelConfig( liger_kernel_patch_func=apply_liger_kernel_to_qwen3_vl, @@ -1872,6 +1942,25 @@ def run_mini_model( ), ], ), + pytest.param( + "mini_muse_glimmer", + 32, + 1e-5, + torch.bfloat16, + 5e-3, + 1e-2, + 1e-1, + 1e-2, + 1e-2, + 1e-2, + marks=[ + pytest.mark.skipif(not supports_bfloat16(), reason="bfloat16 not supported on this GPU"), + pytest.mark.skipif( + not MUSE_GLIMMER_AVAILABLE, + reason="MuseGlimmer not available in this version of transformers", + ), + ], + ), pytest.param( "mini_qwen3_vl", 32, diff --git a/test/convergence/fp32/test_mini_models.py b/test/convergence/fp32/test_mini_models.py index eff183b36..a7f1dcb33 100644 --- a/test/convergence/fp32/test_mini_models.py +++ b/test/convergence/fp32/test_mini_models.py @@ -45,6 +45,7 @@ from liger_kernel.transformers import apply_liger_kernel_to_mistral from liger_kernel.transformers import apply_liger_kernel_to_mixtral from liger_kernel.transformers import apply_liger_kernel_to_mllama +from liger_kernel.transformers import apply_liger_kernel_to_muse_glimmer from liger_kernel.transformers import apply_liger_kernel_to_nemotron from liger_kernel.transformers import apply_liger_kernel_to_olmo2 from liger_kernel.transformers import apply_liger_kernel_to_olmo3 @@ -88,6 +89,7 @@ from test.utils import revert_liger_kernel_to_mistral from test.utils import revert_liger_kernel_to_mixtral from test.utils import revert_liger_kernel_to_mllama +from test.utils import revert_liger_kernel_to_muse_glimmer from test.utils import revert_liger_kernel_to_nemotron from test.utils import revert_liger_kernel_to_olmo2 from test.utils import revert_liger_kernel_to_olmo3 @@ -171,6 +173,21 @@ QWEN3_VL_AVAILABLE = False +try: + # MuseGlimmer is only available in transformers>=5.15.0 + import transformers + + from packaging import version + from transformers.models.muse_glimmer.configuration_muse_glimmer import MuseGlimmerConfig + from transformers.models.muse_glimmer.configuration_muse_glimmer import MuseGlimmerTextConfig + from transformers.models.muse_glimmer.configuration_muse_glimmer import MuseGlimmerVisionConfig + from transformers.models.muse_glimmer.modeling_muse_glimmer import MuseGlimmerForConditionalGeneration + + MUSE_GLIMMER_AVAILABLE = version.parse(transformers.__version__) >= version.parse("5.15.0") +except ImportError: + MUSE_GLIMMER_AVAILABLE = False + + try: # Qwen3-VL-MoE is only available in transformers>=4.57.0 import transformers @@ -901,6 +918,59 @@ ), ) + +if MUSE_GLIMMER_AVAILABLE: + MINI_MODEL_SETUPS["mini_muse_glimmer"] = MiniModelConfig( + liger_kernel_patch_func=apply_liger_kernel_to_muse_glimmer, + liger_kernel_patch_revert_func=revert_liger_kernel_to_muse_glimmer, + model_class=MuseGlimmerForConditionalGeneration, + mini_model_config=MuseGlimmerConfig( + attn_implementation="sdpa", + image_token_id=32768, + video_token_id=32769, + # The vision adapter consumes `vision hidden_size * merge_size ** 2` after pixel shuffle + out_hidden_size=128 * 2**2, + projector_hidden_size=256, + projector_hidden_act="gelu", + text_config=MuseGlimmerTextConfig( + bos_token_id=1, + eos_token_id=2, + pad_token_id=None, + vocab_size=32000, + hidden_size=896, + intermediate_size=2176, + num_hidden_layers=4, + num_attention_heads=8, + num_key_value_heads=2, + head_dim=112, + hidden_activation="silu", + max_position_embeddings=4096, + initializer_range=0.02, + rms_norm_eps=1e-5, + post_norm_eps=1e-8, + sliding_window=128, + attention_dropout=0.0, + attention_bias=False, + tie_word_embeddings=False, + use_cache=True, + ), + vision_config=MuseGlimmerVisionConfig( + hidden_size=128, + intermediate_size=256, + num_hidden_layers=2, + num_attention_heads=4, + hidden_act="gelu", + patch_size=14, + pos_emb_height=4, + pos_emb_width=4, + max_position_embeddings=16, + merge_size=2, + layer_norm_eps=1e-5, + ), + ), + ) + + if QWEN3_VL_AVAILABLE: MINI_MODEL_SETUPS["mini_qwen3_vl"] = MiniModelConfig( liger_kernel_patch_func=apply_liger_kernel_to_qwen3_vl, @@ -1922,6 +1992,22 @@ def run_mini_model( reason="Qwen2.5-VL not available in this version of transformers", ), ), + pytest.param( + "mini_muse_glimmer", + 32, + 1e-4, + torch.float32, + 1e-8, + 2e-5, + 5e-3, + 1e-5, + 5e-3, + 1e-5, + marks=pytest.mark.skipif( + not MUSE_GLIMMER_AVAILABLE, + reason="MuseGlimmer not available in this version of transformers", + ), + ), pytest.param( "mini_qwen3_vl", 32, diff --git a/test/convergence/fp32/test_mini_models_multimodal.py b/test/convergence/fp32/test_mini_models_multimodal.py index e9e3e6361..bce0ff225 100644 --- a/test/convergence/fp32/test_mini_models_multimodal.py +++ b/test/convergence/fp32/test_mini_models_multimodal.py @@ -19,6 +19,7 @@ from liger_kernel.transformers import apply_liger_kernel_to_llama4 from liger_kernel.transformers import apply_liger_kernel_to_llava from liger_kernel.transformers import apply_liger_kernel_to_mllama +from liger_kernel.transformers import apply_liger_kernel_to_muse_glimmer from liger_kernel.transformers import apply_liger_kernel_to_paligemma from liger_kernel.transformers import apply_liger_kernel_to_pixtral from liger_kernel.transformers import apply_liger_kernel_to_qwen2_5_vl @@ -47,6 +48,7 @@ from test.utils import revert_liger_kernel_to_llama4 from test.utils import revert_liger_kernel_to_llava from test.utils import revert_liger_kernel_to_mllama +from test.utils import revert_liger_kernel_to_muse_glimmer from test.utils import revert_liger_kernel_to_Paligemma from test.utils import revert_liger_kernel_to_pixtral from test.utils import revert_liger_kernel_to_qwen2_5_vl @@ -153,6 +155,24 @@ QWEN3_VL_AVAILABLE = False +try: + # MuseGlimmer is only available in transformers>=5.15.0 + import transformers + + from packaging import version + from transformers.models.muse_glimmer.configuration_muse_glimmer import MuseGlimmerConfig + from transformers.models.muse_glimmer.configuration_muse_glimmer import MuseGlimmerTextConfig + from transformers.models.muse_glimmer.configuration_muse_glimmer import MuseGlimmerVisionConfig + from transformers.models.muse_glimmer.image_processing_muse_glimmer import MuseGlimmerImageProcessor + from transformers.models.muse_glimmer.modeling_muse_glimmer import MuseGlimmerForConditionalGeneration + from transformers.models.muse_glimmer.processing_muse_glimmer import MuseGlimmerProcessor + from transformers.models.muse_glimmer.video_processing_muse_glimmer import MuseGlimmerVideoProcessor + + MUSE_GLIMMER_AVAILABLE = version.parse(transformers.__version__) >= version.parse("5.15.0") +except ImportError: + MUSE_GLIMMER_AVAILABLE = False + + try: from transformers.models.qwen3_vl_moe.configuration_qwen3_vl_moe import Qwen3VLMoeConfig from transformers.models.qwen3_vl_moe.configuration_qwen3_vl_moe import Qwen3VLMoeTextConfig @@ -842,6 +862,60 @@ ), ) + +if MUSE_GLIMMER_AVAILABLE: + MINI_MODEL_SETUPS["mini_muse_glimmer"] = MiniModelConfig( + liger_kernel_patch_func=apply_liger_kernel_to_muse_glimmer, + liger_kernel_patch_revert_func=revert_liger_kernel_to_muse_glimmer, + model_class=MuseGlimmerForConditionalGeneration, + mini_model_config=MuseGlimmerConfig( + attn_implementation="sdpa", + # Must match the ids the fake tokenizer assigns to `<|patch|>` / `<|video|>` + image_token_id=11, + video_token_id=10, + # The vision adapter consumes `vision hidden_size * merge_size ** 2` after pixel shuffle + out_hidden_size=128 * 2**2, + projector_hidden_size=256, + projector_hidden_act="gelu", + text_config=MuseGlimmerTextConfig( + bos_token_id=0, + eos_token_id=1, + pad_token_id=2, + vocab_size=32000, + hidden_size=512, + intermediate_size=1024, + num_hidden_layers=4, + num_attention_heads=8, + num_key_value_heads=2, + head_dim=64, + hidden_activation="silu", + max_position_embeddings=4096, + initializer_range=0.02, + rms_norm_eps=1e-5, + post_norm_eps=1e-8, + sliding_window=128, + attention_dropout=0.0, + attention_bias=False, + tie_word_embeddings=False, + use_cache=False, + ), + vision_config=MuseGlimmerVisionConfig( + hidden_size=128, + intermediate_size=256, + num_hidden_layers=2, + num_attention_heads=4, + hidden_act="gelu", + patch_size=14, + pos_emb_height=8, + pos_emb_width=8, + max_position_embeddings=64, + merge_size=2, + layer_norm_eps=1e-5, + ), + ), + ) + + if QWEN3_VL_AVAILABLE: MINI_MODEL_SETUPS["mini_qwen3_vl"] = MiniModelConfig( liger_kernel_patch_func=apply_liger_kernel_to_qwen3_vl, @@ -1223,6 +1297,29 @@ def create_processor(model_name: str): tokenizer=qwen_tokenizer, ) + elif model_name == "mini_muse_glimmer": + tokenizer_config = load_tokenizer_config( + os.path.join(FAKE_CONFIGS_PATH, "meta-models/Muse-Glimmer-30B/tokenizer_config.json") + ) + tokenizer_base = train_bpe_tokenizer( + [ + token.content + for key, token in sorted( + tokenizer_config["added_tokens_decoder"].items(), + key=lambda x: int(x[0]), + ) + ] + ) + muse_tokenizer = PreTrainedTokenizerFast(tokenizer_object=tokenizer_base, **tokenizer_config) + # `max_image_tokens` is capped so the 64x64 procedural test image stays cheap. + image_processor = MuseGlimmerImageProcessor(patch_size=14, merge_size=2, max_image_tokens=256) + video_processor = MuseGlimmerVideoProcessor(patch_size=14, merge_size=2) + return MuseGlimmerProcessor( + image_processor=image_processor, + video_processor=video_processor, + tokenizer=muse_tokenizer, + ) + elif model_name in ("mini_qwen3_vl", "mini_qwen3_vl_moe", "mini_qwen3_5", "mini_qwen3_5_moe"): tokenizer_config = load_tokenizer_config( os.path.join(FAKE_CONFIGS_PATH, "Qwen/Qwen3-VL-4B-Instruct/tokenizer_config.json") @@ -1744,6 +1841,28 @@ def run_mini_model_multimodal( pytest.mark.skipif(not is_torchvision_available(), reason="Qwen2VLVideoProcessor requires torchvision"), ], ), + pytest.param( + "mini_muse_glimmer", + 32, + 1e-4, + torch.float32, + 1e-8, + 1e-5, + 5e-3, + 1e-5, + 5e-3, + 1e-5, + marks=[ + pytest.mark.skipif( + not MUSE_GLIMMER_AVAILABLE, + reason="MuseGlimmer not available in this version of transformers", + ), + pytest.mark.skipif( + not is_torchvision_available(), + reason="MuseGlimmerImageProcessor requires torchvision", + ), + ], + ), pytest.param( "mini_qwen3_vl", 32, diff --git a/test/convergence/fp32/test_mini_models_with_logits.py b/test/convergence/fp32/test_mini_models_with_logits.py index e6b9ba3e1..3326dbdc2 100644 --- a/test/convergence/fp32/test_mini_models_with_logits.py +++ b/test/convergence/fp32/test_mini_models_with_logits.py @@ -44,6 +44,7 @@ from liger_kernel.transformers import apply_liger_kernel_to_mistral from liger_kernel.transformers import apply_liger_kernel_to_mixtral from liger_kernel.transformers import apply_liger_kernel_to_mllama +from liger_kernel.transformers import apply_liger_kernel_to_muse_glimmer from liger_kernel.transformers import apply_liger_kernel_to_nemotron from liger_kernel.transformers import apply_liger_kernel_to_olmo2 from liger_kernel.transformers import apply_liger_kernel_to_olmo3 @@ -85,6 +86,7 @@ from test.utils import revert_liger_kernel_to_mistral from test.utils import revert_liger_kernel_to_mixtral from test.utils import revert_liger_kernel_to_mllama +from test.utils import revert_liger_kernel_to_muse_glimmer from test.utils import revert_liger_kernel_to_nemotron from test.utils import revert_liger_kernel_to_olmo2 from test.utils import revert_liger_kernel_to_olmo3 @@ -197,6 +199,21 @@ except ImportError: QWEN3_VL_AVAILABLE = False + +try: + # MuseGlimmer is only available in transformers>=5.15.0 + import transformers + + from packaging import version + from transformers.models.muse_glimmer.configuration_muse_glimmer import MuseGlimmerConfig + from transformers.models.muse_glimmer.configuration_muse_glimmer import MuseGlimmerTextConfig + from transformers.models.muse_glimmer.configuration_muse_glimmer import MuseGlimmerVisionConfig + from transformers.models.muse_glimmer.modeling_muse_glimmer import MuseGlimmerForConditionalGeneration + + MUSE_GLIMMER_AVAILABLE = version.parse(transformers.__version__) >= version.parse("5.15.0") +except ImportError: + MUSE_GLIMMER_AVAILABLE = False + try: from transformers.models.qwen3_vl_moe.configuration_qwen3_vl_moe import Qwen3VLMoeConfig from transformers.models.qwen3_vl_moe.modeling_qwen3_vl_moe import Qwen3VLMoeForConditionalGeneration @@ -864,6 +881,59 @@ ), ) + +if MUSE_GLIMMER_AVAILABLE: + MINI_MODEL_SETUPS["mini_muse_glimmer"] = MiniModelConfig( + liger_kernel_patch_func=apply_liger_kernel_to_muse_glimmer, + liger_kernel_patch_revert_func=revert_liger_kernel_to_muse_glimmer, + model_class=MuseGlimmerForConditionalGeneration, + mini_model_config=MuseGlimmerConfig( + attn_implementation="sdpa", + image_token_id=32768, + video_token_id=32769, + # The vision adapter consumes `vision hidden_size * merge_size ** 2` after pixel shuffle + out_hidden_size=128 * 2**2, + projector_hidden_size=256, + projector_hidden_act="gelu", + text_config=MuseGlimmerTextConfig( + bos_token_id=1, + eos_token_id=2, + pad_token_id=None, + vocab_size=32000, + hidden_size=896, + intermediate_size=2176, + num_hidden_layers=4, + num_attention_heads=8, + num_key_value_heads=2, + head_dim=112, + hidden_activation="silu", + max_position_embeddings=4096, + initializer_range=0.02, + rms_norm_eps=1e-5, + post_norm_eps=1e-8, + sliding_window=128, + attention_dropout=0.0, + attention_bias=False, + tie_word_embeddings=False, + use_cache=True, + ), + vision_config=MuseGlimmerVisionConfig( + hidden_size=128, + intermediate_size=256, + num_hidden_layers=2, + num_attention_heads=4, + hidden_act="gelu", + patch_size=14, + pos_emb_height=4, + pos_emb_width=4, + max_position_embeddings=16, + merge_size=2, + layer_norm_eps=1e-5, + ), + ), + ) + + if QWEN3_VL_AVAILABLE: MINI_MODEL_SETUPS["mini_qwen3_vl"] = MiniModelConfig( liger_kernel_patch_func=apply_liger_kernel_to_qwen3_vl, @@ -1822,6 +1892,22 @@ def run_mini_model( reason="Qwen2.5-VL not available in this version of transformers", ), ), + pytest.param( + "mini_muse_glimmer", + 32, + 1e-4, + torch.float32, + 1e-8, + 2e-5, + 5e-3, + 1e-5, + 5e-3, + 1e-5, + marks=pytest.mark.skipif( + not MUSE_GLIMMER_AVAILABLE, + reason="MuseGlimmer not available in this version of transformers", + ), + ), pytest.param( "mini_qwen3_vl", 32, diff --git a/test/resources/fake_configs/meta-models/Muse-Glimmer-30B/tokenizer_config.json b/test/resources/fake_configs/meta-models/Muse-Glimmer-30B/tokenizer_config.json new file mode 100644 index 000000000..ba4478c20 --- /dev/null +++ b/test/resources/fake_configs/meta-models/Muse-Glimmer-30B/tokenizer_config.json @@ -0,0 +1,132 @@ +{ + "added_tokens_decoder": { + "0": { + "content": "<|begin_of_text|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "1": { + "content": "<|end_of_text|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "2": { + "content": "<|finetune_right_pad|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "3": { + "content": "<|eom|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "4": { + "content": "<|eot|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "5": { + "content": "<|start|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "6": { + "content": "<|message|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "7": { + "content": "<|image_start|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "8": { + "content": "<|image_end|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "9": { + "content": "<|image|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "10": { + "content": "<|video|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "11": { + "content": "<|patch|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "12": { + "content": "<|vid_frame_separator|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "13": { + "content": "<|vid_start|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + }, + "14": { + "content": "<|vid_end|>", + "lstrip": false, + "normalized": false, + "rstrip": false, + "single_word": false, + "special": true + } + }, + "bos_token": "<|begin_of_text|>", + "chat_template": "{{- '<|begin_of_text|>' -}}{%- for message in messages -%}{{- '<|start|>' + message['role'] + '<|message|>' -}}{%- if message['content'] is string -%}{{- message['content'] -}}{%- else -%}{%- for part in message['content'] -%}{%- if part['type'] == 'image' -%}{{- '<|patch|>' -}}{%- elif part['type'] == 'video' -%}{{- '<|video|>' -}}{%- elif part['type'] == 'text' -%}{{- part['text'] -}}{%- endif -%}{%- endfor -%}{%- endif -%}{{- '<|eot|>' -}}{%- endfor -%}{%- if add_generation_prompt -%}{{- '<|start|>assistant<|message|>' -}}{%- endif -%}", + "clean_up_tokenization_spaces": false, + "eos_token": "<|end_of_text|>", + "extra_special_tokens": {}, + "model_max_length": 8192, + "pad_token": "<|finetune_right_pad|>", + "processor_class": "MuseGlimmerProcessor" +} diff --git a/test/transformers/test_monkey_patch.py b/test/transformers/test_monkey_patch.py index 25099f7be..c4fd94562 100755 --- a/test/transformers/test_monkey_patch.py +++ b/test/transformers/test_monkey_patch.py @@ -102,6 +102,15 @@ def is_ministral_available(): return False +def is_muse_glimmer_available(): + try: + import transformers.models.muse_glimmer # noqa: F401 + + return True + except ImportError: + return False + + def is_qwen3_available(): try: import transformers.models.qwen3 # noqa: F401 @@ -513,6 +522,129 @@ def test_apply_liger_kernel_to_instance_for_llama(): pytest.fail(f"An exception occured in extra_expr: {type(e).__name__} - {e}") +@pytest.mark.skipif(not is_muse_glimmer_available(), reason="muse_glimmer module not available") +def test_apply_liger_kernel_to_instance_for_muse_glimmer(): + # Ensure any monkey patching is cleaned up for subsequent tests + with patch("transformers.models.muse_glimmer.modeling_muse_glimmer"): + from transformers.models.muse_glimmer.modeling_muse_glimmer import MuseGlimmerForConditionalGeneration + + from liger_kernel.transformers.model.muse_glimmer import lce_forward as muse_glimmer_lce_forward + from liger_kernel.transformers.rms_norm import LigerRMSNormForMuseGlimmer + + # Instantiate a dummy model + config = transformers.models.muse_glimmer.configuration_muse_glimmer.MuseGlimmerConfig( + attn_implementation="sdpa", + out_hidden_size=128, + projector_hidden_size=64, + vision_config=transformers.models.muse_glimmer.configuration_muse_glimmer.MuseGlimmerVisionConfig( + hidden_size=32, + intermediate_size=64, + num_hidden_layers=2, + num_attention_heads=2, + patch_size=14, + pos_emb_height=4, + pos_emb_width=4, + max_position_embeddings=16, + ), + text_config=transformers.models.muse_glimmer.configuration_muse_glimmer.MuseGlimmerTextConfig( + vocab_size=512, + hidden_size=64, + intermediate_size=128, + num_hidden_layers=4, + num_attention_heads=4, + num_key_value_heads=2, + head_dim=16, + max_position_embeddings=128, + sliding_window=16, + ), + ) + dummy_model_instance = MuseGlimmerForConditionalGeneration._from_config(config) + + assert isinstance(dummy_model_instance, MuseGlimmerForConditionalGeneration) + + text_model = dummy_model_instance.model.language_model + vision_model = dummy_model_instance.model.vision_tower + + # Check that model instance variables are not yet patched with Liger modules + assert inspect.getsource(dummy_model_instance.forward) != inspect.getsource(muse_glimmer_lce_forward) + assert inspect.getsource(text_model.norm.forward) != inspect.getsource(LigerRMSNorm.forward) + assert inspect.getsource(text_model.embed_tokens.embed_norm.forward) != inspect.getsource( + LigerRMSNormForMuseGlimmer.forward + ) + assert inspect.getsource(dummy_model_instance.model.perception_emb_norm.forward) != inspect.getsource( + LigerRMSNormForMuseGlimmer.forward + ) + for decoder_layer in text_model.layers: + assert inspect.getsource(decoder_layer.mlp.forward) != inspect.getsource(LigerSwiGLUMLP.forward) + assert inspect.getsource(decoder_layer.input_layernorm.forward) != inspect.getsource(LigerRMSNorm.forward) + assert inspect.getsource(decoder_layer.post_attention_layernorm.forward) != inspect.getsource( + LigerRMSNorm.forward + ) + assert inspect.getsource(decoder_layer.pre_feedforward_layernorm.forward) != inspect.getsource( + LigerRMSNorm.forward + ) + assert inspect.getsource(decoder_layer.post_feedforward_layernorm.forward) != inspect.getsource( + LigerRMSNorm.forward + ) + assert inspect.getsource(decoder_layer.self_attn.qk_norm.forward) != inspect.getsource( + LigerRMSNormForMuseGlimmer.forward + ) + assert inspect.getsource(vision_model.ln_pre.forward) != inspect.getsource(LigerLayerNorm.forward) + assert inspect.getsource(vision_model.ln_post.forward) != inspect.getsource(LigerLayerNorm.forward) + for vision_layer in vision_model.layers: + assert inspect.getsource(vision_layer.norm1.forward) != inspect.getsource(LigerLayerNorm.forward) + assert inspect.getsource(vision_layer.norm2.forward) != inspect.getsource(LigerLayerNorm.forward) + + # Test applying kernels to the model instance + _apply_liger_kernel_to_instance(model=dummy_model_instance) + + # Check that the model's instance variables were correctly patched with Liger modules + assert inspect.getsource(dummy_model_instance.forward) == inspect.getsource(muse_glimmer_lce_forward) + assert inspect.getsource(text_model.norm.forward) == inspect.getsource(LigerRMSNorm.forward) + assert inspect.getsource(text_model.embed_tokens.embed_norm.forward) == inspect.getsource( + LigerRMSNormForMuseGlimmer.forward + ) + assert inspect.getsource(dummy_model_instance.model.perception_emb_norm.forward) == inspect.getsource( + LigerRMSNormForMuseGlimmer.forward + ) + for decoder_layer in text_model.layers: + assert inspect.getsource(decoder_layer.mlp.forward) == inspect.getsource(LigerSwiGLUMLP.forward) + assert inspect.getsource(decoder_layer.input_layernorm.forward) == inspect.getsource(LigerRMSNorm.forward) + assert inspect.getsource(decoder_layer.post_attention_layernorm.forward) == inspect.getsource( + LigerRMSNorm.forward + ) + assert inspect.getsource(decoder_layer.pre_feedforward_layernorm.forward) == inspect.getsource( + LigerRMSNorm.forward + ) + assert inspect.getsource(decoder_layer.post_feedforward_layernorm.forward) == inspect.getsource( + LigerRMSNorm.forward + ) + assert inspect.getsource(decoder_layer.self_attn.qk_norm.forward) == inspect.getsource( + LigerRMSNormForMuseGlimmer.forward + ) + # MuseGlimmerTextCenteredRMSNorm scales by (1 + weight) and the post-norms use + # a distinct, much smaller epsilon -- both must survive patching. + assert decoder_layer.input_layernorm.offset == 1.0 + assert decoder_layer.input_layernorm.casting_mode == "gemma" + assert decoder_layer.input_layernorm.in_place is False + assert decoder_layer.post_attention_layernorm.variance_epsilon == config.text_config.post_norm_eps + assert decoder_layer.input_layernorm.variance_epsilon == config.text_config.rms_norm_eps + # The scale-free QK norm has no weight to scale by + assert decoder_layer.self_attn.qk_norm.with_scale is False + # The final norm scales by weight directly (no +1 offset) + assert text_model.norm.offset == 0.0 + assert inspect.getsource(vision_model.ln_pre.forward) == inspect.getsource(LigerLayerNorm.forward) + assert inspect.getsource(vision_model.ln_post.forward) == inspect.getsource(LigerLayerNorm.forward) + for vision_layer in vision_model.layers: + assert inspect.getsource(vision_layer.norm1.forward) == inspect.getsource(LigerLayerNorm.forward) + assert inspect.getsource(vision_layer.norm2.forward) == inspect.getsource(LigerLayerNorm.forward) + + try: + print(dummy_model_instance) + except Exception as e: + pytest.fail(f"An exception occured in extra_expr: {type(e).__name__} - {e}") + + @pytest.mark.skipif(not is_qwen3_vl_available(), reason="qwen3_vl module not available") def test_apply_liger_kernel_to_instance_for_qwen3_vl_for_conditional_generation(): # Ensure any monkey patching is cleaned up for subsequent tests diff --git a/test/transformers/test_muse_glimmer.py b/test/transformers/test_muse_glimmer.py new file mode 100644 index 000000000..0efe67ead --- /dev/null +++ b/test/transformers/test_muse_glimmer.py @@ -0,0 +1,300 @@ +"""Numerical and contract tests for the Muse Glimmer Liger integration. + +These cover what `test_monkey_patch.py::test_apply_liger_kernel_to_instance_for_muse_glimmer` +and the convergence suite do not: + + * The convergence suite patches at *module* level (it calls the apply fn without + `model=`). TRL and `HFTrainer(use_liger_kernel=True)` go through + `_apply_liger_kernel_to_instance(model=...)` instead, which runs a different code + path -- three bespoke RMSNorm helpers plus the vision LayerNorm patch. Nothing + validated that path numerically. + * `return_token_accuracy=True -> outputs.token_accuracy` is the subject of the feature + request these patches implement, and had no model-level coverage at all. + +Requires a GPU: the Liger modules dispatch to Triton kernels. +""" + +import copy + +import pytest +import torch + +from test.utils import assert_verbose_allclose +from test.utils import supports_bfloat16 + +from liger_kernel.transformers.monkey_patch import _apply_liger_kernel_to_instance + + +def is_muse_glimmer_available(): + try: + import transformers.models.muse_glimmer # noqa: F401 + + return True + except ImportError: + return False + + +pytestmark = [ + pytest.mark.skipif(not is_muse_glimmer_available(), reason="muse_glimmer module not available"), + pytest.mark.skipif(not torch.cuda.is_available(), reason="Liger kernels require a GPU"), +] + + +def _mini_config(): + from transformers.models.muse_glimmer.configuration_muse_glimmer import MuseGlimmerConfig + from transformers.models.muse_glimmer.configuration_muse_glimmer import MuseGlimmerTextConfig + from transformers.models.muse_glimmer.configuration_muse_glimmer import MuseGlimmerVisionConfig + + # Text parameters must go in `text_config`; top-level fields are ignored. + config = MuseGlimmerConfig( + text_config=MuseGlimmerTextConfig( + vocab_size=512, + hidden_size=64, + intermediate_size=128, + num_hidden_layers=2, + num_attention_heads=4, + num_key_value_heads=2, + head_dim=16, + hidden_activation="silu", + max_position_embeddings=128, + sliding_window=16, + rms_norm_eps=1e-5, + post_norm_eps=1e-8, + attention_dropout=0.0, + ), + vision_config=MuseGlimmerVisionConfig( + hidden_size=32, + intermediate_size=64, + num_hidden_layers=1, + num_attention_heads=2, + patch_size=14, + pos_emb_height=4, + pos_emb_width=4, + max_position_embeddings=16, + merge_size=2, + ), + ) + config._attn_implementation = "sdpa" + return config + + +def _build_pair(dtype): + """An unpatched model and an instance-patched deepcopy of it, sharing weights.""" + from transformers.models.muse_glimmer.modeling_muse_glimmer import MuseGlimmerForConditionalGeneration + + torch.manual_seed(42) + model_hf = MuseGlimmerForConditionalGeneration(_mini_config()).to("cuda").to(dtype) + model_liger = copy.deepcopy(model_hf) + _apply_liger_kernel_to_instance(model=model_liger) + return model_hf, model_liger + + +def _assert_instance_patch_applied(model_liger): + """Guard against a silently no-op patch, which would make every parity assert vacuous.""" + from liger_kernel.transformers.rms_norm import LigerRMSNorm + from liger_kernel.transformers.swiglu import LigerSwiGLUMLPForMuseGlimmer + + text_model = model_liger.model.language_model + assert text_model.norm._get_name() == LigerRMSNorm.__name__ + for layer in text_model.layers: + assert layer.mlp._get_name() == LigerSwiGLUMLPForMuseGlimmer.__name__ + assert layer.input_layernorm._get_name() == LigerRMSNorm.__name__ + assert layer.input_layernorm.offset == 1.0 + assert not layer.input_layernorm.in_place + assert getattr(layer.self_attn.qk_norm, "weight", None) is None + + assert getattr(model_liger.model.perception_emb_norm, "weight", None) is None + embed_norm = text_model.embed_tokens.embed_norm + assert getattr(embed_norm, "weight", None) is None + + +def _text_batch(batch_size=2, seq_len=16, vocab_size=512): + input_ids = torch.randint(0, vocab_size, (batch_size, seq_len), device="cuda") + labels = input_ids.clone() + labels[:, :2] = -100 # exercise the ignore_index path + return input_ids, labels + + +def _expected_token_accuracy(logits, labels): + """Reference accuracy over the same shifted, non-ignored positions Liger uses.""" + shift_logits = logits[..., :-1, :] + shift_labels = labels[..., 1:] + mask = shift_labels != -100 + correct = shift_logits.argmax(dim=-1)[mask] == shift_labels[mask] + return correct.float().mean() + + +# --------------------------------------------------------------------------- +# Gap A -- instance-patch numerical parity (the path TRL actually takes) +# --------------------------------------------------------------------------- +@pytest.mark.parametrize( + "dtype, atol, rtol", + [ + (torch.float32, 1e-4, 1e-4), + pytest.param( + torch.bfloat16, + 1e-2, + 1e-2, + marks=pytest.mark.skipif(not supports_bfloat16(), reason="bfloat16 not supported on this GPU"), + ), + ], +) +def test_muse_glimmer_instance_patch_numerical_parity(dtype, atol, rtol): + model_hf, model_liger = _build_pair(dtype) + _assert_instance_patch_applied(model_liger) + + model_hf.train() + model_liger.train() + + input_ids, labels = _text_batch() + + out_hf = model_hf(input_ids=input_ids, labels=labels) + out_liger = model_liger(input_ids=input_ids, labels=labels) + + # Training + labels means the patched forward should default to skip_logits. + assert out_hf.logits is not None + assert out_liger.logits is None + + assert_verbose_allclose(out_hf.loss, out_liger.loss, atol=atol, rtol=rtol, extra_info="[Loss]") + + out_hf.loss.backward() + out_liger.loss.backward() + + # lm_head.weight is tied to embed_tokens.weight, so this also covers the fused + # linear cross entropy weight gradient. + assert_verbose_allclose( + model_hf.get_input_embeddings().weight.grad, + model_liger.get_input_embeddings().weight.grad, + atol=atol, + rtol=rtol, + extra_info="[Embedding grad]", + ) + + hf_layer = model_hf.model.language_model.layers[0] + liger_layer = model_liger.model.language_model.layers[0] + assert_verbose_allclose( + hf_layer.mlp.down_proj.weight.grad, + liger_layer.mlp.down_proj.weight.grad, + atol=atol, + rtol=rtol, + extra_info="[SwiGLU down_proj grad]", + ) + assert_verbose_allclose( + hf_layer.input_layernorm.weight.grad, + liger_layer.input_layernorm.weight.grad, + atol=atol, + rtol=rtol, + extra_info="[TextCentered RMSNorm grad]", + ) + + +def test_muse_glimmer_instance_patch_parity_without_skip_logits(): + """The non-fused branch reapplies output_multiplier and the tanh softcap by hand.""" + dtype, atol, rtol = torch.float32, 1e-4, 1e-4 + model_hf, model_liger = _build_pair(dtype) + _assert_instance_patch_applied(model_liger) + + model_hf.eval() + model_liger.eval() + + input_ids, labels = _text_batch() + + with torch.no_grad(): + out_hf = model_hf(input_ids=input_ids, labels=labels) + out_liger = model_liger(input_ids=input_ids, labels=labels, skip_logits=False) + + assert out_liger.logits is not None + assert_verbose_allclose(out_hf.logits, out_liger.logits, atol=atol, rtol=rtol, extra_info="[Logits]") + assert_verbose_allclose(out_hf.loss, out_liger.loss, atol=atol, rtol=rtol, extra_info="[Loss]") + + +def test_muse_glimmer_vision_layer_norm_patch_parity(): + """Vision LayerNorm is only patched on the instance path, so nothing else covers it.""" + from liger_kernel.transformers.layer_norm import LigerLayerNorm + + dtype, atol, rtol = torch.float32, 1e-4, 1e-4 + model_hf, model_liger = _build_pair(dtype) + + vision_hf = model_hf.model.vision_tower + vision_liger = model_liger.model.vision_tower + assert vision_liger.ln_pre._get_name() == LigerLayerNorm.__name__ + assert vision_liger.layers[0].norm1._get_name() == LigerLayerNorm.__name__ + + x = torch.randn(4, vision_hf.config.hidden_size, device="cuda", dtype=dtype) + for name, mod_hf, mod_liger in [ + ("ln_pre", vision_hf.ln_pre, vision_liger.ln_pre), + ("ln_post", vision_hf.ln_post, vision_liger.ln_post), + ("norm1", vision_hf.layers[0].norm1, vision_liger.layers[0].norm1), + ("norm2", vision_hf.layers[0].norm2, vision_liger.layers[0].norm2), + ]: + assert_verbose_allclose(mod_hf(x), mod_liger(x), atol=atol, rtol=rtol, extra_info=f"[vision {name}]") + + +# --------------------------------------------------------------------------- +# Gap B -- the token_accuracy contract TRL relies on +# --------------------------------------------------------------------------- +def test_muse_glimmer_returns_token_accuracy(): + model_hf, model_liger = _build_pair(torch.float32) + model_hf.train() + model_liger.train() + + input_ids, labels = _text_batch() + + out_liger = model_liger(input_ids=input_ids, labels=labels, return_token_accuracy=True) + + assert out_liger.token_accuracy is not None, ( + "return_token_accuracy=True must populate outputs.token_accuracy -- this is what " + "TRL reads, and it warns and drops the metric when it is missing." + ) + assert out_liger.logits is None, "token accuracy must not come at the cost of materializing logits" + assert 0.0 <= float(out_liger.token_accuracy) <= 1.0 + + # output_multiplier > 0 and tanh are monotonic, so argmax over HF's softcapped + # logits matches the argmax Liger takes over the raw ones. + with torch.no_grad(): + out_hf = model_hf(input_ids=input_ids, labels=labels) + expected = _expected_token_accuracy(out_hf.logits, labels) + assert_verbose_allclose( + expected, out_liger.token_accuracy.float(), atol=1e-4, rtol=1e-4, extra_info="[token_accuracy]" + ) + + +def test_muse_glimmer_returns_predicted_tokens(): + model_hf, model_liger = _build_pair(torch.float32) + model_hf.train() + model_liger.train() + + input_ids, labels = _text_batch() + batch_size, seq_len = input_ids.shape + + out_liger = model_liger(input_ids=input_ids, labels=labels, return_predicted_tokens=True) + + assert out_liger.predicted_tokens is not None + assert out_liger.predicted_tokens.dtype == torch.int64 + assert out_liger.predicted_tokens.numel() == batch_size * seq_len + + with torch.no_grad(): + out_hf = model_hf(input_ids=input_ids, labels=labels) + + # Liger pads labels then shifts, so position i predicts labels[i+1] and the final + # position is always ignored; ignored positions are left as -1. + predicted = out_liger.predicted_tokens.view(batch_size, seq_len) + shift_labels = torch.nn.functional.pad(labels, (0, 1), value=-100)[..., 1:] + mask = shift_labels != -100 + expected = out_hf.logits.argmax(dim=-1) + assert torch.equal(predicted[mask], expected[mask]) + assert (predicted[~mask] == -1).all() + + +def test_muse_glimmer_token_accuracy_absent_when_not_requested(): + _, model_liger = _build_pair(torch.float32) + model_liger.train() + + input_ids, labels = _text_batch() + out = model_liger(input_ids=input_ids, labels=labels) + + # The field must exist on the output class (TRL uses getattr), but stay None so the + # accuracy kernel is not paid for when nobody asked for it. + assert hasattr(out, "token_accuracy") + assert out.token_accuracy is None + assert out.predicted_tokens is None diff --git a/test/utils.py b/test/utils.py index 52142ac71..50cd4ec77 100644 --- a/test/utils.py +++ b/test/utils.py @@ -450,6 +450,18 @@ def revert_liger_kernel_to_mixtral(model_config: MiniModelConfig): print("Liger kernel patches have been reverted.") +def revert_liger_kernel_to_muse_glimmer(model_config: MiniModelConfig): + """ + Revert all Liger kernel patches applied to MuseGlimmer. + """ + + from transformers.models.muse_glimmer import modeling_muse_glimmer + + importlib.reload(modeling_muse_glimmer) + model_config.model_class = modeling_muse_glimmer.MuseGlimmerForConditionalGeneration + print("Liger kernel patches have been reverted.") + + def revert_liger_kernel_to_gemma(model_config: MiniModelConfig): """ Revert all Liger kernel patches applied to Gemma.