From ad5358d083237ed5770edee9f80c1748a08b734a Mon Sep 17 00:00:00 2001 From: zRzRzRzRzRzRzR <2448370773@qq.com> Date: Wed, 28 Jan 2026 15:21:30 +0800 Subject: [PATCH 01/19] draft --- docs/source/en/_toctree.yml | 2 + docs/source/en/model_doc/glm_moe_dsa.md | 59 ++ src/transformers/models/__init__.py | 1 + .../models/auto/configuration_auto.py | 2 + src/transformers/models/auto/modeling_auto.py | 2 + .../models/glm_moe_dsa/__init__.py | 28 + .../glm_moe_dsa/configuration_glm_moe_dsa.py | 252 ++++++ .../glm_moe_dsa/modeling_glm_moe_dsa.py | 789 ++++++++++++++++++ .../models/glm_moe_dsa/modular_glm_moe_dsa.py | 392 +++++++++ tests/models/glm_moe_dsa/__init__.py | 0 .../glm_moe_dsa/test_modeling_glm_moe_dsa.py | 112 +++ 11 files changed, 1639 insertions(+) create mode 100644 docs/source/en/model_doc/glm_moe_dsa.md create mode 100644 src/transformers/models/glm_moe_dsa/__init__.py create mode 100644 src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py create mode 100644 src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py create mode 100644 src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py create mode 100644 tests/models/glm_moe_dsa/__init__.py create mode 100644 tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py diff --git a/docs/source/en/_toctree.yml b/docs/source/en/_toctree.yml index 27778e780f91..c4181d2ddea1 100644 --- a/docs/source/en/_toctree.yml +++ b/docs/source/en/_toctree.yml @@ -545,6 +545,8 @@ title: GLM-4.7-Flash - local: model_doc/glm_image title: GLM-Image + - local: model_doc/glm_moe_dsa + title: GlmMoeDsa - local: model_doc/openai-gpt title: GPT - local: model_doc/gpt_neo diff --git a/docs/source/en/model_doc/glm_moe_dsa.md b/docs/source/en/model_doc/glm_moe_dsa.md new file mode 100644 index 000000000000..53682e45610f --- /dev/null +++ b/docs/source/en/model_doc/glm_moe_dsa.md @@ -0,0 +1,59 @@ + + + +# GlmMoeDsa + +## Overview + +The GlmMoeDsa model was proposed in []() by . + + +The abstract from the paper is the following: + + + +Tips: + + + +This model was contributed by [INSERT YOUR HF USERNAME HERE](https://huggingface.co/). +The original code can be found [here](). + +## Usage examples + + + +## GlmMoeDsaConfig + +[[autodoc]] GlmMoeDsaConfig + +## GlmMoeDsaPreTrainedModel + +[[autodoc]] GlmMoeDsaPreTrainedModel + - forward + +## GlmMoeDsaModel + +[[autodoc]] GlmMoeDsaModel + - forward + +## GlmMoeDsaForCausalLM + +[[autodoc]] GlmMoeDsaForCausalLM \ No newline at end of file diff --git a/src/transformers/models/__init__.py b/src/transformers/models/__init__.py index c6bef6db17d0..c8d6bfd94685 100644 --- a/src/transformers/models/__init__.py +++ b/src/transformers/models/__init__.py @@ -156,6 +156,7 @@ from .glm4v_moe import * from .glm46v import * from .glm_image import * + from .glm_moe_dsa import * from .glm_ocr import * from .glmasr import * from .glpn import * diff --git a/src/transformers/models/auto/configuration_auto.py b/src/transformers/models/auto/configuration_auto.py index 938450e5a1b2..e46b4d7f1ea0 100644 --- a/src/transformers/models/auto/configuration_auto.py +++ b/src/transformers/models/auto/configuration_auto.py @@ -184,6 +184,7 @@ ("glm_image_text", "GlmImageTextConfig"), ("glm_image_vision", "GlmImageVisionConfig"), ("glm_image_vqmodel", "GlmImageVQVAEConfig"), + ("glm_moe_dsa", "GlmMoeDsaConfig"), ("glm_ocr", "GlmOcrConfig"), ("glm_ocr_text", "GlmOcrTextConfig"), ("glm_ocr_vision", "GlmOcrVisionConfig"), @@ -651,6 +652,7 @@ ("glm_image_text", "GlmImageText"), ("glm_image_vision", "GlmImageVisionModel"), ("glm_image_vqmodel", "GlmImageVQVAE"), + ("glm_moe_dsa", "GlmMoeDsa"), ("glm_ocr", "Glmocr"), ("glm_ocr_text", "GlmOcrText"), ("glm_ocr_vision", "GlmOcrVisionModel"), diff --git a/src/transformers/models/auto/modeling_auto.py b/src/transformers/models/auto/modeling_auto.py index 9e31732bebea..fcf8793dab16 100644 --- a/src/transformers/models/auto/modeling_auto.py +++ b/src/transformers/models/auto/modeling_auto.py @@ -186,6 +186,7 @@ class _BaseModelWithGenerate(PreTrainedModel, GenerationMixin): ("glm_image_text", "GlmImageTextModel"), ("glm_image_vision", "GlmImageVisionModel"), ("glm_image_vqmodel", "GlmImageVQVAE"), + ("glm_moe_dsa", "GlmMoeDsaModel"), ("glm_ocr", "GlmOcrModel"), ("glm_ocr_text", "GlmOcrTextModel"), ("glm_ocr_vision", "GlmOcrVisionModel"), @@ -610,6 +611,7 @@ class _BaseModelWithGenerate(PreTrainedModel, GenerationMixin): ("glm4", "Glm4ForCausalLM"), ("glm4_moe", "Glm4MoeForCausalLM"), ("glm4_moe_lite", "Glm4MoeLiteForCausalLM"), + ("glm_moe_dsa", "GlmMoeDsaForCausalLM"), ("got_ocr2", "GotOcr2ForConditionalGeneration"), ("gpt-sw3", "GPT2LMHeadModel"), ("gpt2", "GPT2LMHeadModel"), diff --git a/src/transformers/models/glm_moe_dsa/__init__.py b/src/transformers/models/glm_moe_dsa/__init__.py new file mode 100644 index 000000000000..d874c6c21832 --- /dev/null +++ b/src/transformers/models/glm_moe_dsa/__init__.py @@ -0,0 +1,28 @@ +# Copyright 2026 the HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from typing import TYPE_CHECKING + +from ...utils import _LazyModule +from ...utils.import_utils import define_import_structure + + +if TYPE_CHECKING: + from .configuration_glm_moe_dsa import * + from .modeling_glm_moe_dsa import * +else: + import sys + + _file = globals()["__file__"] + sys.modules[__name__] = _LazyModule(__name__, _file, define_import_structure(_file), module_spec=__spec__) diff --git a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py new file mode 100644 index 000000000000..143ae4d7ea8d --- /dev/null +++ b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py @@ -0,0 +1,252 @@ +# 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨 +# This file was automatically generated from src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py. +# Do NOT edit this file manually as any edits will be overwritten by the generation of +# the file from the modular. If any change should be done, please apply the change to the +# modular_glm_moe_dsa.py file directly. One of our CI enforces this. +# 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨 +# Copyright 2026 the HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +from ...configuration_utils import PreTrainedConfig, layer_type_validation +from ...modeling_rope_utils import RopeParameters + + +class GlmMoeDsaConfig(PreTrainedConfig): + r""" + This is the configuration class to store the configuration of a [`GlmMoeDsaModel`]. It is used to instantiate an DeepSeek + model according to the specified arguments, defining the model architecture. Instantiating a configuration with the + defaults will yield a similar configuration to that of the DeepSeek-V3. + e.g. [bzantium/tiny-deepseek-v3](https://huggingface.co/bzantium/tiny-deepseek-v3) + Configuration objects inherit from [`PreTrainedConfig`] and can be used to control the model outputs. Read the + documentation from [`PreTrainedConfig`] for more information. + + + Args: + vocab_size (`int`, *optional*, defaults to 154880): + Vocabulary size of the Deep model. Defines the number of different tokens that can be represented by the + `inputs_ids` passed when calling [`Glm4MoeLiteModel`] + hidden_size (`int`, *optional*, defaults to 6144): + Dimension of the hidden representations. + intermediate_size (`int`, *optional*, defaults to 12288): + Dimension of the MLP representations. + moe_intermediate_size (`int`, *optional*, defaults to 2048): + Dimension of the MoE representations. + num_hidden_layers (`int`, *optional*, defaults to 78): + Number of hidden layers in the Transformer decoder. + num_attention_heads (`int`, *optional*, defaults to 64): + Number of attention heads for each attention layer in the Transformer decoder. + num_key_value_heads (`int`, *optional*, defaults to 64): + This is the number of key_value heads that should be used to implement Grouped Query Attention. If + `num_key_value_heads=num_attention_heads`, the model will use Multi Head Attention (MHA), if + `num_key_value_heads=1 the model will use Multi Query Attention (MQA) otherwise GQA is used. When + converting a multi-head checkpoint to a GQA checkpoint, each group key and value head should be constructed + by meanpooling all the original heads within that group. For more details, check out [this + paper](https://huggingface.co/papers/2305.13245). If it is not specified, will default to + `num_attention_heads`. + n_shared_experts (`int`, *optional*, defaults to 1): + Number of shared experts. + n_routed_experts (`int`, *optional*, defaults to 256): + Number of routed experts. + routed_scaling_factor (`float`, *optional*, defaults to 2.5): + Scaling factor or routed experts. + kv_lora_rank (`int`, *optional*, defaults to 512): + Rank of the LoRA matrices for key and value projections. + q_lora_rank (`int`, *optional*, defaults to 2048): + Rank of the LoRA matrices for query projections. + qk_rope_head_dim (`int`, *optional*, defaults to 64): + Dimension of the query/key heads that use rotary position embeddings. + v_head_dim (`int`, *optional*, defaults to 256): + Dimension of the value heads. + qk_nope_head_dim (`int`, *optional*, defaults to 192): + Dimension of the query/key heads that don't use rotary position embeddings. + n_group (`int`, *optional*, defaults to 1): + Number of groups for routed experts. + topk_group (`int`, *optional*, defaults to 1): + Number of selected groups for each token(for each token, ensuring the selected experts is only within `topk_group` groups). + num_experts_per_tok (`int`, *optional*, defaults to 8): + Number of selected experts, None means dense model. + norm_topk_prob (`bool`, *optional*, defaults to `True`): + Whether to normalize the weights of the routed experts. + hidden_act (`str` or `function`, *optional*, defaults to `"silu"`): + The non-linear activation function (function or string) in the decoder. + max_position_embeddings (`int`, *optional*, defaults to 202752): + The maximum sequence length that this model might ever be used with. + initializer_range (`float`, *optional*, defaults to 0.02): + The standard deviation of the truncated_normal_initializer for initializing all weight matrices. + rms_norm_eps (`float`, *optional*, defaults to 1e-05): + The epsilon used by the rms normalization layers. + use_cache (`bool`, *optional*, defaults to `True`): + Whether or not the model should return the last key/values attentions (not used by all models). Only + relevant if `config.is_decoder=True`. + pad_token_id (`int`, *optional*): + Padding token id. + bos_token_id (`int`, *optional*, defaults to 0): + Beginning of stream token id. + eos_token_id (`int`, *optional*, defaults to 1): + End of stream token id. + pretraining_tp (`int`, *optional*, defaults to 1): + Experimental feature. Tensor parallelism rank used during pretraining. Please refer to [this + document](https://huggingface.co/docs/transformers/parallelism) to understand more about it. This value is + necessary to ensure exact reproducibility of the pretraining results. Please refer to [this + issue](https://github.com/pytorch/pytorch/issues/76232). + tie_word_embeddings (`bool`, *optional*, defaults to `False`): + Whether to tie weight embeddings + rope_parameters (`RopeParameters`, *optional*): + Dictionary containing the configuration parameters for the RoPE embeddings. The dictionary should contain + a value for `rope_theta` and optionally parameters used for scaling in case you want to use RoPE + with longer `max_position_embeddings`. + rope_interleave (`bool`, *optional*, defaults to `True`): + Whether to interleave the rotary position embeddings. + mlp_layer_types (`list`, *optional*): + MLP (Moe vs Dense) pattern for each layer. + attention_bias (`bool`, defaults to `False`, *optional*, defaults to `False`): + Whether to use a bias in the query, key, value and output projection layers during self-attention. + attention_dropout (`float`, *optional*, defaults to 0.0): + The dropout ratio for the attention probabilities. + + ```python + >>> from transformers import Glm4MoeLiteModel, Glm4MoeLiteConfig + + >>> # Initializing a GLM-MOE-DSA style configuration + >>> configuration = GlmMoeDsaConfig() + + >>> # Accessing the model configuration + >>> configuration = model.config + ```""" + + model_type = "glm_moe_dsa" + keys_to_ignore_at_inference = ["past_key_values"] + base_model_tp_plan = { + "layers.*.self_attn.o_proj": "rowwise", + "layers.*.mlp.experts.gate_up_proj": "local_rowwise", + "layers.*.mlp.experts.down_proj": "local_rowwise", + "layers.*.mlp.experts": "gather", + "layers.*.mlp.gate_proj": "colwise", + "layers.*.mlp.up_proj": "colwise", + "layers.*.mlp.down_proj": "rowwise", + } + base_model_pp_plan = { + "embed_tokens": (["input_ids"], ["inputs_embeds"]), + "layers": (["hidden_states", "attention_mask"], ["hidden_states"]), + "norm": (["hidden_states"], ["hidden_states"]), + } + attribute_map = { + "num_local_experts": "n_routed_experts", + } + + def __init__( + self, + vocab_size: int | None = 154880, + hidden_size: int | None = 6144, + intermediate_size: int | None = 12288, + moe_intermediate_size: int | None = 2048, + num_hidden_layers: int | None = 78, + num_attention_heads: int | None = 64, + num_key_value_heads: int | None = 64, + n_shared_experts: int | None = 1, + n_routed_experts: int | None = 256, + routed_scaling_factor: float | None = 2.5, + kv_lora_rank: int | None = 512, + q_lora_rank: int | None = 2048, + qk_rope_head_dim: int | None = 64, + v_head_dim: int | None = 256, + qk_nope_head_dim: int | None = 192, + n_group: int | None = 1, + topk_group: int | None = 1, + num_experts_per_tok: int | None = 8, + norm_topk_prob: bool | None = True, + hidden_act: str | None = "silu", + max_position_embeddings: int | None = 202752, + initializer_range: float | None = 0.02, + rms_norm_eps: int | None = 1e-5, + use_cache: bool | None = True, + pad_token_id: int | None = None, + bos_token_id: int | None = 0, + eos_token_id: int | None = 1, + pretraining_tp: int | None = 1, + tie_word_embeddings: bool | None = False, + rope_parameters: RopeParameters | dict[str, RopeParameters] | None = None, + rope_interleave: bool | None = True, + mlp_layer_types=None, + attention_bias: bool | None = False, + attention_dropout: float | None = 0.0, + index_topk: int | None = 2048, + **kwargs, + ): + self.hidden_size = hidden_size + self.intermediate_size = intermediate_size + self.num_hidden_layers = num_hidden_layers + self.moe_intermediate_size = moe_intermediate_size + self.num_attention_heads = num_attention_heads + self.n_shared_experts = n_shared_experts + self.n_routed_experts = n_routed_experts + self.routed_scaling_factor = routed_scaling_factor + self.kv_lora_rank = kv_lora_rank + self.q_lora_rank = q_lora_rank + self.qk_rope_head_dim = qk_rope_head_dim + self.v_head_dim = v_head_dim + self.qk_nope_head_dim = qk_nope_head_dim + self.qk_head_dim = qk_nope_head_dim + qk_rope_head_dim + self.head_dim = qk_rope_head_dim + self.num_experts_per_tok = num_experts_per_tok + self.num_key_value_heads = num_key_value_heads + self.initializer_range = initializer_range + self.index_topk = index_topk + self.vocab_size = vocab_size + self.max_position_embeddings = max_position_embeddings + self.hidden_size = hidden_size + self.intermediate_size = intermediate_size + self.num_hidden_layers = num_hidden_layers + + # Default to MoE from the second layer and on + self.mlp_layer_types = mlp_layer_types + if self.mlp_layer_types is None: + self.mlp_layer_types = ["dense"] + ["sparse"] * (self.num_hidden_layers - 1) + layer_type_validation(self.mlp_layer_types, self.num_hidden_layers, attention=False) + + self.moe_intermediate_size = moe_intermediate_size + self.num_attention_heads = num_attention_heads + self.n_shared_experts = n_shared_experts + self.n_routed_experts = n_routed_experts + self.routed_scaling_factor = routed_scaling_factor + self.kv_lora_rank = kv_lora_rank + self.q_lora_rank = q_lora_rank + self.qk_rope_head_dim = qk_rope_head_dim + self.v_head_dim = v_head_dim + self.qk_nope_head_dim = qk_nope_head_dim + self.qk_head_dim = qk_nope_head_dim + qk_rope_head_dim + self.head_dim = qk_rope_head_dim + self.n_group = n_group + self.topk_group = topk_group + self.num_experts_per_tok = num_experts_per_tok + self.norm_topk_prob = norm_topk_prob + self.rope_interleave = rope_interleave + self.num_key_value_heads = num_key_value_heads + self.hidden_act = hidden_act + self.initializer_range = initializer_range + self.rms_norm_eps = rms_norm_eps + self.pretraining_tp = pretraining_tp + self.use_cache = use_cache + self.attention_bias = attention_bias + self.attention_dropout = attention_dropout + self.rope_parameters = rope_parameters + self.pad_token_id = pad_token_id + self.bos_token_id = bos_token_id + self.eos_token_id = eos_token_id + self.tie_word_embeddings = tie_word_embeddings + + super().__init__(**kwargs) + + +__all__ = ["GlmMoeDsaConfig"] diff --git a/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py new file mode 100644 index 000000000000..ff360a75c0c4 --- /dev/null +++ b/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py @@ -0,0 +1,789 @@ +# 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨 +# This file was automatically generated from src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py. +# Do NOT edit this file manually as any edits will be overwritten by the generation of +# the file from the modular. If any change should be done, please apply the change to the +# modular_glm_moe_dsa.py file directly. One of our CI enforces this. +# 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨 +# Copyright 2026 the HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import math +import warnings +from collections.abc import Callable +from typing import Optional + +import torch +import torch.nn.functional as F +from torch import nn + +from ... import initialization as init +from ...activations import ACT2FN +from ...cache_utils import Cache, DynamicCache +from ...generation import GenerationMixin +from ...integrations import use_experts_implementation, use_kernel_forward_from_hub, use_kernel_func_from_hub +from ...masking_utils import create_causal_mask +from ...modeling_flash_attention_utils import FlashAttentionKwargs +from ...modeling_layers import GradientCheckpointingLayer +from ...modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast +from ...modeling_rope_utils import ROPE_INIT_FUNCTIONS, dynamic_rope_update +from ...modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel +from ...processing_utils import Unpack +from ...utils import TransformersKwargs, auto_docstring, can_return_tuple, is_grouped_mm_available +from ...utils.generic import check_model_inputs, maybe_autocast +from .configuration_glm_moe_dsa import GlmMoeDsaConfig + + +@use_kernel_forward_from_hub("RMSNorm") +class GlmMoeDsaRMSNorm(nn.Module): + def __init__(self, hidden_size, eps=1e-6): + """ + GlmMoeDsaRMSNorm is equivalent to T5LayerNorm + """ + super().__init__() + self.weight = nn.Parameter(torch.ones(hidden_size)) + self.variance_epsilon = eps + + def forward(self, hidden_states): + input_dtype = hidden_states.dtype + hidden_states = hidden_states.to(torch.float32) + variance = hidden_states.pow(2).mean(-1, keepdim=True) + hidden_states = hidden_states * torch.rsqrt(variance + self.variance_epsilon) + return self.weight * hidden_states.to(input_dtype) + + def extra_repr(self): + return f"{tuple(self.weight.shape)}, eps={self.variance_epsilon}" + + +def rotate_half(x): + """Rotates half the hidden dims of the input.""" + x1 = x[..., : x.shape[-1] // 2] + x2 = x[..., x.shape[-1] // 2 :] + return torch.cat((-x2, x1), dim=-1) + + +@use_kernel_func_from_hub("rotary_pos_emb") +def apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=1): + """Applies Rotary Position Embedding to the query and key tensors. + + Args: + q (`torch.Tensor`): The query tensor. + k (`torch.Tensor`): The key tensor. + cos (`torch.Tensor`): The cosine part of the rotary embedding. + sin (`torch.Tensor`): The sine part of the rotary embedding. + unsqueeze_dim (`int`, *optional*, defaults to 1): + The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and + sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note + that cos[position_ids] and sin[position_ids] have the shape [batch_size, seq_len, head_dim]. Then, if q and + k have the shape [batch_size, heads, seq_len, head_dim], then setting unsqueeze_dim=1 makes + cos[position_ids] and sin[position_ids] broadcastable to the shapes of q and k. Similarly, if q and k have + the shape [batch_size, seq_len, heads, head_dim], then set unsqueeze_dim=2. + Returns: + `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding. + """ + cos = cos.unsqueeze(unsqueeze_dim) + sin = sin.unsqueeze(unsqueeze_dim) + q_embed = (q * cos) + (rotate_half(q) * sin) + k_embed = (k * cos) + (rotate_half(k) * sin) + return q_embed, k_embed + + +def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor: + """ + This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch, + num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim) + """ + batch, num_key_value_heads, slen, head_dim = hidden_states.shape + if n_rep == 1: + return hidden_states + hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim) + return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim) + + +def eager_attention_forward( + module: nn.Module, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + attention_mask: torch.Tensor | None, + scaling: float, + dropout: float = 0.0, + **kwargs: Unpack[TransformersKwargs], +): + key_states = repeat_kv(key, module.num_key_value_groups) + value_states = repeat_kv(value, module.num_key_value_groups) + + attn_weights = torch.matmul(query, key_states.transpose(2, 3)) * scaling + if attention_mask is not None: + causal_mask = attention_mask[:, :, :, : key_states.shape[-2]] + attn_weights = attn_weights + causal_mask + + attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query.dtype) + attn_weights = nn.functional.dropout(attn_weights, p=dropout, training=module.training) + attn_output = torch.matmul(attn_weights, value_states) + attn_output = attn_output.transpose(1, 2).contiguous() + + return attn_output, attn_weights + + +def apply_rotary_pos_emb_interleave(q, k, cos, sin, position_ids=None, unsqueeze_dim=1): + r""" + TODO let's just use the original freqcis computation to not have the view + transpose + reshape! This is not optimized! + Applies Rotary Position Embedding to the query and key tensors. + + Args: + q (`torch.Tensor`): The query tensor. + k (`torch.Tensor`): The key tensor. + cos (`torch.Tensor`): The cosine part of the rotary embedding. + sin (`torch.Tensor`): The sine part of the rotary embedding. + position_ids (`torch.Tensor`): + The position indices of the tokens corresponding to the query and key tensors. For example, this can be + used to pass offsetted position ids when working with a KV-cache. + unsqueeze_dim (`int`, *optional*, defaults to 1): + The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and + sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note + that cos[position_ids] and sin[position_ids] have the shape [batch_size, seq_len, head_dim]. Then, if q and + k have the shape [batch_size, heads, seq_len, head_dim], then setting unsqueeze_dim=1 makes + cos[position_ids] and sin[position_ids] broadcastable to the shapes of q and k. Similarly, if q and k have + the shape [batch_size, seq_len, heads, head_dim], then set unsqueeze_dim=2. + Returns: + `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding. + """ + cos = cos.unsqueeze(unsqueeze_dim) + sin = sin.unsqueeze(unsqueeze_dim) + + b, h, s, d = q.shape + q = q.view(b, h, s, d // 2, 2).transpose(4, 3).reshape(b, h, s, d) + + b, h, s, d = k.shape + k = k.view(b, h, s, d // 2, 2).transpose(4, 3).reshape(b, h, s, d) + + q_embed = (q * cos) + (rotate_half(q) * sin) + k_embed = (k * cos) + (rotate_half(k) * sin) + return q_embed, k_embed + + +def yarn_get_mscale(scale=1, mscale=1): + if scale <= 1: + return 1.0 + return 0.1 * mscale * math.log(scale) + 1.0 + + +class GlmMoeDsaAttention(nn.Module): + """ + DeepSeek V3.2 sparse attention mechanism with indexer. + + This implements the native sparse attention from [DeepSeek V3.2](https://huggingface.co/deepseek-ai/DeepSeek-V3.2) which uses + an indexer to select top-k tokens for attention computation, making it more efficient for long sequences. + + Switch to the implementation from this [PR](https://github.com/huggingface/transformers/pull/41251) as soon as it’s merged. + """ + + def __init__(self, config: GlmMoeDsaConfig, layer_idx: int): + super().__init__() + self.config = config + self.layer_idx = layer_idx + self.num_key_value_groups = config.num_attention_heads // config.num_key_value_heads + self.attention_dropout = config.attention_dropout + self.num_heads = config.num_attention_heads + + self.q_lora_rank = config.q_lora_rank + self.qk_rope_head_dim = config.qk_rope_head_dim + self.kv_lora_rank = config.kv_lora_rank + self.v_head_dim = config.v_head_dim + self.qk_nope_head_dim = config.qk_nope_head_dim + self.qk_head_dim = config.qk_head_dim + self.index_topk = config.index_topk + + self.is_causal = True + + # Query projection + if self.q_lora_rank is None: + self.q_proj = nn.Linear(config.hidden_size, self.num_heads * self.qk_head_dim, bias=False) + else: + self.q_a_proj = nn.Linear(config.hidden_size, config.q_lora_rank, bias=config.attention_bias) + self.q_a_layernorm = GlmMoeDsaRMSNorm(config.q_lora_rank) + self.q_b_proj = nn.Linear(config.q_lora_rank, self.num_heads * self.qk_head_dim, bias=False) + + # Key-Value projections + self.kv_a_proj_with_mqa = nn.Linear( + config.hidden_size, + self.kv_lora_rank + self.qk_rope_head_dim, + bias=config.attention_bias, + ) + self.kv_a_layernorm = GlmMoeDsaRMSNorm(self.kv_lora_rank) + self.kv_b_proj = nn.Linear( + self.kv_lora_rank, + self.num_heads * (self.qk_nope_head_dim + self.v_head_dim), + bias=False, + ) + + # Output projection + self.o_proj = nn.Linear( + self.num_heads * self.v_head_dim, + config.hidden_size, + bias=config.attention_bias, + ) + + # Indexer components for sparse attention + self.wq_b = nn.Linear(config.q_lora_rank, self.num_heads * self.qk_head_dim, bias=False) + self.wk = nn.Linear(config.hidden_size, self.qk_head_dim, bias=config.attention_bias) + self.k_norm = GlmMoeDsaRMSNorm(self.qk_head_dim) + self.weights_proj = nn.Linear(config.hidden_size, self.num_heads, bias=False) + + self.scaling = self.qk_head_dim ** (-0.5) + if self.config.rope_parameters.get("rope_type", "default") != "default": + mscale_all_dim = self.config.rope_parameters.get("mscale_all_dim", 0) + scaling_factor = self.config.rope_parameters["factor"] + if mscale_all_dim: + mscale = yarn_get_mscale(scaling_factor, mscale_all_dim) + self.scaling = self.scaling * mscale * mscale + + def forward( + self, + hidden_states: torch.Tensor, + position_embeddings: tuple[torch.Tensor, torch.Tensor], + attention_mask: torch.Tensor | None, + past_key_values: Cache | None = None, + cache_position: torch.LongTensor | None = None, + **kwargs: Unpack[FlashAttentionKwargs], + ) -> tuple[torch.Tensor, torch.Tensor | None, tuple[torch.Tensor] | None]: + batch_size, seq_length = hidden_states.shape[:-1] + + # For training or when index_topk is not effective, fall back to standard attention + # This is a simplified implementation - in practice, you'd implement the full sparse indexer + if self.training or seq_length <= self.index_topk: + warnings.warn( + "DeepSeek V3.2 sparse attention is not fully implemented in this version. " + "Falling back to standard attention. For production use, please use vLLM or " + "other optimized inference engines.", + UserWarning, + ) + return self._standard_attention( + hidden_states, position_embeddings, attention_mask, past_key_values, cache_position, **kwargs + ) + + # Sparse attention implementation would go here + # This requires custom CUDA kernels for efficient top-k selection and indexing + return self._standard_attention( + hidden_states, position_embeddings, attention_mask, past_key_values, cache_position, **kwargs + ) + + def _standard_attention( + self, + hidden_states: torch.Tensor, + position_embeddings: tuple[torch.Tensor, torch.Tensor], + attention_mask: torch.Tensor | None, + past_key_values: Cache | None = None, + cache_position: torch.LongTensor | None = None, + **kwargs: Unpack[FlashAttentionKwargs], + ) -> tuple[torch.Tensor, torch.Tensor | None, tuple[torch.Tensor] | None]: + """Standard attention fallback (same as DeepSeek V3)""" + batch_size, seq_length = hidden_states.shape[:-1] + query_shape = (batch_size, seq_length, -1, self.qk_head_dim) + key_shape = (batch_size, seq_length, -1, self.qk_nope_head_dim + self.v_head_dim) + + if self.q_lora_rank is None: + q_states = self.q_proj(hidden_states) + else: + q_states = self.q_b_proj(self.q_a_layernorm(self.q_a_proj(hidden_states))) + q_states = q_states.view(query_shape).transpose(1, 2) + q_pass, q_rot = torch.split(q_states, [self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1) + + compressed_kv = self.kv_a_proj_with_mqa(hidden_states) + k_pass, k_rot = torch.split(compressed_kv, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1) + + k_pass = self.kv_b_proj(self.kv_a_layernorm(k_pass)).view(key_shape).transpose(1, 2) + k_pass, value_states = torch.split(k_pass, [self.qk_nope_head_dim, self.v_head_dim], dim=-1) + + k_rot = k_rot.view(batch_size, 1, seq_length, self.qk_rope_head_dim) + + cos, sin = position_embeddings + if self.config.rope_interleave: + q_rot, k_rot = apply_rotary_pos_emb_interleave(q_rot, k_rot, cos, sin) + else: + q_rot, k_rot = apply_rotary_pos_emb(q_rot, k_rot, cos, sin) + k_rot = k_rot.expand(*k_pass.shape[:-1], -1) + + query_states = torch.cat((q_pass, q_rot), dim=-1) + key_states = torch.cat((k_pass, k_rot), dim=-1) + + if past_key_values is not None: + cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position} + key_states, value_states = past_key_values.update(key_states, value_states, self.layer_idx, cache_kwargs) + + if self.config._attn_implementation == "flash_attention_2" and self.qk_head_dim != self.v_head_dim: + value_states = F.pad(value_states, [0, self.qk_head_dim - self.v_head_dim]) + + attention_interface: Callable = eager_attention_forward + if self.config._attn_implementation != "eager": + attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation] + + attn_output, attn_weights = attention_interface( + self, + query_states, + key_states, + value_states, + attention_mask, + dropout=0.0 if not self.training else self.attention_dropout, + scaling=self.scaling, + **kwargs, + ) + + if self.config._attn_implementation == "flash_attention_2" and self.qk_head_dim != self.v_head_dim: + attn_output = attn_output[:, :, :, : self.v_head_dim] + + attn_output = attn_output.reshape(batch_size, seq_length, -1).contiguous() + attn_output = self.o_proj(attn_output) + return attn_output, attn_weights + + +class GlmMoeDsaMLP(nn.Module): + def __init__(self, config, intermediate_size=None): + super().__init__() + self.config = config + self.hidden_size = config.hidden_size + self.intermediate_size = config.intermediate_size if intermediate_size is None else 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) + self.act_fn = ACT2FN[config.hidden_act] + + def forward(self, x): + down_proj = self.down_proj(self.act_fn(self.gate_proj(x)) * self.up_proj(x)) + return down_proj + + +class GlmMoeDsaTopkRouter(nn.Module): + def __init__(self, config: GlmMoeDsaConfig): + super().__init__() + self.config = config + self.top_k = config.num_experts_per_tok + self.n_routed_experts = config.n_routed_experts + self.routed_scaling_factor = config.routed_scaling_factor + self.n_group = config.n_group + self.topk_group = config.topk_group + self.norm_topk_prob = config.norm_topk_prob + + self.weight = nn.Parameter(torch.empty((self.n_routed_experts, config.hidden_size))) + self.register_buffer("e_score_correction_bias", torch.zeros((self.n_routed_experts), dtype=torch.float32)) + + def forward(self, hidden_states): + hidden_states = hidden_states.view(-1, self.config.hidden_size) + router_logits = F.linear(hidden_states.type(torch.float32), self.weight.type(torch.float32)) + return router_logits + + +@use_experts_implementation +class GlmMoeDsaNaiveMoe(nn.Module): + """Collection of expert weights stored as 3D tensors.""" + + def __init__(self, config): + super().__init__() + self.num_experts = config.num_local_experts + self.hidden_dim = config.hidden_size + self.intermediate_dim = config.moe_intermediate_size + self.gate_up_proj = nn.Parameter(torch.empty(self.num_experts, 2 * self.intermediate_dim, self.hidden_dim)) + self.down_proj = nn.Parameter(torch.empty(self.num_experts, self.hidden_dim, self.intermediate_dim)) + self.act_fn = ACT2FN[config.hidden_act] + + def forward( + self, + hidden_states: torch.Tensor, + top_k_index: torch.Tensor, + top_k_weights: torch.Tensor, + ) -> torch.Tensor: + final_hidden_states = torch.zeros_like(hidden_states) + with torch.no_grad(): + expert_mask = torch.nn.functional.one_hot(top_k_index, num_classes=self.num_experts) + expert_mask = expert_mask.permute(2, 1, 0) + expert_hit = torch.greater(expert_mask.sum(dim=(-1, -2)), 0).nonzero() + + for expert_idx in expert_hit: + expert_idx = expert_idx[0] + if expert_idx == self.num_experts: + continue + top_k_pos, token_idx = torch.where(expert_mask[expert_idx]) + current_state = hidden_states[token_idx] + gate, up = nn.functional.linear(current_state, self.gate_up_proj[expert_idx]).chunk(2, dim=-1) + current_hidden_states = self.act_fn(gate) * up + current_hidden_states = nn.functional.linear(current_hidden_states, self.down_proj[expert_idx]) + current_hidden_states = current_hidden_states * top_k_weights[token_idx, top_k_pos, None] + final_hidden_states.index_add_(0, token_idx, current_hidden_states.to(final_hidden_states.dtype)) + + return final_hidden_states + + +class GlmMoeDsaMoE(nn.Module): + """ + A mixed expert module containing shared experts. + """ + + def __init__(self, config): + super().__init__() + self.config = config + self.experts = GlmMoeDsaNaiveMoe(config) + self.gate = GlmMoeDsaTopkRouter(config) + self.shared_experts = GlmMoeDsaMLP( + config=config, intermediate_size=config.moe_intermediate_size * config.n_shared_experts + ) + self.n_routed_experts = config.n_routed_experts + self.n_group = config.n_group + self.topk_group = config.topk_group + self.norm_topk_prob = config.norm_topk_prob + self.routed_scaling_factor = config.routed_scaling_factor + self.top_k = config.num_experts_per_tok + + def route_tokens_to_experts(self, router_logits): + router_logits = router_logits.sigmoid() + router_logits_for_choice = router_logits + self.gate.e_score_correction_bias + group_scores = ( + router_logits_for_choice.view(-1, self.n_group, self.n_routed_experts // self.n_group) + .topk(2, dim=-1)[0] + .sum(dim=-1) + ) + group_idx = torch.topk(group_scores, k=self.topk_group, dim=-1, sorted=False)[1] + group_mask = torch.zeros_like(group_scores) + group_mask.scatter_(1, group_idx, 1) + score_mask = ( + group_mask.unsqueeze(-1) + .expand(-1, self.n_group, self.n_routed_experts // self.n_group) + .reshape(-1, self.n_routed_experts) + ) + scores_for_choice = router_logits_for_choice.masked_fill(~score_mask.bool(), 0.0) + topk_indices = torch.topk(scores_for_choice, k=self.top_k, dim=-1, sorted=False)[1] + topk_weights = router_logits.gather(1, topk_indices) + if self.norm_topk_prob: + denominator = topk_weights.sum(dim=-1, keepdim=True) + 1e-20 + topk_weights /= denominator + topk_weights = topk_weights * self.routed_scaling_factor + return topk_indices, topk_weights + + def forward(self, hidden_states): + residuals = hidden_states + orig_shape = hidden_states.shape + router_logits = self.gate(hidden_states) + topk_indices, topk_weights = self.route_tokens_to_experts(router_logits) + hidden_states = hidden_states.view(-1, hidden_states.shape[-1]) + hidden_states = self.experts(hidden_states, topk_indices, topk_weights).view(*orig_shape) + hidden_states = hidden_states + self.shared_experts(residuals) + return hidden_states + + +class GlmMoeDsaDecoderLayer(GradientCheckpointingLayer): + def __init__(self, config: GlmMoeDsaConfig, layer_idx: int): + super().__init__() + self.hidden_size = config.hidden_size + + self.self_attn = GlmMoeDsaAttention(config=config, layer_idx=layer_idx) + + if layer_idx >= config.first_k_dense_replace: + self.mlp = GlmMoeDsaMoE(config) + else: + self.mlp = GlmMoeDsaMLP(config) + + self.input_layernorm = GlmMoeDsaRMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.post_attention_layernorm = GlmMoeDsaRMSNorm(config.hidden_size, eps=config.rms_norm_eps) + + def forward( + self, + hidden_states: torch.Tensor, + attention_mask: torch.Tensor | None = None, + position_ids: torch.LongTensor | None = None, + past_key_values: Cache | None = None, + use_cache: bool | None = False, + cache_position: torch.LongTensor | None = None, + position_embeddings: tuple[torch.Tensor, torch.Tensor] | None = None, + **kwargs: Unpack[TransformersKwargs], + ) -> torch.Tensor: + residual = hidden_states + hidden_states = self.input_layernorm(hidden_states) + # Self Attention + hidden_states, _ = self.self_attn( + hidden_states=hidden_states, + attention_mask=attention_mask, + position_ids=position_ids, + past_key_values=past_key_values, + use_cache=use_cache, + cache_position=cache_position, + position_embeddings=position_embeddings, + **kwargs, + ) + hidden_states = residual + hidden_states + + # Fully Connected + residual = hidden_states + hidden_states = self.post_attention_layernorm(hidden_states) + hidden_states = self.mlp(hidden_states) + hidden_states = residual + hidden_states + return hidden_states + + +@auto_docstring +class GlmMoeDsaPreTrainedModel(PreTrainedModel): + config: GlmMoeDsaConfig + base_model_prefix = "model" + supports_gradient_checkpointing = True + _no_split_modules = ["GlmMoeDsaDecoderLayer"] + _skip_keys_device_placement = ["past_key_values"] + _supports_flash_attn = True + _supports_sdpa = True + _supports_flex_attn = True + _can_compile_fullgraph = ( + is_grouped_mm_available() + ) # https://huggingface.co/docs/transformers/experts_interface#torchcompile + _supports_attention_backend = True + _can_record_outputs = { + "hidden_states": GlmMoeDsaDecoderLayer, + "attentions": GlmMoeDsaAttention, + } + _keep_in_fp32_modules_strict = ["e_score_correction_bias"] + + @torch.no_grad() + def _init_weights(self, module): + super()._init_weights(module) + if isinstance(module, GlmMoeDsaTopkRouter): + init.normal_(module.weight, mean=0.0, std=self.config.initializer_range) + init.zeros_(module.e_score_correction_bias) + elif isinstance(module, GlmMoeDsaNaiveMoe): + init.normal_(module.gate_up_proj, mean=0.0, std=self.config.initializer_range) + init.normal_(module.down_proj, mean=0.0, std=self.config.initializer_range) + + +class GlmMoeDsaRotaryEmbedding(nn.Module): + inv_freq: torch.Tensor # fix linting for `register_buffer` + + def __init__(self, config: GlmMoeDsaConfig, device=None): + super().__init__() + self.max_seq_len_cached = config.max_position_embeddings + self.original_max_seq_len = config.max_position_embeddings + + self.config = config + + self.rope_type = self.config.rope_parameters["rope_type"] + rope_init_fn: Callable = self.compute_default_rope_parameters + if self.rope_type != "default": + rope_init_fn = ROPE_INIT_FUNCTIONS[self.rope_type] + inv_freq, self.attention_scaling = rope_init_fn(self.config, device) + + self.register_buffer("inv_freq", inv_freq, persistent=False) + self.register_buffer("original_inv_freq", inv_freq.clone(), persistent=False) + + @staticmethod + def compute_default_rope_parameters( + config: GlmMoeDsaConfig | None = None, + device: Optional["torch.device"] = None, + seq_len: int | None = None, + ) -> tuple["torch.Tensor", float]: + """ + Computes the inverse frequencies according to the original RoPE implementation + Args: + config ([`~transformers.PreTrainedConfig`]): + The model configuration. + device (`torch.device`): + The device to use for initialization of the inverse frequencies. + seq_len (`int`, *optional*): + The current sequence length. Unused for this type of RoPE. + Returns: + Tuple of (`torch.Tensor`, `float`), containing the inverse frequencies for the RoPE embeddings and the + post-processing scaling factor applied to the computed cos/sin (unused in this type of RoPE). + """ + base = config.rope_parameters["rope_theta"] + partial_rotary_factor = config.rope_parameters.get("partial_rotary_factor", 1.0) + head_dim = getattr(config, "head_dim", None) or config.hidden_size // config.num_attention_heads + dim = int(head_dim * partial_rotary_factor) + + attention_factor = 1.0 # Unused in this type of RoPE + + # Compute the inverse frequencies + inv_freq = 1.0 / ( + base ** (torch.arange(0, dim, 2, dtype=torch.int64).to(device=device, dtype=torch.float) / dim) + ) + return inv_freq, attention_factor + + @torch.no_grad() + @dynamic_rope_update # power user: used with advanced RoPE types (e.g. dynamic rope) + def forward(self, x, position_ids): + inv_freq_expanded = self.inv_freq[None, :, None].float().expand(position_ids.shape[0], -1, 1).to(x.device) + position_ids_expanded = position_ids[:, None, :].float() + + device_type = x.device.type if isinstance(x.device.type, str) and x.device.type != "mps" else "cpu" + with maybe_autocast(device_type=device_type, enabled=False): # Force float32 + freqs = (inv_freq_expanded.float() @ position_ids_expanded.float()).transpose(1, 2) + emb = torch.cat((freqs, freqs), dim=-1) + cos = emb.cos() * self.attention_scaling + sin = emb.sin() * self.attention_scaling + + return cos.to(dtype=x.dtype), sin.to(dtype=x.dtype) + + +@auto_docstring +class GlmMoeDsaModel(GlmMoeDsaPreTrainedModel): + _keys_to_ignore_on_load_unexpected = [r"model\.layers\.92.*", r"model\.layers\.46.*"] + + def __init__(self, config: GlmMoeDsaConfig): + super().__init__(config) + self.padding_idx = config.pad_token_id + self.vocab_size = config.vocab_size + + self.embed_tokens = nn.Embedding(config.vocab_size, config.hidden_size, self.padding_idx) + self.layers = nn.ModuleList( + [GlmMoeDsaDecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)] + ) + self.norm = GlmMoeDsaRMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.rotary_emb = GlmMoeDsaRotaryEmbedding(config=config) + self.gradient_checkpointing = False + + # Initialize weights and apply final processing + self.post_init() + + @check_model_inputs + @auto_docstring + def forward( + self, + input_ids: torch.LongTensor | None = None, + attention_mask: torch.Tensor | None = None, + position_ids: torch.LongTensor | None = None, + past_key_values: Cache | None = None, + inputs_embeds: torch.FloatTensor | None = None, + cache_position: torch.LongTensor | None = None, + use_cache: bool | None = None, + **kwargs: Unpack[TransformersKwargs], + ) -> BaseModelOutputWithPast: + if (input_ids is None) ^ (inputs_embeds is not None): + raise ValueError("You must specify exactly one of input_ids or inputs_embeds") + + if inputs_embeds is None: + inputs_embeds: torch.Tensor = self.embed_tokens(input_ids) + + if use_cache and past_key_values is None: + past_key_values = DynamicCache(config=self.config) + + if cache_position is None: + past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0 + cache_position: torch.Tensor = ( + torch.arange(inputs_embeds.shape[1], device=inputs_embeds.device) + past_seen_tokens + ) + + if position_ids is None: + position_ids = cache_position.unsqueeze(0) + + causal_mask = create_causal_mask( + config=self.config, + input_embeds=inputs_embeds, + attention_mask=attention_mask, + cache_position=cache_position, + past_key_values=past_key_values, + position_ids=position_ids, + ) + + hidden_states = inputs_embeds + position_embeddings = self.rotary_emb(hidden_states, position_ids=position_ids) + + for decoder_layer in self.layers[: self.config.num_hidden_layers]: + hidden_states = decoder_layer( + hidden_states, + attention_mask=causal_mask, + position_embeddings=position_embeddings, + position_ids=position_ids, + past_key_values=past_key_values, + use_cache=use_cache, + cache_position=cache_position, + **kwargs, + ) + + hidden_states = self.norm(hidden_states) + return BaseModelOutputWithPast( + last_hidden_state=hidden_states, + past_key_values=past_key_values, + ) + + +@auto_docstring +class GlmMoeDsaForCausalLM(GlmMoeDsaPreTrainedModel, GenerationMixin): + _tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"} + _tp_plan = {"lm_head": "colwise_rep"} + _pp_plan = {"lm_head": (["hidden_states"], ["logits"])} + + def __init__(self, config): + super().__init__(config) + self.model = GlmMoeDsaModel(config) + self.vocab_size = config.vocab_size + self.lm_head = nn.Linear(config.hidden_size, config.vocab_size, bias=False) + + # Initialize weights and apply final processing + self.post_init() + + @can_return_tuple + @auto_docstring + def forward( + self, + input_ids: torch.LongTensor | None = None, + attention_mask: torch.Tensor | None = None, + position_ids: torch.LongTensor | None = None, + past_key_values: Cache | None = None, + inputs_embeds: torch.FloatTensor | None = None, + labels: torch.LongTensor | None = None, + use_cache: bool | None = None, + cache_position: torch.LongTensor | None = None, + logits_to_keep: int | torch.Tensor = 0, + **kwargs: Unpack[TransformersKwargs], + ) -> CausalLMOutputWithPast: + r""" + Example: + + ```python + >>> from transformers import AutoTokenizer, GlmMoeDsaForCausalLM + + >>> model = GlmMoeDsaForCausalLM.from_pretrained("meta-glm_moe_dsa/GlmMoeDsa-2-7b-hf") + >>> tokenizer = AutoTokenizer.from_pretrained("meta-glm_moe_dsa/GlmMoeDsa-2-7b-hf") + + >>> prompt = "Hey, are you conscious? Can you talk to me?" + >>> inputs = tokenizer(prompt, return_tensors="pt") + + >>> # Generate + >>> generate_ids = model.generate(inputs.input_ids, max_length=30) + >>> tokenizer.batch_decode(generate_ids, skip_special_tokens=True, clean_up_tokenization_spaces=False)[0] + "Hey, are you conscious? Can you talk to me?\nI'm not conscious, but I can talk to you." + ```""" + outputs: BaseModelOutputWithPast = self.model( + input_ids=input_ids, + attention_mask=attention_mask, + position_ids=position_ids, + past_key_values=past_key_values, + inputs_embeds=inputs_embeds, + use_cache=use_cache, + cache_position=cache_position, + **kwargs, + ) + + hidden_states = outputs.last_hidden_state + # Only compute necessary logits, and do not upcast them to float if we are not computing the loss + slice_indices = slice(-logits_to_keep, None) if isinstance(logits_to_keep, int) else logits_to_keep + logits = self.lm_head(hidden_states[:, slice_indices, :]) + + loss = None + if labels is not None: + loss = self.loss_function(logits=logits, labels=labels, vocab_size=self.config.vocab_size, **kwargs) + + return CausalLMOutputWithPast( + loss=loss, + logits=logits, + past_key_values=outputs.past_key_values, + hidden_states=outputs.hidden_states, + attentions=outputs.attentions, + ) + + +__all__ = ["GlmMoeDsaPreTrainedModel", "GlmMoeDsaModel", "GlmMoeDsaForCausalLM"] diff --git a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py new file mode 100644 index 000000000000..e8d9936b6dd3 --- /dev/null +++ b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py @@ -0,0 +1,392 @@ +# Copyright 2026 the HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. + +import warnings +from collections.abc import Callable + +import torch +import torch.nn.functional as F +from torch import nn + +from ...cache_utils import Cache +from ...modeling_flash_attention_utils import FlashAttentionKwargs +from ...modeling_utils import ALL_ATTENTION_FUNCTIONS +from ...models.deepseek_v3.modeling_deepseek_v3 import ( + apply_rotary_pos_emb_interleave, + yarn_get_mscale, +) +from ...models.llama.modeling_llama import ( + apply_rotary_pos_emb, + eager_attention_forward, +) +from ...processing_utils import Unpack +from ...utils import logging +from ..glm4_moe.modeling_glm4_moe import ( + Glm4MoeDecoderLayer, + Glm4MoeForCausalLM, + Glm4MoeModel, + Glm4MoePreTrainedModel, + Glm4MoeRMSNorm, +) +from ..glm4_moe_lite.configuration_glm4_moe_lite import Glm4MoeLiteConfig + + +logger = logging.get_logger(__name__) + + +class GlmMoeDsaConfig(Glm4MoeLiteConfig): + r""" + This is the configuration class to store the configuration of a [`GlmMoeDsaModel`]. It is used to instantiate an DeepSeek + model according to the specified arguments, defining the model architecture. Instantiating a configuration with the + defaults will yield a similar configuration to that of the DeepSeek-V3. + e.g. [bzantium/tiny-deepseek-v3](https://huggingface.co/bzantium/tiny-deepseek-v3) + Configuration objects inherit from [`PreTrainedConfig`] and can be used to control the model outputs. Read the + documentation from [`PreTrainedConfig`] for more information. + + + Args: + vocab_size (`int`, *optional*, defaults to 154880): + Vocabulary size of the Deep model. Defines the number of different tokens that can be represented by the + `inputs_ids` passed when calling [`Glm4MoeLiteModel`] + hidden_size (`int`, *optional*, defaults to 6144): + Dimension of the hidden representations. + intermediate_size (`int`, *optional*, defaults to 12288): + Dimension of the MLP representations. + moe_intermediate_size (`int`, *optional*, defaults to 2048): + Dimension of the MoE representations. + num_hidden_layers (`int`, *optional*, defaults to 78): + Number of hidden layers in the Transformer decoder. + num_attention_heads (`int`, *optional*, defaults to 64): + Number of attention heads for each attention layer in the Transformer decoder. + num_key_value_heads (`int`, *optional*, defaults to 64): + This is the number of key_value heads that should be used to implement Grouped Query Attention. If + `num_key_value_heads=num_attention_heads`, the model will use Multi Head Attention (MHA), if + `num_key_value_heads=1 the model will use Multi Query Attention (MQA) otherwise GQA is used. When + converting a multi-head checkpoint to a GQA checkpoint, each group key and value head should be constructed + by meanpooling all the original heads within that group. For more details, check out [this + paper](https://huggingface.co/papers/2305.13245). If it is not specified, will default to + `num_attention_heads`. + n_shared_experts (`int`, *optional*, defaults to 1): + Number of shared experts. + n_routed_experts (`int`, *optional*, defaults to 256): + Number of routed experts. + routed_scaling_factor (`float`, *optional*, defaults to 2.5): + Scaling factor or routed experts. + kv_lora_rank (`int`, *optional*, defaults to 512): + Rank of the LoRA matrices for key and value projections. + q_lora_rank (`int`, *optional*, defaults to 2048): + Rank of the LoRA matrices for query projections. + qk_rope_head_dim (`int`, *optional*, defaults to 64): + Dimension of the query/key heads that use rotary position embeddings. + v_head_dim (`int`, *optional*, defaults to 256): + Dimension of the value heads. + qk_nope_head_dim (`int`, *optional*, defaults to 192): + Dimension of the query/key heads that don't use rotary position embeddings. + n_group (`int`, *optional*, defaults to 1): + Number of groups for routed experts. + topk_group (`int`, *optional*, defaults to 1): + Number of selected groups for each token(for each token, ensuring the selected experts is only within `topk_group` groups). + num_experts_per_tok (`int`, *optional*, defaults to 8): + Number of selected experts, None means dense model. + norm_topk_prob (`bool`, *optional*, defaults to `True`): + Whether to normalize the weights of the routed experts. + hidden_act (`str` or `function`, *optional*, defaults to `"silu"`): + The non-linear activation function (function or string) in the decoder. + max_position_embeddings (`int`, *optional*, defaults to 202752): + The maximum sequence length that this model might ever be used with. + initializer_range (`float`, *optional*, defaults to 0.02): + The standard deviation of the truncated_normal_initializer for initializing all weight matrices. + rms_norm_eps (`float`, *optional*, defaults to 1e-05): + The epsilon used by the rms normalization layers. + use_cache (`bool`, *optional*, defaults to `True`): + Whether or not the model should return the last key/values attentions (not used by all models). Only + relevant if `config.is_decoder=True`. + pad_token_id (`int`, *optional*): + Padding token id. + bos_token_id (`int`, *optional*, defaults to 0): + Beginning of stream token id. + eos_token_id (`int`, *optional*, defaults to 1): + End of stream token id. + pretraining_tp (`int`, *optional*, defaults to 1): + Experimental feature. Tensor parallelism rank used during pretraining. Please refer to [this + document](https://huggingface.co/docs/transformers/parallelism) to understand more about it. This value is + necessary to ensure exact reproducibility of the pretraining results. Please refer to [this + issue](https://github.com/pytorch/pytorch/issues/76232). + tie_word_embeddings (`bool`, *optional*, defaults to `False`): + Whether to tie weight embeddings + rope_parameters (`RopeParameters`, *optional*): + Dictionary containing the configuration parameters for the RoPE embeddings. The dictionary should contain + a value for `rope_theta` and optionally parameters used for scaling in case you want to use RoPE + with longer `max_position_embeddings`. + rope_interleave (`bool`, *optional*, defaults to `True`): + Whether to interleave the rotary position embeddings. + mlp_layer_types (`list`, *optional*): + MLP (Moe vs Dense) pattern for each layer. + attention_bias (`bool`, defaults to `False`, *optional*, defaults to `False`): + Whether to use a bias in the query, key, value and output projection layers during self-attention. + attention_dropout (`float`, *optional*, defaults to 0.0): + The dropout ratio for the attention probabilities. + + ```python + >>> from transformers import Glm4MoeLiteModel, Glm4MoeLiteConfig + + >>> # Initializing a GLM-MOE-DSA style configuration + >>> configuration = GlmMoeDsaConfig() + + >>> # Accessing the model configuration + >>> configuration = model.config + ```""" + + def __init__( + self, + hidden_size: int | None = 6144, + intermediate_size: int | None = 12288, + moe_intermediate_size: int | None = 2048, + num_hidden_layers: int | None = 78, + num_attention_heads: int | None = 64, + num_key_value_heads: int | None = 64, + n_shared_experts: int | None = 1, + n_routed_experts: int | None = 256, + routed_scaling_factor: float | None = 2.5, + kv_lora_rank: int | None = 512, + q_lora_rank: int | None = 2048, + qk_rope_head_dim: int | None = 64, + v_head_dim: int | None = 256, + qk_nope_head_dim: int | None = 192, + num_experts_per_tok: int | None = 8, + initializer_range: float | None = 0.02, + index_topk: int | None = 2048, + **super_kwargs, + ): + self.hidden_size = hidden_size + self.intermediate_size = intermediate_size + self.num_hidden_layers = num_hidden_layers + self.moe_intermediate_size = moe_intermediate_size + self.num_attention_heads = num_attention_heads + self.n_shared_experts = n_shared_experts + self.n_routed_experts = n_routed_experts + self.routed_scaling_factor = routed_scaling_factor + self.kv_lora_rank = kv_lora_rank + self.q_lora_rank = q_lora_rank + self.qk_rope_head_dim = qk_rope_head_dim + self.v_head_dim = v_head_dim + self.qk_nope_head_dim = qk_nope_head_dim + self.qk_head_dim = qk_nope_head_dim + qk_rope_head_dim + self.head_dim = qk_rope_head_dim + self.num_experts_per_tok = num_experts_per_tok + self.num_key_value_heads = num_key_value_heads + self.initializer_range = initializer_range + self.index_topk = index_topk + + super().__init__(**super_kwargs) + + +class GlmMoeDsaRMSNorm(Glm4MoeRMSNorm): + pass + + +class GlmMoeDsaAttention(nn.Module): + """ + DeepSeek V3.2 sparse attention mechanism with indexer. + + This implements the native sparse attention from [DeepSeek V3.2](https://huggingface.co/deepseek-ai/DeepSeek-V3.2) which uses + an indexer to select top-k tokens for attention computation, making it more efficient for long sequences. + + Switch to the implementation from this [PR](https://github.com/huggingface/transformers/pull/41251) as soon as it’s merged. + """ + + def __init__(self, config: GlmMoeDsaConfig, layer_idx: int): + super().__init__() + self.config = config + self.layer_idx = layer_idx + self.num_key_value_groups = config.num_attention_heads // config.num_key_value_heads + self.attention_dropout = config.attention_dropout + self.num_heads = config.num_attention_heads + + self.q_lora_rank = config.q_lora_rank + self.qk_rope_head_dim = config.qk_rope_head_dim + self.kv_lora_rank = config.kv_lora_rank + self.v_head_dim = config.v_head_dim + self.qk_nope_head_dim = config.qk_nope_head_dim + self.qk_head_dim = config.qk_head_dim + self.index_topk = config.index_topk + + self.is_causal = True + + # Query projection + if self.q_lora_rank is None: + self.q_proj = nn.Linear(config.hidden_size, self.num_heads * self.qk_head_dim, bias=False) + else: + self.q_a_proj = nn.Linear(config.hidden_size, config.q_lora_rank, bias=config.attention_bias) + self.q_a_layernorm = GlmMoeDsaRMSNorm(config.q_lora_rank) + self.q_b_proj = nn.Linear(config.q_lora_rank, self.num_heads * self.qk_head_dim, bias=False) + + # Key-Value projections + self.kv_a_proj_with_mqa = nn.Linear( + config.hidden_size, + self.kv_lora_rank + self.qk_rope_head_dim, + bias=config.attention_bias, + ) + self.kv_a_layernorm = GlmMoeDsaRMSNorm(self.kv_lora_rank) + self.kv_b_proj = nn.Linear( + self.kv_lora_rank, + self.num_heads * (self.qk_nope_head_dim + self.v_head_dim), + bias=False, + ) + + # Output projection + self.o_proj = nn.Linear( + self.num_heads * self.v_head_dim, + config.hidden_size, + bias=config.attention_bias, + ) + + # Indexer components for sparse attention + self.wq_b = nn.Linear(config.q_lora_rank, self.num_heads * self.qk_head_dim, bias=False) + self.wk = nn.Linear(config.hidden_size, self.qk_head_dim, bias=config.attention_bias) + self.k_norm = GlmMoeDsaRMSNorm(self.qk_head_dim) + self.weights_proj = nn.Linear(config.hidden_size, self.num_heads, bias=False) + + self.scaling = self.qk_head_dim ** (-0.5) + if self.config.rope_parameters.get("rope_type", "default") != "default": + mscale_all_dim = self.config.rope_parameters.get("mscale_all_dim", 0) + scaling_factor = self.config.rope_parameters["factor"] + if mscale_all_dim: + mscale = yarn_get_mscale(scaling_factor, mscale_all_dim) + self.scaling = self.scaling * mscale * mscale + + def forward( + self, + hidden_states: torch.Tensor, + position_embeddings: tuple[torch.Tensor, torch.Tensor], + attention_mask: torch.Tensor | None, + past_key_values: Cache | None = None, + cache_position: torch.LongTensor | None = None, + **kwargs: Unpack[FlashAttentionKwargs], + ) -> tuple[torch.Tensor, torch.Tensor | None, tuple[torch.Tensor] | None]: + batch_size, seq_length = hidden_states.shape[:-1] + + # For training or when index_topk is not effective, fall back to standard attention + # This is a simplified implementation - in practice, you'd implement the full sparse indexer + if self.training or seq_length <= self.index_topk: + warnings.warn( + "DeepSeek V3.2 sparse attention is not fully implemented in this version. " + "Falling back to standard attention. For production use, please use vLLM or " + "other optimized inference engines.", + UserWarning, + ) + return self._standard_attention( + hidden_states, position_embeddings, attention_mask, past_key_values, cache_position, **kwargs + ) + + # Sparse attention implementation would go here + # This requires custom CUDA kernels for efficient top-k selection and indexing + return self._standard_attention( + hidden_states, position_embeddings, attention_mask, past_key_values, cache_position, **kwargs + ) + + def _standard_attention( + self, + hidden_states: torch.Tensor, + position_embeddings: tuple[torch.Tensor, torch.Tensor], + attention_mask: torch.Tensor | None, + past_key_values: Cache | None = None, + cache_position: torch.LongTensor | None = None, + **kwargs: Unpack[FlashAttentionKwargs], + ) -> tuple[torch.Tensor, torch.Tensor | None, tuple[torch.Tensor] | None]: + """Standard attention fallback (same as DeepSeek V3)""" + batch_size, seq_length = hidden_states.shape[:-1] + query_shape = (batch_size, seq_length, -1, self.qk_head_dim) + key_shape = (batch_size, seq_length, -1, self.qk_nope_head_dim + self.v_head_dim) + + if self.q_lora_rank is None: + q_states = self.q_proj(hidden_states) + else: + q_states = self.q_b_proj(self.q_a_layernorm(self.q_a_proj(hidden_states))) + q_states = q_states.view(query_shape).transpose(1, 2) + q_pass, q_rot = torch.split(q_states, [self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1) + + compressed_kv = self.kv_a_proj_with_mqa(hidden_states) + k_pass, k_rot = torch.split(compressed_kv, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1) + + k_pass = self.kv_b_proj(self.kv_a_layernorm(k_pass)).view(key_shape).transpose(1, 2) + k_pass, value_states = torch.split(k_pass, [self.qk_nope_head_dim, self.v_head_dim], dim=-1) + + k_rot = k_rot.view(batch_size, 1, seq_length, self.qk_rope_head_dim) + + cos, sin = position_embeddings + if self.config.rope_interleave: + q_rot, k_rot = apply_rotary_pos_emb_interleave(q_rot, k_rot, cos, sin) + else: + q_rot, k_rot = apply_rotary_pos_emb(q_rot, k_rot, cos, sin) + k_rot = k_rot.expand(*k_pass.shape[:-1], -1) + + query_states = torch.cat((q_pass, q_rot), dim=-1) + key_states = torch.cat((k_pass, k_rot), dim=-1) + + if past_key_values is not None: + cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position} + key_states, value_states = past_key_values.update(key_states, value_states, self.layer_idx, cache_kwargs) + + if self.config._attn_implementation == "flash_attention_2" and self.qk_head_dim != self.v_head_dim: + value_states = F.pad(value_states, [0, self.qk_head_dim - self.v_head_dim]) + + attention_interface: Callable = eager_attention_forward + if self.config._attn_implementation != "eager": + attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation] + + attn_output, attn_weights = attention_interface( + self, + query_states, + key_states, + value_states, + attention_mask, + dropout=0.0 if not self.training else self.attention_dropout, + scaling=self.scaling, + **kwargs, + ) + + if self.config._attn_implementation == "flash_attention_2" and self.qk_head_dim != self.v_head_dim: + attn_output = attn_output[:, :, :, : self.v_head_dim] + + attn_output = attn_output.reshape(batch_size, seq_length, -1).contiguous() + attn_output = self.o_proj(attn_output) + return attn_output, attn_weights + + +class GlmMoeDsaDecoderLayer(Glm4MoeDecoderLayer): + def __init__(self, config: GlmMoeDsaConfig, layer_idx: int): + super().__init__(config, layer_idx) + + self.self_attn = GlmMoeDsaAttention(config=config, layer_idx=layer_idx) + + +class GlmMoeDsaPreTrainedModel(Glm4MoePreTrainedModel): + pass + + +class GlmMoeDsaModel(Glm4MoeModel): + pass + + +class GlmMoeDsaForCausalLM(Glm4MoeForCausalLM): + pass + + +__all__ = [ + "GlmMoeDsaConfig", + "GlmMoeDsaPreTrainedModel", + "GlmMoeDsaModel", + "GlmMoeDsaForCausalLM", +] diff --git a/tests/models/glm_moe_dsa/__init__.py b/tests/models/glm_moe_dsa/__init__.py new file mode 100644 index 000000000000..e69de29bb2d1 diff --git a/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py b/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py new file mode 100644 index 000000000000..bef53bb769be --- /dev/null +++ b/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py @@ -0,0 +1,112 @@ +# Copyright 2026 the HuggingFace Team. All rights reserved. +# +# Licensed under the Apache License, Version 2.0 (the "License"); +# you may not use this file except in compliance with the License. +# You may obtain a copy of the License at +# +# http://www.apache.org/licenses/LICENSE-2.0 +# +# Unless required by applicable law or agreed to in writing, software +# distributed under the License is distributed on an "AS IS" BASIS, +# WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. +# See the License for the specific language governing permissions and +# limitations under the License. +"""Testing suite for the PyTorch GLM-4.5, GLM-4.6, GLM-4.7 model.""" + +import unittest + +import pytest +import torch +from packaging import version + +from transformers import is_torch_available +from transformers.testing_utils import ( + cleanup, + require_torch, + require_torch_accelerator, + slow, + torch_device, +) + +from ...causal_lm_tester import CausalLMModelTest, CausalLMModelTester + + +if is_torch_available(): + from transformers import AutoTokenizer, GlmMoeDsaForCausalLM, GlmMoeDsaModel + + +class GlmMoeDsaModelTester(CausalLMModelTester): + if is_torch_available(): + base_model_class = GlmMoeDsaModel + + def __init__( + self, + parent, + n_routed_experts=8, + n_shared_experts=1, + n_group=1, + topk_group=1, + num_experts_per_tok=8, + ): + super().__init__(parent=parent, num_experts_per_tok=num_experts_per_tok) + self.n_routed_experts = n_routed_experts + self.n_shared_experts = n_shared_experts + self.n_group = n_group + self.topk_group = topk_group + + +@require_torch +class GlmMoeDsaModelTest(CausalLMModelTest, unittest.TestCase): + model_tester_class = GlmMoeDsaModelTester + # used in `test_torch_compile_for_training`. Skip as "Dynamic control flow in MoE" + _torch_compile_train_cls = None + model_split_percents = [0.5, 0.85, 0.9] # it tries to offload everything with the default value + + +@require_torch_accelerator +@slow +class GlmMoeDsaIntegrationTest(unittest.TestCase): + def tearDown(self): + # See LlamaIntegrationTest.tearDown(). Can be removed once LlamaIntegrationTest.tearDown() is removed. + cleanup(torch_device, gc_collect=False) + + @slow + @require_torch_accelerator + @pytest.mark.torch_compile_test + def test_compile_static_cache(self): + # `torch==2.2` will throw an error on this test (as in other compilation tests), but torch==2.1.2 and torch>2.2 + # work as intended. See https://github.com/pytorch/pytorch/issues/121943 + if version.parse(torch.__version__) < version.parse("2.3.0"): + self.skipTest(reason="This test requires torch >= 2.3 to run.") + + NUM_TOKENS_TO_GENERATE = 40 + EXPECTED_TEXT_COMPLETION = [ + 'hello, world!\'\'\')\nprint(\'hello, world!\')\nprint("hello, world!")\nprint("hello, world!")\nprint("hello, world!")\nprint("hello, world!")\nprint("hello, world!")\n', + "tell me the story of the first Thanksgiving. commonly known as the Pilgrims, arrived in the autumn of 1620. They were seeking religious freedom and a new life in the Plymouth Colony. Their first", + ] + + prompts = ["[gMASK]hello", "[gMASK]tell me"] + tokenizer = AutoTokenizer.from_pretrained("zai-org/GLM-4.5") + model = GlmMoeDsaForCausalLM.from_pretrained("zai-org/GLM-4.5", device_map=torch_device, dtype=torch.bfloat16) + inputs = tokenizer(prompts, return_tensors="pt", padding=True).to(model.device) + + # Dynamic Cache + generated_ids = model.generate(**inputs, max_new_tokens=NUM_TOKENS_TO_GENERATE, do_sample=False) + dynamic_text = tokenizer.batch_decode(generated_ids, skip_special_tokens=True) + self.assertEqual(EXPECTED_TEXT_COMPLETION, dynamic_text) + + # Static Cache + generated_ids = model.generate( + **inputs, max_new_tokens=NUM_TOKENS_TO_GENERATE, do_sample=False, cache_implementation="static" + ) + static_text = tokenizer.batch_decode(generated_ids, skip_special_tokens=True) + self.assertEqual(EXPECTED_TEXT_COMPLETION, static_text) + + # Static Cache + compile + model._cache = None # clear cache object, initialized when we pass `cache_implementation="static"` + model.forward = torch.compile(model.forward, mode="reduce-overhead", fullgraph=True) + generated_ids = model.generate( + **inputs, max_new_tokens=NUM_TOKENS_TO_GENERATE, do_sample=False, cache_implementation="static" + ) + static_compiled_text = tokenizer.batch_decode(generated_ids, skip_special_tokens=True) + self.assertEqual(EXPECTED_TEXT_COMPLETION, static_compiled_text) From d77015003e6265f9a0be5f0140a96bf41c9ca42b Mon Sep 17 00:00:00 2001 From: zRzRzRzRzRzRzR <2448370773@qq.com> Date: Wed, 28 Jan 2026 19:32:10 +0800 Subject: [PATCH 02/19] for review only --- src/transformers/conversion_mapping.py | 1 + .../models/glm_moe_dsa/configuration_glm_moe_dsa.py | 5 +++++ src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py | 3 +++ src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py | 5 +++++ 4 files changed, 14 insertions(+) diff --git a/src/transformers/conversion_mapping.py b/src/transformers/conversion_mapping.py index a413692383f7..ee1769024667 100644 --- a/src/transformers/conversion_mapping.py +++ b/src/transformers/conversion_mapping.py @@ -52,6 +52,7 @@ "ernie4_5_moe": "qwen2_moe", "glm4_moe": "qwen2_moe", "glm4_moe_lite": "qwen2_moe", + "glm_moe_dsa": "qwen2_moe", "glm4v_moe": "qwen2_moe", "longcat_flash": "qwen2_moe", "solar_open": "qwen2_moe", diff --git a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py index 143ae4d7ea8d..39aa1e7ae474 100644 --- a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py @@ -102,6 +102,9 @@ class GlmMoeDsaConfig(PreTrainedConfig): issue](https://github.com/pytorch/pytorch/issues/76232). tie_word_embeddings (`bool`, *optional*, defaults to `False`): Whether to tie weight embeddings + first_k_dense_replace (`int`, *optional*, defaults to 3): + Number of dense layers in shallow layers(embed->dense->dense->...->dense->moe->moe...->lm_head). + \--k dense layers--/ rope_parameters (`RopeParameters`, *optional*): Dictionary containing the configuration parameters for the RoPE embeddings. The dictionary should contain a value for `rope_theta` and optionally parameters used for scaling in case you want to use RoPE @@ -181,6 +184,7 @@ def __init__( mlp_layer_types=None, attention_bias: bool | None = False, attention_dropout: float | None = 0.0, + first_k_dense_replace: int | None = 3, index_topk: int | None = 2048, **kwargs, ): @@ -196,6 +200,7 @@ def __init__( self.q_lora_rank = q_lora_rank self.qk_rope_head_dim = qk_rope_head_dim self.v_head_dim = v_head_dim + self.first_k_dense_replace = first_k_dense_replace self.qk_nope_head_dim = qk_nope_head_dim self.qk_head_dim = qk_nope_head_dim + qk_rope_head_dim self.head_dim = qk_rope_head_dim diff --git a/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py index ff360a75c0c4..2aef55f2ccf6 100644 --- a/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py @@ -243,6 +243,9 @@ def __init__(self, config: GlmMoeDsaConfig, layer_idx: int): self.weights_proj = nn.Linear(config.hidden_size, self.num_heads, bias=False) self.scaling = self.qk_head_dim ** (-0.5) + print(self.config) + print("===") + print(self.config.rope_parameters) if self.config.rope_parameters.get("rope_type", "default") != "default": mscale_all_dim = self.config.rope_parameters.get("mscale_all_dim", 0) scaling_factor = self.config.rope_parameters["factor"] diff --git a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py index e8d9936b6dd3..952fe83ecbf9 100644 --- a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py @@ -125,6 +125,9 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): issue](https://github.com/pytorch/pytorch/issues/76232). tie_word_embeddings (`bool`, *optional*, defaults to `False`): Whether to tie weight embeddings + first_k_dense_replace (`int`, *optional*, defaults to 3): + Number of dense layers in shallow layers(embed->dense->dense->...->dense->moe->moe...->lm_head). + \--k dense layers--/ rope_parameters (`RopeParameters`, *optional*): Dictionary containing the configuration parameters for the RoPE embeddings. The dictionary should contain a value for `rope_theta` and optionally parameters used for scaling in case you want to use RoPE @@ -156,6 +159,7 @@ def __init__( num_hidden_layers: int | None = 78, num_attention_heads: int | None = 64, num_key_value_heads: int | None = 64, + first_k_dense_replace: int | None = 3, n_shared_experts: int | None = 1, n_routed_experts: int | None = 256, routed_scaling_factor: float | None = 2.5, @@ -181,6 +185,7 @@ def __init__( self.q_lora_rank = q_lora_rank self.qk_rope_head_dim = qk_rope_head_dim self.v_head_dim = v_head_dim + self.first_k_dense_replace = first_k_dense_replace self.qk_nope_head_dim = qk_nope_head_dim self.qk_head_dim = qk_nope_head_dim + qk_rope_head_dim self.head_dim = qk_rope_head_dim From 63450b98b6eb31ed3036775fe621f5dcbc6a1a3a Mon Sep 17 00:00:00 2001 From: zRzRzRzRzRzRzR <2448370773@qq.com> Date: Wed, 28 Jan 2026 20:05:15 +0800 Subject: [PATCH 03/19] update ignore layers --- src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py | 5 +---- src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py | 2 +- 2 files changed, 2 insertions(+), 5 deletions(-) diff --git a/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py index 2aef55f2ccf6..4eefdef10d6e 100644 --- a/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py @@ -243,9 +243,6 @@ def __init__(self, config: GlmMoeDsaConfig, layer_idx: int): self.weights_proj = nn.Linear(config.hidden_size, self.num_heads, bias=False) self.scaling = self.qk_head_dim ** (-0.5) - print(self.config) - print("===") - print(self.config.rope_parameters) if self.config.rope_parameters.get("rope_type", "default") != "default": mscale_all_dim = self.config.rope_parameters.get("mscale_all_dim", 0) scaling_factor = self.config.rope_parameters["factor"] @@ -633,7 +630,7 @@ def forward(self, x, position_ids): @auto_docstring class GlmMoeDsaModel(GlmMoeDsaPreTrainedModel): - _keys_to_ignore_on_load_unexpected = [r"model\.layers\.92.*", r"model\.layers\.46.*"] + _keys_to_ignore_on_load_unexpected = [r"model\.layers\.78.*"] def __init__(self, config: GlmMoeDsaConfig): super().__init__(config) diff --git a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py index 952fe83ecbf9..511e263ebfbd 100644 --- a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py @@ -382,7 +382,7 @@ class GlmMoeDsaPreTrainedModel(Glm4MoePreTrainedModel): class GlmMoeDsaModel(Glm4MoeModel): - pass + _keys_to_ignore_on_load_unexpected = [r"model\.layers\.78.*"] class GlmMoeDsaForCausalLM(Glm4MoeForCausalLM): From 367ba3596d53621d6a3f615caec8d93746cfb733 Mon Sep 17 00:00:00 2001 From: zRzRzRzRzRzRzR <2448370773@qq.com> Date: Wed, 28 Jan 2026 20:55:14 +0800 Subject: [PATCH 04/19] add config --- docs/source/en/model_doc/glm_moe_dsa.md | 4 +- .../glm_moe_dsa/configuration_glm_moe_dsa.py | 13 +++-- .../models/glm_moe_dsa/modular_glm_moe_dsa.py | 13 +++-- .../glm_moe_dsa/test_modeling_glm_moe_dsa.py | 51 +++++++++++++------ 4 files changed, 59 insertions(+), 22 deletions(-) diff --git a/docs/source/en/model_doc/glm_moe_dsa.md b/docs/source/en/model_doc/glm_moe_dsa.md index 53682e45610f..cd863d1205c5 100644 --- a/docs/source/en/model_doc/glm_moe_dsa.md +++ b/docs/source/en/model_doc/glm_moe_dsa.md @@ -16,6 +16,7 @@ limitations under the License. ⚠️ Note that this file is in Markdown but contain specific syntax for our doc-builder (similar to MDX) that may not be rendered properly in your Markdown viewer. --> +*This model was released on {release_date} and added to Hugging Face Transformers on 2026-01-28.* # GlmMoeDsa @@ -56,4 +57,5 @@ The original code can be found [here](). ## GlmMoeDsaForCausalLM -[[autodoc]] GlmMoeDsaForCausalLM \ No newline at end of file +[[autodoc]] GlmMoeDsaForCausalLM + - forward \ No newline at end of file diff --git a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py index 39aa1e7ae474..d8198231bb34 100644 --- a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py @@ -102,9 +102,6 @@ class GlmMoeDsaConfig(PreTrainedConfig): issue](https://github.com/pytorch/pytorch/issues/76232). tie_word_embeddings (`bool`, *optional*, defaults to `False`): Whether to tie weight embeddings - first_k_dense_replace (`int`, *optional*, defaults to 3): - Number of dense layers in shallow layers(embed->dense->dense->...->dense->moe->moe...->lm_head). - \--k dense layers--/ rope_parameters (`RopeParameters`, *optional*): Dictionary containing the configuration parameters for the RoPE embeddings. The dictionary should contain a value for `rope_theta` and optionally parameters used for scaling in case you want to use RoPE @@ -117,6 +114,12 @@ class GlmMoeDsaConfig(PreTrainedConfig): Whether to use a bias in the query, key, value and output projection layers during self-attention. attention_dropout (`float`, *optional*, defaults to 0.0): The dropout ratio for the attention probabilities. + first_k_dense_replace (`int`, *optional*, defaults to 3): + Number of dense layers in shallow layers(embed->dense->dense->...->dense->moe->moe...->lm_head). + \--k dense layers--/ + index_topk (`int`, *optional*, defaults to 2048): + index_head_dim (`int`, *optional*, defaults to 128): + index_n_heads (`int`, *optional*, defaults to 32): ```python >>> from transformers import Glm4MoeLiteModel, Glm4MoeLiteConfig @@ -186,6 +189,8 @@ def __init__( attention_dropout: float | None = 0.0, first_k_dense_replace: int | None = 3, index_topk: int | None = 2048, + index_head_dim: int | None = 128, + index_n_heads: int | None = 32, **kwargs, ): self.hidden_size = hidden_size @@ -208,6 +213,8 @@ def __init__( self.num_key_value_heads = num_key_value_heads self.initializer_range = initializer_range self.index_topk = index_topk + self.index_head_dim = index_head_dim + self.index_n_heads = index_n_heads self.vocab_size = vocab_size self.max_position_embeddings = max_position_embeddings self.hidden_size = hidden_size diff --git a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py index 511e263ebfbd..ec8e29f0368a 100644 --- a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py @@ -125,9 +125,6 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): issue](https://github.com/pytorch/pytorch/issues/76232). tie_word_embeddings (`bool`, *optional*, defaults to `False`): Whether to tie weight embeddings - first_k_dense_replace (`int`, *optional*, defaults to 3): - Number of dense layers in shallow layers(embed->dense->dense->...->dense->moe->moe...->lm_head). - \--k dense layers--/ rope_parameters (`RopeParameters`, *optional*): Dictionary containing the configuration parameters for the RoPE embeddings. The dictionary should contain a value for `rope_theta` and optionally parameters used for scaling in case you want to use RoPE @@ -140,6 +137,12 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): Whether to use a bias in the query, key, value and output projection layers during self-attention. attention_dropout (`float`, *optional*, defaults to 0.0): The dropout ratio for the attention probabilities. + first_k_dense_replace (`int`, *optional*, defaults to 3): + Number of dense layers in shallow layers(embed->dense->dense->...->dense->moe->moe...->lm_head). + \--k dense layers--/ + index_topk (`int`, *optional*, defaults to 2048): + index_head_dim (`int`, *optional*, defaults to 128): + index_n_heads (`int`, *optional*, defaults to 32): ```python >>> from transformers import Glm4MoeLiteModel, Glm4MoeLiteConfig @@ -171,6 +174,8 @@ def __init__( num_experts_per_tok: int | None = 8, initializer_range: float | None = 0.02, index_topk: int | None = 2048, + index_head_dim: int | None = 128, + index_n_heads: int | None = 32, **super_kwargs, ): self.hidden_size = hidden_size @@ -193,6 +198,8 @@ def __init__( self.num_key_value_heads = num_key_value_heads self.initializer_range = initializer_range self.index_topk = index_topk + self.index_head_dim = index_head_dim + self.index_n_heads = index_n_heads super().__init__(**super_kwargs) diff --git a/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py b/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py index bef53bb769be..bfcc63ec9268 100644 --- a/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py +++ b/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py @@ -1,4 +1,4 @@ -# Copyright 2026 the HuggingFace Team. All rights reserved. +# Copyright 2025 the HuggingFace Team. All rights reserved. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. @@ -19,7 +19,7 @@ import torch from packaging import version -from transformers import is_torch_available +from transformers import Cache, is_torch_available from transformers.testing_utils import ( cleanup, require_torch, @@ -43,24 +43,43 @@ def __init__( self, parent, n_routed_experts=8, - n_shared_experts=1, - n_group=1, - topk_group=1, - num_experts_per_tok=8, + kv_lora_rank=32, + q_lora_rank=16, + qk_nope_head_dim=64, + qk_rope_head_dim=64, + v_head_dim=128, ): - super().__init__(parent=parent, num_experts_per_tok=num_experts_per_tok) + super().__init__(parent=parent) self.n_routed_experts = n_routed_experts - self.n_shared_experts = n_shared_experts - self.n_group = n_group - self.topk_group = topk_group + self.kv_lora_rank = kv_lora_rank + self.q_lora_rank = q_lora_rank + self.qk_nope_head_dim = qk_nope_head_dim + self.qk_rope_head_dim = qk_rope_head_dim + self.v_head_dim = v_head_dim @require_torch class GlmMoeDsaModelTest(CausalLMModelTest, unittest.TestCase): model_tester_class = GlmMoeDsaModelTester - # used in `test_torch_compile_for_training`. Skip as "Dynamic control flow in MoE" - _torch_compile_train_cls = None - model_split_percents = [0.5, 0.85, 0.9] # it tries to offload everything with the default value + test_all_params_have_gradient = False + model_split_percents = [0.5, 0.7, 0.8] + + def _check_past_key_values_for_generate(self, batch_size, past_key_values, seq_length, config): + """Needs to be overridden as GLM-4.7-Flash has special MLA cache format (though we don't really use the MLA)""" + self.assertIsInstance(past_key_values, Cache) + + # (batch, head, seq_length, head_features) + expected_common_shape = ( + batch_size, + getattr(config, "num_key_value_heads", config.num_attention_heads), + seq_length, + ) + expected_key_shape = expected_common_shape + (config.qk_nope_head_dim + config.qk_rope_head_dim,) + expected_value_shape = expected_common_shape + (config.v_head_dim,) + + for layer in past_key_values.layers: + self.assertEqual(layer.keys.shape, expected_key_shape) + self.assertEqual(layer.values.shape, expected_value_shape) @require_torch_accelerator @@ -86,8 +105,10 @@ def test_compile_static_cache(self): ] prompts = ["[gMASK]hello", "[gMASK]tell me"] - tokenizer = AutoTokenizer.from_pretrained("zai-org/GLM-4.5") - model = GlmMoeDsaForCausalLM.from_pretrained("zai-org/GLM-4.5", device_map=torch_device, dtype=torch.bfloat16) + tokenizer = AutoTokenizer.from_pretrained("zai-org/GLM-4.7-Flash") + model = GlmMoeDsaForCausalLM.from_pretrained( + "zai-org/GLM-4.7-Flash", device_map=torch_device, dtype=torch.bfloat16 + ) inputs = tokenizer(prompts, return_tensors="pt", padding=True).to(model.device) # Dynamic Cache From 270f4696f69ed16eb4d05a5cd2ae012b8298aee9 Mon Sep 17 00:00:00 2001 From: zRzRzRzRzRzRzR <2448370773@qq.com> Date: Wed, 4 Feb 2026 19:36:35 +0800 Subject: [PATCH 05/19] update --- .../glm_moe_dsa/configuration_glm_moe_dsa.py | 9 +- .../glm_moe_dsa/modeling_glm_moe_dsa.py | 423 ++++++++++-------- .../models/glm_moe_dsa/modular_glm_moe_dsa.py | 389 +++++++++------- 3 files changed, 481 insertions(+), 340 deletions(-) diff --git a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py index d8198231bb34..7a8729d2e324 100644 --- a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py @@ -18,16 +18,17 @@ # See the License for the specific language governing permissions and # limitations under the License. + from ...configuration_utils import PreTrainedConfig, layer_type_validation from ...modeling_rope_utils import RopeParameters class GlmMoeDsaConfig(PreTrainedConfig): r""" - This is the configuration class to store the configuration of a [`GlmMoeDsaModel`]. It is used to instantiate an DeepSeek - model according to the specified arguments, defining the model architecture. Instantiating a configuration with the - defaults will yield a similar configuration to that of the DeepSeek-V3. - e.g. [bzantium/tiny-deepseek-v3](https://huggingface.co/bzantium/tiny-deepseek-v3) + This is the configuration class to store the configuration of a [`GlmMoeDsaModel`]. It is used to instantiate a + GLM-5 model according to the specified arguments, defining the model architecture. Instantiating a configuration with the + defaults will yield a similar configuration to that of the GLM-5. + e.g. [zai-org/GLM-5](https://huggingface.co/zai-org/GLM-5) Configuration objects inherit from [`PreTrainedConfig`] and can be used to control the model outputs. Read the documentation from [`PreTrainedConfig`] for more information. diff --git a/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py index 4eefdef10d6e..6df3cc3a4b95 100644 --- a/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py @@ -18,8 +18,8 @@ # See the License for the specific language governing permissions and # limitations under the License. + import math -import warnings from collections.abc import Callable from typing import Optional @@ -33,11 +33,10 @@ from ...generation import GenerationMixin from ...integrations import use_experts_implementation, use_kernel_forward_from_hub, use_kernel_func_from_hub from ...masking_utils import create_causal_mask -from ...modeling_flash_attention_utils import FlashAttentionKwargs from ...modeling_layers import GradientCheckpointingLayer from ...modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast from ...modeling_rope_utils import ROPE_INIT_FUNCTIONS, dynamic_rope_update -from ...modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel +from ...modeling_utils import PreTrainedModel from ...processing_utils import Unpack from ...utils import TransformersKwargs, auto_docstring, can_return_tuple, is_grouped_mm_available from ...utils.generic import check_model_inputs, maybe_autocast @@ -72,15 +71,20 @@ def rotate_half(x): return torch.cat((-x2, x1), dim=-1) -@use_kernel_func_from_hub("rotary_pos_emb") -def apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=1): - """Applies Rotary Position Embedding to the query and key tensors. +def apply_rotary_pos_emb_interleave(q, k, cos, sin, position_ids=None, unsqueeze_dim=1): + r""" + TODO let's just use the original freqcis computation to not have the view + transpose + reshape! This is not optimized! + Applies Rotary Position Embedding to the query and key tensors. Args: q (`torch.Tensor`): The query tensor. k (`torch.Tensor`): The key tensor. cos (`torch.Tensor`): The cosine part of the rotary embedding. sin (`torch.Tensor`): The sine part of the rotary embedding. + position_ids (`torch.Tensor`): + The position indices of the tokens corresponding to the query and key tensors. For example, this can be + used to pass offsetted position ids when working with a KV-cache. unsqueeze_dim (`int`, *optional*, defaults to 1): The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note @@ -93,63 +97,106 @@ def apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=1): """ cos = cos.unsqueeze(unsqueeze_dim) sin = sin.unsqueeze(unsqueeze_dim) - q_embed = (q * cos) + (rotate_half(q) * sin) - k_embed = (k * cos) + (rotate_half(k) * sin) - return q_embed, k_embed + b, h, s, d = q.shape + q = q.view(b, h, s, d // 2, 2).transpose(4, 3).reshape(b, h, s, d) -def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor: - """ - This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch, - num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim) - """ - batch, num_key_value_heads, slen, head_dim = hidden_states.shape - if n_rep == 1: - return hidden_states - hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim) - return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim) + b, h, s, d = k.shape + k = k.view(b, h, s, d // 2, 2).transpose(4, 3).reshape(b, h, s, d) + q_embed = (q * cos) + (rotate_half(q) * sin) + k_embed = (k * cos) + (rotate_half(k) * sin) + return q_embed, k_embed -def eager_attention_forward( - module: nn.Module, - query: torch.Tensor, - key: torch.Tensor, - value: torch.Tensor, - attention_mask: torch.Tensor | None, - scaling: float, - dropout: float = 0.0, - **kwargs: Unpack[TransformersKwargs], -): - key_states = repeat_kv(key, module.num_key_value_groups) - value_states = repeat_kv(value, module.num_key_value_groups) - attn_weights = torch.matmul(query, key_states.transpose(2, 3)) * scaling - if attention_mask is not None: - causal_mask = attention_mask[:, :, :, : key_states.shape[-2]] - attn_weights = attn_weights + causal_mask +class GLmMoeDsaIndexer(nn.Module): + def __init__(self, config: "GlmMoeDsaConfig", index_layer_idx: int): + super().__init__() + self.config = config + self.layer_idx = index_layer_idx + + self.hidden_size: int = config.hidden_size + self.num_heads: int = config.index_n_heads + self.num_local_heads: int = config.index_n_heads # world_size handling can be added as needed + self.head_dim: int = config.index_head_dim + self.qk_rope_head_dim: int = config.qk_rope_head_dim + self.index_topk: int = config.index_topk + self.q_lora_rank: int = config.q_lora_rank + + self.q_b_proj = nn.Linear(self.q_lora_rank, self.num_heads * self.head_dim, bias=False) + self.k_proj = nn.Linear(self.hidden_size, self.head_dim, bias=False) + self.k_layernorm = nn.LayerNorm(self.head_dim) + self.weights_proj = nn.Linear(self.hidden_size, self.num_heads, dtype=torch.get_default_dtype(), bias=False) + self.softmax_scale = self.head_dim**-0.5 - attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query.dtype) - attn_weights = nn.functional.dropout(attn_weights, p=dropout, training=module.training) - attn_output = torch.matmul(attn_weights, value_states) - attn_output = attn_output.transpose(1, 2).contiguous() + @torch.no_grad() + def forward( + self, + hidden_states: torch.Tensor, # [B, S, hidden] + q_resid: torch.Tensor, # [B, S, q_lora_rank] + position_embeddings: tuple[torch.Tensor, torch.Tensor], + attention_mask: torch.Tensor | None, + past_key_values_index: "Cache", + cache_position: torch.LongTensor | None, + ) -> torch.LongTensor: + B, S, _ = hidden_states.shape + cos, sin = position_embeddings - return attn_output, attn_weights + # Queries + q_states = self.q_b_proj(q_resid) # [B, S, H*D] + q_states = q_states.view(B, S, self.num_heads, self.head_dim) # [B, S, H, D] + q_rot, q_pass = torch.split(q_states, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1) + q_rot = apply_rotary_pos_emb_interleave(q_rot, cos, sin) # [B, S, H, rope_D] + q_states = torch.cat([q_rot, q_pass], dim=-1) # [B, S, H, D] + + # Keys + k = self.k_layernorm(self.k_proj(hidden_states)) # [B, S, D] + k_rot, k_pass = torch.split(k, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1) + # MLA uses single-head rope stream, then expands later; keep [B, 1, S, rope_D] here + k_rot = k_rot.unsqueeze(1) # [B, 1, S, rope_D] + k_rot = apply_rotary_pos_emb_interleave(k_rot, cos, sin) # [B, 1, S, rope_D] + k_states = torch.cat( + [ + k_rot.expand(B, self.num_heads, S, -1), # expand rope + k_pass.view(B, 1, S, -1).expand(B, self.num_heads, S, -1), + ], + dim=-1, + ) # [B, H, S, D] + + # Quantize (per provided utilities) + # Update indexer cache (layer idx belongs to the attention layer using this indexer) + # We store as: keys = k_fp8 (as [B, 1, S, D] or [B, H, S, D]? We keep [B, 1, S, D] like original) + # For compactness, collapse heads to 1 for the indexer (you can keep H if your fp8_index expects it). + k_1h = k_states.mean(dim=1, keepdim=True) # [B, 1, S, D] (cheap head merge; adjust if needed) + k_cache = past_key_values_index.update(k_1h, self.layer_idx, cache_kwargs={"cache_position": cache_position}) + + # Weights per head + head_weights = self.weights_proj(hidden_states) * (self.num_heads**-0.5) # [B, S, H] + head_weights = head_weights.unsqueeze(-1) * self.softmax_scale # [B, S, H, *] + logits = torch.matmul(k_cache.unsqueeze(1), q_states.transpose(-1, -2)) # [B, M, N, H] + + # ReLU and sum over heads -> [B, M, N] + logits.clamp_min_(0) + index_scores = logits.sum(dim=-1) # [B, M, N] + + if attention_mask is not None: + index_scores = index_scores + attention_mask + + T = index_scores.shape[-1] + topk = min(self.index_topk, T) + topk_indices = index_scores.topk(topk, dim=-1).indices # [..., topk] + return topk_indices -def apply_rotary_pos_emb_interleave(q, k, cos, sin, position_ids=None, unsqueeze_dim=1): - r""" - TODO let's just use the original freqcis computation to not have the view - transpose + reshape! This is not optimized! - Applies Rotary Position Embedding to the query and key tensors. +@use_kernel_func_from_hub("rotary_pos_emb") +def apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=1): + """Applies Rotary Position Embedding to the query and key tensors. Args: q (`torch.Tensor`): The query tensor. k (`torch.Tensor`): The key tensor. cos (`torch.Tensor`): The cosine part of the rotary embedding. sin (`torch.Tensor`): The sine part of the rotary embedding. - position_ids (`torch.Tensor`): - The position indices of the tokens corresponding to the query and key tensors. For example, this can be - used to pass offsetted position ids when working with a KV-cache. unsqueeze_dim (`int`, *optional*, defaults to 1): The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note @@ -162,24 +209,11 @@ def apply_rotary_pos_emb_interleave(q, k, cos, sin, position_ids=None, unsqueeze """ cos = cos.unsqueeze(unsqueeze_dim) sin = sin.unsqueeze(unsqueeze_dim) - - b, h, s, d = q.shape - q = q.view(b, h, s, d // 2, 2).transpose(4, 3).reshape(b, h, s, d) - - b, h, s, d = k.shape - k = k.view(b, h, s, d // 2, 2).transpose(4, 3).reshape(b, h, s, d) - q_embed = (q * cos) + (rotate_half(q) * sin) k_embed = (k * cos) + (rotate_half(k) * sin) return q_embed, k_embed -def yarn_get_mscale(scale=1, mscale=1): - if scale <= 1: - return 1.0 - return 0.1 * mscale * math.log(scale) + 1.0 - - class GlmMoeDsaAttention(nn.Module): """ DeepSeek V3.2 sparse attention mechanism with indexer. @@ -190,163 +224,202 @@ class GlmMoeDsaAttention(nn.Module): Switch to the implementation from this [PR](https://github.com/huggingface/transformers/pull/41251) as soon as it’s merged. """ - def __init__(self, config: GlmMoeDsaConfig, layer_idx: int): + def __init__(self, config, layer_idx): super().__init__() self.config = config self.layer_idx = layer_idx - self.num_key_value_groups = config.num_attention_heads // config.num_key_value_heads self.attention_dropout = config.attention_dropout + self.hidden_size = config.hidden_size self.num_heads = config.num_attention_heads + self.head_dim = config.head_dim + self.max_position_embeddings = config.max_position_embeddings self.q_lora_rank = config.q_lora_rank self.qk_rope_head_dim = config.qk_rope_head_dim self.kv_lora_rank = config.kv_lora_rank self.v_head_dim = config.v_head_dim self.qk_nope_head_dim = config.qk_nope_head_dim - self.qk_head_dim = config.qk_head_dim - self.index_topk = config.index_topk + self.qk_head_dim = config.qk_nope_head_dim + config.qk_rope_head_dim + self.num_key_value_groups = config.num_attention_heads // config.num_key_value_heads self.is_causal = True - # Query projection if self.q_lora_rank is None: - self.q_proj = nn.Linear(config.hidden_size, self.num_heads * self.qk_head_dim, bias=False) + self.q_proj = nn.Linear(self.hidden_size, self.num_heads * self.qk_head_dim, bias=False) else: - self.q_a_proj = nn.Linear(config.hidden_size, config.q_lora_rank, bias=config.attention_bias) + self.q_a_proj = nn.Linear(self.hidden_size, config.q_lora_rank, bias=config.attention_bias) self.q_a_layernorm = GlmMoeDsaRMSNorm(config.q_lora_rank) self.q_b_proj = nn.Linear(config.q_lora_rank, self.num_heads * self.qk_head_dim, bias=False) - # Key-Value projections self.kv_a_proj_with_mqa = nn.Linear( - config.hidden_size, - self.kv_lora_rank + self.qk_rope_head_dim, + self.hidden_size, + config.kv_lora_rank + config.qk_rope_head_dim, bias=config.attention_bias, ) - self.kv_a_layernorm = GlmMoeDsaRMSNorm(self.kv_lora_rank) + self.kv_a_layernorm = GlmMoeDsaRMSNorm(config.kv_lora_rank) self.kv_b_proj = nn.Linear( - self.kv_lora_rank, - self.num_heads * (self.qk_nope_head_dim + self.v_head_dim), + config.kv_lora_rank, + self.num_heads * (self.qk_head_dim - self.qk_rope_head_dim + self.v_head_dim), bias=False, ) - # Output projection self.o_proj = nn.Linear( self.num_heads * self.v_head_dim, - config.hidden_size, + self.hidden_size, bias=config.attention_bias, ) - # Indexer components for sparse attention - self.wq_b = nn.Linear(config.q_lora_rank, self.num_heads * self.qk_head_dim, bias=False) - self.wk = nn.Linear(config.hidden_size, self.qk_head_dim, bias=config.attention_bias) - self.k_norm = GlmMoeDsaRMSNorm(self.qk_head_dim) - self.weights_proj = nn.Linear(config.hidden_size, self.num_heads, bias=False) - self.scaling = self.qk_head_dim ** (-0.5) - if self.config.rope_parameters.get("rope_type", "default") != "default": - mscale_all_dim = self.config.rope_parameters.get("mscale_all_dim", 0) - scaling_factor = self.config.rope_parameters["factor"] - if mscale_all_dim: - mscale = yarn_get_mscale(scaling_factor, mscale_all_dim) - self.scaling = self.scaling * mscale * mscale + self.softmax_scale = self.qk_head_dim**-0.5 + if config.max_seq_len > config.original_seq_len: + mscale = 0.1 * config.mscale * math.log(config.rope_factor) + 1.0 + self.softmax_scale = self.softmax_scale * mscale * mscale + + self.indexer = GLmMoeDsaIndexer(config, layer_idx) def forward( self, - hidden_states: torch.Tensor, - position_embeddings: tuple[torch.Tensor, torch.Tensor], + hidden_states: torch.Tensor, # [B, S, hidden] + position_embeddings: tuple[torch.Tensor, torch.Tensor], # (cos, sin) attention_mask: torch.Tensor | None, - past_key_values: Cache | None = None, + past_key_values: Cache | None = None, # must be Cache with MlaLayer at `layer_idx` cache_position: torch.LongTensor | None = None, - **kwargs: Unpack[FlashAttentionKwargs], + **kwargs, ) -> tuple[torch.Tensor, torch.Tensor | None, tuple[torch.Tensor] | None]: - batch_size, seq_length = hidden_states.shape[:-1] - - # For training or when index_topk is not effective, fall back to standard attention - # This is a simplified implementation - in practice, you'd implement the full sparse indexer - if self.training or seq_length <= self.index_topk: - warnings.warn( - "DeepSeek V3.2 sparse attention is not fully implemented in this version. " - "Falling back to standard attention. For production use, please use vLLM or " - "other optimized inference engines.", - UserWarning, - ) - return self._standard_attention( - hidden_states, position_embeddings, attention_mask, past_key_values, cache_position, **kwargs - ) - - # Sparse attention implementation would go here - # This requires custom CUDA kernels for efficient top-k selection and indexing - return self._standard_attention( - hidden_states, position_embeddings, attention_mask, past_key_values, cache_position, **kwargs - ) - - def _standard_attention( - self, - hidden_states: torch.Tensor, - position_embeddings: tuple[torch.Tensor, torch.Tensor], - attention_mask: torch.Tensor | None, - past_key_values: Cache | None = None, - cache_position: torch.LongTensor | None = None, - **kwargs: Unpack[FlashAttentionKwargs], - ) -> tuple[torch.Tensor, torch.Tensor | None, tuple[torch.Tensor] | None]: - """Standard attention fallback (same as DeepSeek V3)""" - batch_size, seq_length = hidden_states.shape[:-1] - query_shape = (batch_size, seq_length, -1, self.qk_head_dim) - key_shape = (batch_size, seq_length, -1, self.qk_nope_head_dim + self.v_head_dim) - - if self.q_lora_rank is None: - q_states = self.q_proj(hidden_states) - else: - q_states = self.q_b_proj(self.q_a_layernorm(self.q_a_proj(hidden_states))) - q_states = q_states.view(query_shape).transpose(1, 2) - q_pass, q_rot = torch.split(q_states, [self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1) - - compressed_kv = self.kv_a_proj_with_mqa(hidden_states) - k_pass, k_rot = torch.split(compressed_kv, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1) - - k_pass = self.kv_b_proj(self.kv_a_layernorm(k_pass)).view(key_shape).transpose(1, 2) - k_pass, value_states = torch.split(k_pass, [self.qk_nope_head_dim, self.v_head_dim], dim=-1) - - k_rot = k_rot.view(batch_size, 1, seq_length, self.qk_rope_head_dim) - + B, S, _ = hidden_states.shape cos, sin = position_embeddings - if self.config.rope_interleave: - q_rot, k_rot = apply_rotary_pos_emb_interleave(q_rot, k_rot, cos, sin) - else: - q_rot, k_rot = apply_rotary_pos_emb(q_rot, k_rot, cos, sin) - k_rot = k_rot.expand(*k_pass.shape[:-1], -1) - - query_states = torch.cat((q_pass, q_rot), dim=-1) - key_states = torch.cat((k_pass, k_rot), dim=-1) + # ----- Q path ----- + q_resid = self.q_a_layernorm(self.q_a_proj(hidden_states)) # [B, S, q_lora_rank] + q_states = self.q_b_proj(q_resid).view(B, S, self.num_heads, self.qk_head_dim) # [B, S, H, D] + # Split into pass/rot then apply RoPE on q_rot + q_pass, q_rot = torch.split(q_states, [self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1) + q_rot = apply_rotary_pos_emb(q_rot, cos, sin) # [B, S, H, rope_D] + q_states = torch.cat([q_pass, q_rot], dim=-1) # [B, S, H, D] + + # Layout for matmul: [B, H, S, D] + q_states = q_states.transpose(1, 2).contiguous() # [B, H, S, D] + + # ----- KV path (compressed + rope stream) ----- + kv_all = self.kv_a_proj_with_mqa(hidden_states) # [B, S, kv_rank + rope_D] + kv_compressed, k_rot = torch.split(kv_all, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1) + kv_compressed = self.kv_a_layernorm(kv_compressed) # [B, S, kv_rank] + # Pre-project to K_pass and V + kv_proj = self.kv_b_proj(kv_compressed) # [B, S, H*(qk_nope + v)] + kv_proj = kv_proj.view(B, S, self.num_heads, self.qk_nope_head_dim + self.v_head_dim) + k_pass, v_states = torch.split( + kv_proj, [self.qk_nope_head_dim, self.v_head_dim], dim=-1 + ) # [B,S,H,nope], [B,S,H,V] + + # Rope on K side: keep a single-head rope stream like MLA, then expand + k_rot = k_rot.view(B, 1, S, self.qk_rope_head_dim) # [B, 1, S, rope_D] + k_rot = apply_rotary_pos_emb(k_rot, cos, sin) # [B, 1, S, rope_D] + + # Concatenate K = [K_pass, K_rot(expanded)] + k_states = torch.cat( + ( + k_pass.transpose(1, 2), # [B, H, S, nope_D] + k_rot.expand(B, self.num_heads, S, -1), + ), # [B, H, S, rope_D] + dim=-1, + ) # [B, H, S, D] + v_states = v_states.transpose(1, 2).contiguous() # [B, H, S, V] + + # ----- Cache update/usage ----- if past_key_values is not None: - cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position} - key_states, value_states = past_key_values.update(key_states, value_states, self.layer_idx, cache_kwargs) - - if self.config._attn_implementation == "flash_attention_2" and self.qk_head_dim != self.v_head_dim: - value_states = F.pad(value_states, [0, self.qk_head_dim - self.v_head_dim]) - - attention_interface: Callable = eager_attention_forward - if self.config._attn_implementation != "eager": - attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation] - - attn_output, attn_weights = attention_interface( - self, - query_states, - key_states, - value_states, - attention_mask, - dropout=0.0 if not self.training else self.attention_dropout, - scaling=self.scaling, - **kwargs, - ) + # Store compressed stream & rope stream (as in original MLA path) + # We cache `kv_compressed` under `keys` and `k_rot` under `values` in MlaLayer. + # Shapes must be [B, H, t, *] and [B, 1, t, rope_D]. + kv_comp_cache = kv_compressed.view(B, 1, S, self.kv_lora_rank).expand(B, self.num_heads, S, -1) + k_rot_cache = k_rot # [B, 1, S, rope_D] + cached_kv, cached_pe = past_key_values.update( + kv_comp_cache, k_rot_cache, layer_idx=self.layer_idx, cache_kwargs={"cache_position": cache_position} + ) + # Decode path makes use of cached projections; Prefill can use full K/V directly. + + # ----- Two paths (prefill vs decode) ----- + if attention_mask is not None: + # Prefill (full attention over local window): standard scaled dot-product with top-k pruning from indexer + + # Build scores: [B, H, S, S_total] + # K layout already [B, H, T, D] + scores = (q_states.float() @ k_states.float().transpose(-1, -2)) * self.scaling # [B, H, S, T] + + # Indexer top-k + if past_key_values is not None: + topk_idx = self.indexer( + hidden_states, + q_resid, + position_embeddings, + attention_mask, + past_key_values_index=past_key_values, # we reuse same Cache with IndexerLayer? (separate cache recommended) + cache_position=cache_position, + ) + # Build mask to keep only top-k per (B,S,head?) + # Expect topk_idx shape to broadcast to [B, H, S, T]. We scatter along last dim. + keep_mask = torch.full_like(scores, float("-inf")) + # If topk_idx is [B,S,topk], expand for heads: + if topk_idx.dim() == 3: + topk_idx = topk_idx.unsqueeze(1).expand(B, self.num_heads, S, -1) + keep_mask.scatter_(-1, topk_idx, 0.0) + scores = scores + keep_mask + + probs = nn.functional.softmax(scores, dim=-1, dtype=torch.float32).type_as(hidden_states) # [B, H, S, T] + attn_output = probs @ v_states # [B, H, S, V] + + elif past_key_values is not None: + # Decode: use cached compressed KV & rope stream to recompose attention scores efficiently + # Compose q_pass and q_rot pieces as in MLA math, but via matmul + # 1) Rebuild "nope" term via kv_b weights (dequant on the fly) + wkv_b = self.kv_b_proj.weight.view( + self.num_heads, self.qk_nope_head_dim + self.v_head_dim, self.kv_lora_rank + ) + w_k_nope = wkv_b[:, : self.qk_nope_head_dim, :] # [H, nope_D, kv_rank] + w_v = wkv_b[:, self.qk_nope_head_dim :, :] # [H, V, kv_rank] + + # q_pass: [B,H,S,nope_D]; cached_kv: [B,H,T,kv_rank] + q_pass = q_states[..., : self.qk_nope_head_dim] # [B,H,S,nope_D] + kv_comp = past_key_values[self.layer_idx][0] # keys -> [B,H,T,kv_rank] + pe_full = past_key_values[self.layer_idx][1] # values -> [B,1,T,rope_D] + # Project q_pass with w_k_nope: [B,H,S,kv_rank] + qk_nope = torch.matmul(q_pass, w_k_nope.transpose(-1, -2)) # [B,H,S,kv_rank] + # Scores_nope = qk_nope @ kv_comp^T + scores_nope = torch.matmul(qk_nope.float(), kv_comp.float().transpose(-1, -2)) # [B,H,S,T] + + # 2) Rope term: q_rot @ k_rot^T + q_rot_only = q_states[..., -self.qk_rope_head_dim :] # [B,H,S,rope_D] + k_rot_only = pe_full.expand(B, self.num_heads, -1, -1) # [B,H,T,rope_D] + scores_rot = torch.matmul(q_rot_only.float(), k_rot_only.float().transpose(-1, -2)) # [B,H,S,T] + + scores = (scores_nope + scores_rot) * self.scaling + + # Indexer top-k (decode) + topk_idx = self.indexer( + hidden_states, + q_resid, + position_embeddings, + attention_mask, + past_key_values_index=past_key_values, + cache_position=cache_position, + ) + # For decode single-step S==1 typically; build a [B,H,1,T] mask + keep_mask = torch.full_like(scores, float("-inf")) + if topk_idx.dim() == 3: + topk_idx = topk_idx.unsqueeze(1).expand(B, self.num_heads, S, -1) + keep_mask.scatter_(-1, topk_idx, 0.0) + scores = scores + keep_mask - if self.config._attn_implementation == "flash_attention_2" and self.qk_head_dim != self.v_head_dim: - attn_output = attn_output[:, :, :, : self.v_head_dim] + probs = nn.functional.softmax(scores, dim=-1, dtype=torch.float32).type_as(hidden_states) # [B,H,S,T] - attn_output = attn_output.reshape(batch_size, seq_length, -1).contiguous() - attn_output = self.o_proj(attn_output) - return attn_output, attn_weights + # Rebuild V for decode fast-path: v = (kv_comp @ w_v^T) + # kv_comp: [B,H,T,kv_rank], w_v: [H, V, kv_rank] + v_from_comp = torch.matmul(kv_comp, w_v.transpose(-1, -2)) # [B,H,T,V] + attn_output = torch.matmul(probs, v_from_comp) # [B,H,S,V] + + # Output projection + attn_output = attn_output.transpose(1, 2).reshape(B, S, -1).contiguous() # [B,S,H*V] + attn_output = self.o_proj(attn_output) # [B,S,hidden] + return attn_output, None, None class GlmMoeDsaMLP(nn.Module): diff --git a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py index ec8e29f0368a..c4bdfa5d8adc 100644 --- a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py @@ -12,26 +12,19 @@ # See the License for the specific language governing permissions and # limitations under the License. -import warnings -from collections.abc import Callable + +import math import torch -import torch.nn.functional as F from torch import nn from ...cache_utils import Cache -from ...modeling_flash_attention_utils import FlashAttentionKwargs -from ...modeling_utils import ALL_ATTENTION_FUNCTIONS -from ...models.deepseek_v3.modeling_deepseek_v3 import ( - apply_rotary_pos_emb_interleave, - yarn_get_mscale, -) from ...models.llama.modeling_llama import ( apply_rotary_pos_emb, - eager_attention_forward, ) -from ...processing_utils import Unpack from ...utils import logging +from ..deepseek_v2.modeling_deepseek_v2 import DeepseekV2Attention +from ..deepseek_v3.modeling_deepseek_v3 import apply_rotary_pos_emb_interleave from ..glm4_moe.modeling_glm4_moe import ( Glm4MoeDecoderLayer, Glm4MoeForCausalLM, @@ -47,10 +40,10 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): r""" - This is the configuration class to store the configuration of a [`GlmMoeDsaModel`]. It is used to instantiate an DeepSeek - model according to the specified arguments, defining the model architecture. Instantiating a configuration with the - defaults will yield a similar configuration to that of the DeepSeek-V3. - e.g. [bzantium/tiny-deepseek-v3](https://huggingface.co/bzantium/tiny-deepseek-v3) + This is the configuration class to store the configuration of a [`GlmMoeDsaModel`]. It is used to instantiate a + GLM-5 model according to the specified arguments, defining the model architecture. Instantiating a configuration with the + defaults will yield a similar configuration to that of the GLM-5. + e.g. [zai-org/GLM-5](https://huggingface.co/zai-org/GLM-5) Configuration objects inherit from [`PreTrainedConfig`] and can be used to control the model outputs. Read the documentation from [`PreTrainedConfig`] for more information. @@ -208,7 +201,86 @@ class GlmMoeDsaRMSNorm(Glm4MoeRMSNorm): pass -class GlmMoeDsaAttention(nn.Module): +class GLmMoeDsaIndexer(nn.Module): + def __init__(self, config: "GlmMoeDsaConfig", index_layer_idx: int): + super().__init__() + self.config = config + self.layer_idx = index_layer_idx + + self.hidden_size: int = config.hidden_size + self.num_heads: int = config.index_n_heads + self.num_local_heads: int = config.index_n_heads # world_size handling can be added as needed + self.head_dim: int = config.index_head_dim + self.qk_rope_head_dim: int = config.qk_rope_head_dim + self.index_topk: int = config.index_topk + self.q_lora_rank: int = config.q_lora_rank + + self.q_b_proj = nn.Linear(self.q_lora_rank, self.num_heads * self.head_dim, bias=False) + self.k_proj = nn.Linear(self.hidden_size, self.head_dim, bias=False) + self.k_layernorm = nn.LayerNorm(self.head_dim) + self.weights_proj = nn.Linear(self.hidden_size, self.num_heads, dtype=torch.get_default_dtype(), bias=False) + self.softmax_scale = self.head_dim**-0.5 + + @torch.no_grad() + def forward( + self, + hidden_states: torch.Tensor, # [B, S, hidden] + q_resid: torch.Tensor, # [B, S, q_lora_rank] + position_embeddings: tuple[torch.Tensor, torch.Tensor], + attention_mask: torch.Tensor | None, + past_key_values_index: "Cache", + cache_position: torch.LongTensor | None, + ) -> torch.LongTensor: + B, S, _ = hidden_states.shape + cos, sin = position_embeddings + + # Queries + q_states = self.q_b_proj(q_resid) # [B, S, H*D] + q_states = q_states.view(B, S, self.num_heads, self.head_dim) # [B, S, H, D] + q_rot, q_pass = torch.split(q_states, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1) + q_rot = apply_rotary_pos_emb_interleave(q_rot, cos, sin) # [B, S, H, rope_D] + q_states = torch.cat([q_rot, q_pass], dim=-1) # [B, S, H, D] + + # Keys + k = self.k_layernorm(self.k_proj(hidden_states)) # [B, S, D] + k_rot, k_pass = torch.split(k, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1) + # MLA uses single-head rope stream, then expands later; keep [B, 1, S, rope_D] here + k_rot = k_rot.unsqueeze(1) # [B, 1, S, rope_D] + k_rot = apply_rotary_pos_emb_interleave(k_rot, cos, sin) # [B, 1, S, rope_D] + k_states = torch.cat( + [ + k_rot.expand(B, self.num_heads, S, -1), # expand rope + k_pass.view(B, 1, S, -1).expand(B, self.num_heads, S, -1), + ], + dim=-1, + ) # [B, H, S, D] + + # Quantize (per provided utilities) + # Update indexer cache (layer idx belongs to the attention layer using this indexer) + # We store as: keys = k_fp8 (as [B, 1, S, D] or [B, H, S, D]? We keep [B, 1, S, D] like original) + # For compactness, collapse heads to 1 for the indexer (you can keep H if your fp8_index expects it). + k_1h = k_states.mean(dim=1, keepdim=True) # [B, 1, S, D] (cheap head merge; adjust if needed) + k_cache = past_key_values_index.update(k_1h, self.layer_idx, cache_kwargs={"cache_position": cache_position}) + + # Weights per head + head_weights = self.weights_proj(hidden_states) * (self.num_heads**-0.5) # [B, S, H] + head_weights = head_weights.unsqueeze(-1) * self.softmax_scale # [B, S, H, *] + logits = torch.matmul(k_cache.unsqueeze(1), q_states.transpose(-1, -2)) # [B, M, N, H] + + # ReLU and sum over heads -> [B, M, N] + logits.clamp_min_(0) + index_scores = logits.sum(dim=-1) # [B, M, N] + + if attention_mask is not None: + index_scores = index_scores + attention_mask + + T = index_scores.shape[-1] + topk = min(self.index_topk, T) + topk_indices = index_scores.topk(topk, dim=-1).indices # [..., topk] + return topk_indices + + +class GlmMoeDsaAttention(DeepseekV2Attention): """ DeepSeek V3.2 sparse attention mechanism with indexer. @@ -218,163 +290,158 @@ class GlmMoeDsaAttention(nn.Module): Switch to the implementation from this [PR](https://github.com/huggingface/transformers/pull/41251) as soon as it’s merged. """ - def __init__(self, config: GlmMoeDsaConfig, layer_idx: int): - super().__init__() - self.config = config - self.layer_idx = layer_idx - self.num_key_value_groups = config.num_attention_heads // config.num_key_value_heads - self.attention_dropout = config.attention_dropout - self.num_heads = config.num_attention_heads - - self.q_lora_rank = config.q_lora_rank - self.qk_rope_head_dim = config.qk_rope_head_dim - self.kv_lora_rank = config.kv_lora_rank - self.v_head_dim = config.v_head_dim - self.qk_nope_head_dim = config.qk_nope_head_dim - self.qk_head_dim = config.qk_head_dim - self.index_topk = config.index_topk - - self.is_causal = True - - # Query projection - if self.q_lora_rank is None: - self.q_proj = nn.Linear(config.hidden_size, self.num_heads * self.qk_head_dim, bias=False) - else: - self.q_a_proj = nn.Linear(config.hidden_size, config.q_lora_rank, bias=config.attention_bias) - self.q_a_layernorm = GlmMoeDsaRMSNorm(config.q_lora_rank) - self.q_b_proj = nn.Linear(config.q_lora_rank, self.num_heads * self.qk_head_dim, bias=False) - - # Key-Value projections - self.kv_a_proj_with_mqa = nn.Linear( - config.hidden_size, - self.kv_lora_rank + self.qk_rope_head_dim, - bias=config.attention_bias, - ) - self.kv_a_layernorm = GlmMoeDsaRMSNorm(self.kv_lora_rank) - self.kv_b_proj = nn.Linear( - self.kv_lora_rank, - self.num_heads * (self.qk_nope_head_dim + self.v_head_dim), - bias=False, - ) + def __init__(self, config, layer_idx): + super().__init__(config, layer_idx) + self.softmax_scale = self.qk_head_dim**-0.5 + if config.max_seq_len > config.original_seq_len: + mscale = 0.1 * config.mscale * math.log(config.rope_factor) + 1.0 + self.softmax_scale = self.softmax_scale * mscale * mscale - # Output projection - self.o_proj = nn.Linear( - self.num_heads * self.v_head_dim, - config.hidden_size, - bias=config.attention_bias, - ) - - # Indexer components for sparse attention - self.wq_b = nn.Linear(config.q_lora_rank, self.num_heads * self.qk_head_dim, bias=False) - self.wk = nn.Linear(config.hidden_size, self.qk_head_dim, bias=config.attention_bias) - self.k_norm = GlmMoeDsaRMSNorm(self.qk_head_dim) - self.weights_proj = nn.Linear(config.hidden_size, self.num_heads, bias=False) - - self.scaling = self.qk_head_dim ** (-0.5) - if self.config.rope_parameters.get("rope_type", "default") != "default": - mscale_all_dim = self.config.rope_parameters.get("mscale_all_dim", 0) - scaling_factor = self.config.rope_parameters["factor"] - if mscale_all_dim: - mscale = yarn_get_mscale(scaling_factor, mscale_all_dim) - self.scaling = self.scaling * mscale * mscale + self.indexer = GLmMoeDsaIndexer(config, layer_idx) def forward( self, - hidden_states: torch.Tensor, - position_embeddings: tuple[torch.Tensor, torch.Tensor], + hidden_states: torch.Tensor, # [B, S, hidden] + position_embeddings: tuple[torch.Tensor, torch.Tensor], # (cos, sin) attention_mask: torch.Tensor | None, - past_key_values: Cache | None = None, + past_key_values: Cache | None = None, # must be Cache with MlaLayer at `layer_idx` cache_position: torch.LongTensor | None = None, - **kwargs: Unpack[FlashAttentionKwargs], + **kwargs, ) -> tuple[torch.Tensor, torch.Tensor | None, tuple[torch.Tensor] | None]: - batch_size, seq_length = hidden_states.shape[:-1] - - # For training or when index_topk is not effective, fall back to standard attention - # This is a simplified implementation - in practice, you'd implement the full sparse indexer - if self.training or seq_length <= self.index_topk: - warnings.warn( - "DeepSeek V3.2 sparse attention is not fully implemented in this version. " - "Falling back to standard attention. For production use, please use vLLM or " - "other optimized inference engines.", - UserWarning, - ) - return self._standard_attention( - hidden_states, position_embeddings, attention_mask, past_key_values, cache_position, **kwargs - ) - - # Sparse attention implementation would go here - # This requires custom CUDA kernels for efficient top-k selection and indexing - return self._standard_attention( - hidden_states, position_embeddings, attention_mask, past_key_values, cache_position, **kwargs - ) + B, S, _ = hidden_states.shape + cos, sin = position_embeddings - def _standard_attention( - self, - hidden_states: torch.Tensor, - position_embeddings: tuple[torch.Tensor, torch.Tensor], - attention_mask: torch.Tensor | None, - past_key_values: Cache | None = None, - cache_position: torch.LongTensor | None = None, - **kwargs: Unpack[FlashAttentionKwargs], - ) -> tuple[torch.Tensor, torch.Tensor | None, tuple[torch.Tensor] | None]: - """Standard attention fallback (same as DeepSeek V3)""" - batch_size, seq_length = hidden_states.shape[:-1] - query_shape = (batch_size, seq_length, -1, self.qk_head_dim) - key_shape = (batch_size, seq_length, -1, self.qk_nope_head_dim + self.v_head_dim) - - if self.q_lora_rank is None: - q_states = self.q_proj(hidden_states) - else: - q_states = self.q_b_proj(self.q_a_layernorm(self.q_a_proj(hidden_states))) - q_states = q_states.view(query_shape).transpose(1, 2) + # ----- Q path ----- + q_resid = self.q_a_layernorm(self.q_a_proj(hidden_states)) # [B, S, q_lora_rank] + q_states = self.q_b_proj(q_resid).view(B, S, self.num_heads, self.qk_head_dim) # [B, S, H, D] + # Split into pass/rot then apply RoPE on q_rot q_pass, q_rot = torch.split(q_states, [self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1) + q_rot = apply_rotary_pos_emb(q_rot, cos, sin) # [B, S, H, rope_D] + q_states = torch.cat([q_pass, q_rot], dim=-1) # [B, S, H, D] + + # Layout for matmul: [B, H, S, D] + q_states = q_states.transpose(1, 2).contiguous() # [B, H, S, D] + + # ----- KV path (compressed + rope stream) ----- + kv_all = self.kv_a_proj_with_mqa(hidden_states) # [B, S, kv_rank + rope_D] + kv_compressed, k_rot = torch.split(kv_all, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1) + kv_compressed = self.kv_a_layernorm(kv_compressed) # [B, S, kv_rank] + # Pre-project to K_pass and V + kv_proj = self.kv_b_proj(kv_compressed) # [B, S, H*(qk_nope + v)] + kv_proj = kv_proj.view(B, S, self.num_heads, self.qk_nope_head_dim + self.v_head_dim) + k_pass, v_states = torch.split( + kv_proj, [self.qk_nope_head_dim, self.v_head_dim], dim=-1 + ) # [B,S,H,nope], [B,S,H,V] + + # Rope on K side: keep a single-head rope stream like MLA, then expand + k_rot = k_rot.view(B, 1, S, self.qk_rope_head_dim) # [B, 1, S, rope_D] + k_rot = apply_rotary_pos_emb(k_rot, cos, sin) # [B, 1, S, rope_D] + + # Concatenate K = [K_pass, K_rot(expanded)] + k_states = torch.cat( + ( + k_pass.transpose(1, 2), # [B, H, S, nope_D] + k_rot.expand(B, self.num_heads, S, -1), + ), # [B, H, S, rope_D] + dim=-1, + ) # [B, H, S, D] + v_states = v_states.transpose(1, 2).contiguous() # [B, H, S, V] + + # ----- Cache update/usage ----- + if past_key_values is not None: + # Store compressed stream & rope stream (as in original MLA path) + # We cache `kv_compressed` under `keys` and `k_rot` under `values` in MlaLayer. + # Shapes must be [B, H, t, *] and [B, 1, t, rope_D]. + kv_comp_cache = kv_compressed.view(B, 1, S, self.kv_lora_rank).expand(B, self.num_heads, S, -1) + k_rot_cache = k_rot # [B, 1, S, rope_D] + cached_kv, cached_pe = past_key_values.update( + kv_comp_cache, k_rot_cache, layer_idx=self.layer_idx, cache_kwargs={"cache_position": cache_position} + ) + # Decode path makes use of cached projections; Prefill can use full K/V directly. + + # ----- Two paths (prefill vs decode) ----- + if attention_mask is not None: + # Prefill (full attention over local window): standard scaled dot-product with top-k pruning from indexer + + # Build scores: [B, H, S, S_total] + # K layout already [B, H, T, D] + scores = (q_states.float() @ k_states.float().transpose(-1, -2)) * self.scaling # [B, H, S, T] + + # Indexer top-k + if past_key_values is not None: + topk_idx = self.indexer( + hidden_states, + q_resid, + position_embeddings, + attention_mask, + past_key_values_index=past_key_values, # we reuse same Cache with IndexerLayer? (separate cache recommended) + cache_position=cache_position, + ) + # Build mask to keep only top-k per (B,S,head?) + # Expect topk_idx shape to broadcast to [B, H, S, T]. We scatter along last dim. + keep_mask = torch.full_like(scores, float("-inf")) + # If topk_idx is [B,S,topk], expand for heads: + if topk_idx.dim() == 3: + topk_idx = topk_idx.unsqueeze(1).expand(B, self.num_heads, S, -1) + keep_mask.scatter_(-1, topk_idx, 0.0) + scores = scores + keep_mask + + probs = nn.functional.softmax(scores, dim=-1, dtype=torch.float32).type_as(hidden_states) # [B, H, S, T] + attn_output = probs @ v_states # [B, H, S, V] + + elif past_key_values is not None: + # Decode: use cached compressed KV & rope stream to recompose attention scores efficiently + # Compose q_pass and q_rot pieces as in MLA math, but via matmul + # 1) Rebuild "nope" term via kv_b weights (dequant on the fly) + wkv_b = self.kv_b_proj.weight.view( + self.num_heads, self.qk_nope_head_dim + self.v_head_dim, self.kv_lora_rank + ) + w_k_nope = wkv_b[:, : self.qk_nope_head_dim, :] # [H, nope_D, kv_rank] + w_v = wkv_b[:, self.qk_nope_head_dim :, :] # [H, V, kv_rank] + + # q_pass: [B,H,S,nope_D]; cached_kv: [B,H,T,kv_rank] + q_pass = q_states[..., : self.qk_nope_head_dim] # [B,H,S,nope_D] + kv_comp = past_key_values[self.layer_idx][0] # keys -> [B,H,T,kv_rank] + pe_full = past_key_values[self.layer_idx][1] # values -> [B,1,T,rope_D] + # Project q_pass with w_k_nope: [B,H,S,kv_rank] + qk_nope = torch.matmul(q_pass, w_k_nope.transpose(-1, -2)) # [B,H,S,kv_rank] + # Scores_nope = qk_nope @ kv_comp^T + scores_nope = torch.matmul(qk_nope.float(), kv_comp.float().transpose(-1, -2)) # [B,H,S,T] + + # 2) Rope term: q_rot @ k_rot^T + q_rot_only = q_states[..., -self.qk_rope_head_dim :] # [B,H,S,rope_D] + k_rot_only = pe_full.expand(B, self.num_heads, -1, -1) # [B,H,T,rope_D] + scores_rot = torch.matmul(q_rot_only.float(), k_rot_only.float().transpose(-1, -2)) # [B,H,S,T] + + scores = (scores_nope + scores_rot) * self.scaling + + # Indexer top-k (decode) + topk_idx = self.indexer( + hidden_states, + q_resid, + position_embeddings, + attention_mask, + past_key_values_index=past_key_values, + cache_position=cache_position, + ) + # For decode single-step S==1 typically; build a [B,H,1,T] mask + keep_mask = torch.full_like(scores, float("-inf")) + if topk_idx.dim() == 3: + topk_idx = topk_idx.unsqueeze(1).expand(B, self.num_heads, S, -1) + keep_mask.scatter_(-1, topk_idx, 0.0) + scores = scores + keep_mask - compressed_kv = self.kv_a_proj_with_mqa(hidden_states) - k_pass, k_rot = torch.split(compressed_kv, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1) - - k_pass = self.kv_b_proj(self.kv_a_layernorm(k_pass)).view(key_shape).transpose(1, 2) - k_pass, value_states = torch.split(k_pass, [self.qk_nope_head_dim, self.v_head_dim], dim=-1) - - k_rot = k_rot.view(batch_size, 1, seq_length, self.qk_rope_head_dim) - - cos, sin = position_embeddings - if self.config.rope_interleave: - q_rot, k_rot = apply_rotary_pos_emb_interleave(q_rot, k_rot, cos, sin) - else: - q_rot, k_rot = apply_rotary_pos_emb(q_rot, k_rot, cos, sin) - k_rot = k_rot.expand(*k_pass.shape[:-1], -1) + probs = nn.functional.softmax(scores, dim=-1, dtype=torch.float32).type_as(hidden_states) # [B,H,S,T] - query_states = torch.cat((q_pass, q_rot), dim=-1) - key_states = torch.cat((k_pass, k_rot), dim=-1) + # Rebuild V for decode fast-path: v = (kv_comp @ w_v^T) + # kv_comp: [B,H,T,kv_rank], w_v: [H, V, kv_rank] + v_from_comp = torch.matmul(kv_comp, w_v.transpose(-1, -2)) # [B,H,T,V] + attn_output = torch.matmul(probs, v_from_comp) # [B,H,S,V] - if past_key_values is not None: - cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position} - key_states, value_states = past_key_values.update(key_states, value_states, self.layer_idx, cache_kwargs) - - if self.config._attn_implementation == "flash_attention_2" and self.qk_head_dim != self.v_head_dim: - value_states = F.pad(value_states, [0, self.qk_head_dim - self.v_head_dim]) - - attention_interface: Callable = eager_attention_forward - if self.config._attn_implementation != "eager": - attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation] - - attn_output, attn_weights = attention_interface( - self, - query_states, - key_states, - value_states, - attention_mask, - dropout=0.0 if not self.training else self.attention_dropout, - scaling=self.scaling, - **kwargs, - ) - - if self.config._attn_implementation == "flash_attention_2" and self.qk_head_dim != self.v_head_dim: - attn_output = attn_output[:, :, :, : self.v_head_dim] - - attn_output = attn_output.reshape(batch_size, seq_length, -1).contiguous() - attn_output = self.o_proj(attn_output) - return attn_output, attn_weights + # Output projection + attn_output = attn_output.transpose(1, 2).reshape(B, S, -1).contiguous() # [B,S,H*V] + attn_output = self.o_proj(attn_output) # [B,S,hidden] + return attn_output, None, None class GlmMoeDsaDecoderLayer(Glm4MoeDecoderLayer): From 58f800570504168ee4299e15e99aa59c5e5f1ba5 Mon Sep 17 00:00:00 2001 From: zRzRzRzRzRzRzR <2448370773@qq.com> Date: Sat, 7 Feb 2026 09:41:44 +0100 Subject: [PATCH 06/19] rename --- .../models/glm_moe_dsa/configuration_glm_moe_dsa.py | 9 ++++++--- .../models/glm_moe_dsa/modular_glm_moe_dsa.py | 9 ++++++--- 2 files changed, 12 insertions(+), 6 deletions(-) diff --git a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py index 7a8729d2e324..f62d08ca50b5 100644 --- a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py @@ -118,9 +118,12 @@ class GlmMoeDsaConfig(PreTrainedConfig): first_k_dense_replace (`int`, *optional*, defaults to 3): Number of dense layers in shallow layers(embed->dense->dense->...->dense->moe->moe...->lm_head). \--k dense layers--/ - index_topk (`int`, *optional*, defaults to 2048): - index_head_dim (`int`, *optional*, defaults to 128): - index_n_heads (`int`, *optional*, defaults to 32): + index_topk (`int`, *optional*, defaults to 2048): + Number of top tokens selected by the indexer for retrieval/attention in each step. + index_head_dim (`int`, *optional*, defaults to 128): + Hidden size (per-head dimension) of each indexer attention head. + index_n_heads (`int`, *optional*, defaults to 32): + Number of attention heads used by the indexer module. ```python >>> from transformers import Glm4MoeLiteModel, Glm4MoeLiteConfig diff --git a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py index c4bdfa5d8adc..074553d641a3 100644 --- a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py @@ -133,9 +133,12 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): first_k_dense_replace (`int`, *optional*, defaults to 3): Number of dense layers in shallow layers(embed->dense->dense->...->dense->moe->moe...->lm_head). \--k dense layers--/ - index_topk (`int`, *optional*, defaults to 2048): - index_head_dim (`int`, *optional*, defaults to 128): - index_n_heads (`int`, *optional*, defaults to 32): + index_topk (`int`, *optional*, defaults to 2048): + Number of top tokens selected by the indexer for retrieval/attention in each step. + index_head_dim (`int`, *optional*, defaults to 128): + Hidden size (per-head dimension) of each indexer attention head. + index_n_heads (`int`, *optional*, defaults to 32): + Number of attention heads used by the indexer module. ```python >>> from transformers import Glm4MoeLiteModel, Glm4MoeLiteConfig From 9861b580c6e9caf4923f3d06f68b4aa304f5a5d7 Mon Sep 17 00:00:00 2001 From: zRzRzRzRzRzRzR <2448370773@qq.com> Date: Sat, 7 Feb 2026 18:57:48 +0100 Subject: [PATCH 07/19] fallback --- .../glm_moe_dsa/modeling_glm_moe_dsa.py | 445 ++++++++---------- .../models/glm_moe_dsa/modular_glm_moe_dsa.py | 386 +++++++-------- 2 files changed, 345 insertions(+), 486 deletions(-) diff --git a/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py index 6df3cc3a4b95..c47688bc1a0e 100644 --- a/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py @@ -20,12 +20,13 @@ import math +import warnings from collections.abc import Callable from typing import Optional import torch +import torch.nn as nn import torch.nn.functional as F -from torch import nn from ... import initialization as init from ...activations import ACT2FN @@ -33,10 +34,11 @@ from ...generation import GenerationMixin from ...integrations import use_experts_implementation, use_kernel_forward_from_hub, use_kernel_func_from_hub from ...masking_utils import create_causal_mask +from ...modeling_flash_attention_utils import FlashAttentionKwargs from ...modeling_layers import GradientCheckpointingLayer from ...modeling_outputs import BaseModelOutputWithPast, CausalLMOutputWithPast from ...modeling_rope_utils import ROPE_INIT_FUNCTIONS, dynamic_rope_update -from ...modeling_utils import PreTrainedModel +from ...modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel from ...processing_utils import Unpack from ...utils import TransformersKwargs, auto_docstring, can_return_tuple, is_grouped_mm_available from ...utils.generic import check_model_inputs, maybe_autocast @@ -64,6 +66,32 @@ def extra_repr(self): return f"{tuple(self.weight.shape)}, eps={self.variance_epsilon}" +@use_kernel_func_from_hub("rotary_pos_emb") +def apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=1): + """Applies Rotary Position Embedding to the query and key tensors. + + Args: + q (`torch.Tensor`): The query tensor. + k (`torch.Tensor`): The key tensor. + cos (`torch.Tensor`): The cosine part of the rotary embedding. + sin (`torch.Tensor`): The sine part of the rotary embedding. + unsqueeze_dim (`int`, *optional*, defaults to 1): + The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and + sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note + that cos[position_ids] and sin[position_ids] have the shape [batch_size, seq_len, head_dim]. Then, if q and + k have the shape [batch_size, heads, seq_len, head_dim], then setting unsqueeze_dim=1 makes + cos[position_ids] and sin[position_ids] broadcastable to the shapes of q and k. Similarly, if q and k have + the shape [batch_size, seq_len, heads, head_dim], then set unsqueeze_dim=2. + Returns: + `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding. + """ + cos = cos.unsqueeze(unsqueeze_dim) + sin = sin.unsqueeze(unsqueeze_dim) + q_embed = (q * cos) + (rotate_half(q) * sin) + k_embed = (k * cos) + (rotate_half(k) * sin) + return q_embed, k_embed + + def rotate_half(x): """Rotates half the hidden dims of the input.""" x1 = x[..., : x.shape[-1] // 2] @@ -109,109 +137,48 @@ def apply_rotary_pos_emb_interleave(q, k, cos, sin, position_ids=None, unsqueeze return q_embed, k_embed -class GLmMoeDsaIndexer(nn.Module): - def __init__(self, config: "GlmMoeDsaConfig", index_layer_idx: int): - super().__init__() - self.config = config - self.layer_idx = index_layer_idx - - self.hidden_size: int = config.hidden_size - self.num_heads: int = config.index_n_heads - self.num_local_heads: int = config.index_n_heads # world_size handling can be added as needed - self.head_dim: int = config.index_head_dim - self.qk_rope_head_dim: int = config.qk_rope_head_dim - self.index_topk: int = config.index_topk - self.q_lora_rank: int = config.q_lora_rank - - self.q_b_proj = nn.Linear(self.q_lora_rank, self.num_heads * self.head_dim, bias=False) - self.k_proj = nn.Linear(self.hidden_size, self.head_dim, bias=False) - self.k_layernorm = nn.LayerNorm(self.head_dim) - self.weights_proj = nn.Linear(self.hidden_size, self.num_heads, dtype=torch.get_default_dtype(), bias=False) - self.softmax_scale = self.head_dim**-0.5 +def yarn_get_mscale(scale=1, mscale=1): + if scale <= 1: + return 1.0 + return 0.1 * mscale * math.log(scale) + 1.0 - @torch.no_grad() - def forward( - self, - hidden_states: torch.Tensor, # [B, S, hidden] - q_resid: torch.Tensor, # [B, S, q_lora_rank] - position_embeddings: tuple[torch.Tensor, torch.Tensor], - attention_mask: torch.Tensor | None, - past_key_values_index: "Cache", - cache_position: torch.LongTensor | None, - ) -> torch.LongTensor: - B, S, _ = hidden_states.shape - cos, sin = position_embeddings - # Queries - q_states = self.q_b_proj(q_resid) # [B, S, H*D] - q_states = q_states.view(B, S, self.num_heads, self.head_dim) # [B, S, H, D] - q_rot, q_pass = torch.split(q_states, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1) - q_rot = apply_rotary_pos_emb_interleave(q_rot, cos, sin) # [B, S, H, rope_D] - q_states = torch.cat([q_rot, q_pass], dim=-1) # [B, S, H, D] - - # Keys - k = self.k_layernorm(self.k_proj(hidden_states)) # [B, S, D] - k_rot, k_pass = torch.split(k, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1) - # MLA uses single-head rope stream, then expands later; keep [B, 1, S, rope_D] here - k_rot = k_rot.unsqueeze(1) # [B, 1, S, rope_D] - k_rot = apply_rotary_pos_emb_interleave(k_rot, cos, sin) # [B, 1, S, rope_D] - k_states = torch.cat( - [ - k_rot.expand(B, self.num_heads, S, -1), # expand rope - k_pass.view(B, 1, S, -1).expand(B, self.num_heads, S, -1), - ], - dim=-1, - ) # [B, H, S, D] - - # Quantize (per provided utilities) - # Update indexer cache (layer idx belongs to the attention layer using this indexer) - # We store as: keys = k_fp8 (as [B, 1, S, D] or [B, H, S, D]? We keep [B, 1, S, D] like original) - # For compactness, collapse heads to 1 for the indexer (you can keep H if your fp8_index expects it). - k_1h = k_states.mean(dim=1, keepdim=True) # [B, 1, S, D] (cheap head merge; adjust if needed) - k_cache = past_key_values_index.update(k_1h, self.layer_idx, cache_kwargs={"cache_position": cache_position}) - - # Weights per head - head_weights = self.weights_proj(hidden_states) * (self.num_heads**-0.5) # [B, S, H] - head_weights = head_weights.unsqueeze(-1) * self.softmax_scale # [B, S, H, *] - logits = torch.matmul(k_cache.unsqueeze(1), q_states.transpose(-1, -2)) # [B, M, N, H] - - # ReLU and sum over heads -> [B, M, N] - logits.clamp_min_(0) - index_scores = logits.sum(dim=-1) # [B, M, N] - - if attention_mask is not None: - index_scores = index_scores + attention_mask - - T = index_scores.shape[-1] - topk = min(self.index_topk, T) - topk_indices = index_scores.topk(topk, dim=-1).indices # [..., topk] - return topk_indices +def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor: + """ + This is the equivalent of torch.repeat_interleave(x, dim=1, repeats=n_rep). The hidden states go from (batch, + num_key_value_heads, seqlen, head_dim) to (batch, num_attention_heads, seqlen, head_dim) + """ + batch, num_key_value_heads, slen, head_dim = hidden_states.shape + if n_rep == 1: + return hidden_states + hidden_states = hidden_states[:, :, None, :, :].expand(batch, num_key_value_heads, n_rep, slen, head_dim) + return hidden_states.reshape(batch, num_key_value_heads * n_rep, slen, head_dim) -@use_kernel_func_from_hub("rotary_pos_emb") -def apply_rotary_pos_emb(q, k, cos, sin, unsqueeze_dim=1): - """Applies Rotary Position Embedding to the query and key tensors. +def eager_attention_forward( + module: nn.Module, + query: torch.Tensor, + key: torch.Tensor, + value: torch.Tensor, + attention_mask: torch.Tensor | None, + scaling: float, + dropout: float = 0.0, + **kwargs: Unpack[TransformersKwargs], +): + key_states = repeat_kv(key, module.num_key_value_groups) + value_states = repeat_kv(value, module.num_key_value_groups) - Args: - q (`torch.Tensor`): The query tensor. - k (`torch.Tensor`): The key tensor. - cos (`torch.Tensor`): The cosine part of the rotary embedding. - sin (`torch.Tensor`): The sine part of the rotary embedding. - unsqueeze_dim (`int`, *optional*, defaults to 1): - The 'unsqueeze_dim' argument specifies the dimension along which to unsqueeze cos[position_ids] and - sin[position_ids] so that they can be properly broadcasted to the dimensions of q and k. For example, note - that cos[position_ids] and sin[position_ids] have the shape [batch_size, seq_len, head_dim]. Then, if q and - k have the shape [batch_size, heads, seq_len, head_dim], then setting unsqueeze_dim=1 makes - cos[position_ids] and sin[position_ids] broadcastable to the shapes of q and k. Similarly, if q and k have - the shape [batch_size, seq_len, heads, head_dim], then set unsqueeze_dim=2. - Returns: - `tuple(torch.Tensor)` comprising of the query and key tensors rotated using the Rotary Position Embedding. - """ - cos = cos.unsqueeze(unsqueeze_dim) - sin = sin.unsqueeze(unsqueeze_dim) - q_embed = (q * cos) + (rotate_half(q) * sin) - k_embed = (k * cos) + (rotate_half(k) * sin) - return q_embed, k_embed + attn_weights = torch.matmul(query, key_states.transpose(2, 3)) * scaling + if attention_mask is not None: + causal_mask = attention_mask[:, :, :, : key_states.shape[-2]] + attn_weights = attn_weights + causal_mask + + attn_weights = nn.functional.softmax(attn_weights, dim=-1, dtype=torch.float32).to(query.dtype) + attn_weights = nn.functional.dropout(attn_weights, p=dropout, training=module.training) + attn_output = torch.matmul(attn_weights, value_states) + attn_output = attn_output.transpose(1, 2).contiguous() + + return attn_output, attn_weights class GlmMoeDsaAttention(nn.Module): @@ -221,205 +188,168 @@ class GlmMoeDsaAttention(nn.Module): This implements the native sparse attention from [DeepSeek V3.2](https://huggingface.co/deepseek-ai/DeepSeek-V3.2) which uses an indexer to select top-k tokens for attention computation, making it more efficient for long sequences. + In GLM-5, the indexer RoPE uses neox_style = false. Therefore, we introduced the indexer_rope_interleave parameter: when indexer_rope_interleave is set to True, RoPE is computed using the same neox_style = false behavior as in the GLMMOEDSA model. This part has not yet been implemented in transformers. + Switch to the implementation from this [PR](https://github.com/huggingface/transformers/pull/41251) as soon as it’s merged. """ - def __init__(self, config, layer_idx): + def __init__(self, config: GlmMoeDsaConfig, layer_idx: int): super().__init__() self.config = config self.layer_idx = layer_idx + self.num_key_value_groups = config.num_attention_heads // config.num_key_value_heads self.attention_dropout = config.attention_dropout - self.hidden_size = config.hidden_size self.num_heads = config.num_attention_heads - self.head_dim = config.head_dim - self.max_position_embeddings = config.max_position_embeddings self.q_lora_rank = config.q_lora_rank self.qk_rope_head_dim = config.qk_rope_head_dim self.kv_lora_rank = config.kv_lora_rank self.v_head_dim = config.v_head_dim self.qk_nope_head_dim = config.qk_nope_head_dim - self.qk_head_dim = config.qk_nope_head_dim + config.qk_rope_head_dim - self.num_key_value_groups = config.num_attention_heads // config.num_key_value_heads + self.qk_head_dim = config.qk_head_dim + self.index_topk = config.index_topk self.is_causal = True + # Query projection if self.q_lora_rank is None: - self.q_proj = nn.Linear(self.hidden_size, self.num_heads * self.qk_head_dim, bias=False) + self.q_proj = nn.Linear(config.hidden_size, self.num_heads * self.qk_head_dim, bias=False) else: - self.q_a_proj = nn.Linear(self.hidden_size, config.q_lora_rank, bias=config.attention_bias) + self.q_a_proj = nn.Linear(config.hidden_size, config.q_lora_rank, bias=config.attention_bias) self.q_a_layernorm = GlmMoeDsaRMSNorm(config.q_lora_rank) self.q_b_proj = nn.Linear(config.q_lora_rank, self.num_heads * self.qk_head_dim, bias=False) + # Key-Value projections self.kv_a_proj_with_mqa = nn.Linear( - self.hidden_size, - config.kv_lora_rank + config.qk_rope_head_dim, + config.hidden_size, + self.kv_lora_rank + self.qk_rope_head_dim, bias=config.attention_bias, ) - self.kv_a_layernorm = GlmMoeDsaRMSNorm(config.kv_lora_rank) + self.kv_a_layernorm = GlmMoeDsaRMSNorm(self.kv_lora_rank) self.kv_b_proj = nn.Linear( - config.kv_lora_rank, - self.num_heads * (self.qk_head_dim - self.qk_rope_head_dim + self.v_head_dim), + self.kv_lora_rank, + self.num_heads * (self.qk_nope_head_dim + self.v_head_dim), bias=False, ) + # Output projection self.o_proj = nn.Linear( self.num_heads * self.v_head_dim, - self.hidden_size, + config.hidden_size, bias=config.attention_bias, ) - self.scaling = self.qk_head_dim ** (-0.5) - self.softmax_scale = self.qk_head_dim**-0.5 - if config.max_seq_len > config.original_seq_len: - mscale = 0.1 * config.mscale * math.log(config.rope_factor) + 1.0 - self.softmax_scale = self.softmax_scale * mscale * mscale + # Indexer components for sparse attention + self.wq_b = nn.Linear(config.q_lora_rank, self.num_heads * self.qk_head_dim, bias=False) + self.wk = nn.Linear(config.hidden_size, self.qk_head_dim, bias=config.attention_bias) + self.k_norm = GlmMoeDsaRMSNorm(self.qk_head_dim) + self.weights_proj = nn.Linear(config.hidden_size, self.num_heads, bias=False) - self.indexer = GLmMoeDsaIndexer(config, layer_idx) + self.scaling = self.qk_head_dim ** (-0.5) + if self.config.rope_parameters.get("rope_type", "default") != "default": + mscale_all_dim = self.config.rope_parameters.get("mscale_all_dim", 0) + scaling_factor = self.config.rope_parameters["factor"] + if mscale_all_dim: + mscale = yarn_get_mscale(scaling_factor, mscale_all_dim) + self.scaling = self.scaling * mscale * mscale def forward( self, - hidden_states: torch.Tensor, # [B, S, hidden] - position_embeddings: tuple[torch.Tensor, torch.Tensor], # (cos, sin) + hidden_states: torch.Tensor, + position_embeddings: tuple[torch.Tensor, torch.Tensor], attention_mask: torch.Tensor | None, - past_key_values: Cache | None = None, # must be Cache with MlaLayer at `layer_idx` + past_key_values: Cache | None = None, cache_position: torch.LongTensor | None = None, - **kwargs, + **kwargs: Unpack[FlashAttentionKwargs], ) -> tuple[torch.Tensor, torch.Tensor | None, tuple[torch.Tensor] | None]: - B, S, _ = hidden_states.shape - cos, sin = position_embeddings - - # ----- Q path ----- - q_resid = self.q_a_layernorm(self.q_a_proj(hidden_states)) # [B, S, q_lora_rank] - q_states = self.q_b_proj(q_resid).view(B, S, self.num_heads, self.qk_head_dim) # [B, S, H, D] - # Split into pass/rot then apply RoPE on q_rot - q_pass, q_rot = torch.split(q_states, [self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1) - q_rot = apply_rotary_pos_emb(q_rot, cos, sin) # [B, S, H, rope_D] - q_states = torch.cat([q_pass, q_rot], dim=-1) # [B, S, H, D] - - # Layout for matmul: [B, H, S, D] - q_states = q_states.transpose(1, 2).contiguous() # [B, H, S, D] - - # ----- KV path (compressed + rope stream) ----- - kv_all = self.kv_a_proj_with_mqa(hidden_states) # [B, S, kv_rank + rope_D] - kv_compressed, k_rot = torch.split(kv_all, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1) - kv_compressed = self.kv_a_layernorm(kv_compressed) # [B, S, kv_rank] - # Pre-project to K_pass and V - kv_proj = self.kv_b_proj(kv_compressed) # [B, S, H*(qk_nope + v)] - kv_proj = kv_proj.view(B, S, self.num_heads, self.qk_nope_head_dim + self.v_head_dim) - k_pass, v_states = torch.split( - kv_proj, [self.qk_nope_head_dim, self.v_head_dim], dim=-1 - ) # [B,S,H,nope], [B,S,H,V] - - # Rope on K side: keep a single-head rope stream like MLA, then expand - k_rot = k_rot.view(B, 1, S, self.qk_rope_head_dim) # [B, 1, S, rope_D] - k_rot = apply_rotary_pos_emb(k_rot, cos, sin) # [B, 1, S, rope_D] - - # Concatenate K = [K_pass, K_rot(expanded)] - k_states = torch.cat( - ( - k_pass.transpose(1, 2), # [B, H, S, nope_D] - k_rot.expand(B, self.num_heads, S, -1), - ), # [B, H, S, rope_D] - dim=-1, - ) # [B, H, S, D] - v_states = v_states.transpose(1, 2).contiguous() # [B, H, S, V] - - # ----- Cache update/usage ----- - if past_key_values is not None: - # Store compressed stream & rope stream (as in original MLA path) - # We cache `kv_compressed` under `keys` and `k_rot` under `values` in MlaLayer. - # Shapes must be [B, H, t, *] and [B, 1, t, rope_D]. - kv_comp_cache = kv_compressed.view(B, 1, S, self.kv_lora_rank).expand(B, self.num_heads, S, -1) - k_rot_cache = k_rot # [B, 1, S, rope_D] - cached_kv, cached_pe = past_key_values.update( - kv_comp_cache, k_rot_cache, layer_idx=self.layer_idx, cache_kwargs={"cache_position": cache_position} + batch_size, seq_length = hidden_states.shape[:-1] + + # For training or when index_topk is not effective, fall back to standard attention + # This is a simplified implementation - in practice, you'd implement the full sparse indexer + if self.training or seq_length <= self.index_topk: + warnings.warn( + "DeepSeek V3.2 sparse attention is not fully implemented in this version. " + "Falling back to standard attention. For production use, please use vLLM or " + "other optimized inference engines.", + UserWarning, ) - # Decode path makes use of cached projections; Prefill can use full K/V directly. - - # ----- Two paths (prefill vs decode) ----- - if attention_mask is not None: - # Prefill (full attention over local window): standard scaled dot-product with top-k pruning from indexer - - # Build scores: [B, H, S, S_total] - # K layout already [B, H, T, D] - scores = (q_states.float() @ k_states.float().transpose(-1, -2)) * self.scaling # [B, H, S, T] - - # Indexer top-k - if past_key_values is not None: - topk_idx = self.indexer( - hidden_states, - q_resid, - position_embeddings, - attention_mask, - past_key_values_index=past_key_values, # we reuse same Cache with IndexerLayer? (separate cache recommended) - cache_position=cache_position, - ) - # Build mask to keep only top-k per (B,S,head?) - # Expect topk_idx shape to broadcast to [B, H, S, T]. We scatter along last dim. - keep_mask = torch.full_like(scores, float("-inf")) - # If topk_idx is [B,S,topk], expand for heads: - if topk_idx.dim() == 3: - topk_idx = topk_idx.unsqueeze(1).expand(B, self.num_heads, S, -1) - keep_mask.scatter_(-1, topk_idx, 0.0) - scores = scores + keep_mask - - probs = nn.functional.softmax(scores, dim=-1, dtype=torch.float32).type_as(hidden_states) # [B, H, S, T] - attn_output = probs @ v_states # [B, H, S, V] - - elif past_key_values is not None: - # Decode: use cached compressed KV & rope stream to recompose attention scores efficiently - # Compose q_pass and q_rot pieces as in MLA math, but via matmul - # 1) Rebuild "nope" term via kv_b weights (dequant on the fly) - wkv_b = self.kv_b_proj.weight.view( - self.num_heads, self.qk_nope_head_dim + self.v_head_dim, self.kv_lora_rank - ) - w_k_nope = wkv_b[:, : self.qk_nope_head_dim, :] # [H, nope_D, kv_rank] - w_v = wkv_b[:, self.qk_nope_head_dim :, :] # [H, V, kv_rank] - - # q_pass: [B,H,S,nope_D]; cached_kv: [B,H,T,kv_rank] - q_pass = q_states[..., : self.qk_nope_head_dim] # [B,H,S,nope_D] - kv_comp = past_key_values[self.layer_idx][0] # keys -> [B,H,T,kv_rank] - pe_full = past_key_values[self.layer_idx][1] # values -> [B,1,T,rope_D] - # Project q_pass with w_k_nope: [B,H,S,kv_rank] - qk_nope = torch.matmul(q_pass, w_k_nope.transpose(-1, -2)) # [B,H,S,kv_rank] - # Scores_nope = qk_nope @ kv_comp^T - scores_nope = torch.matmul(qk_nope.float(), kv_comp.float().transpose(-1, -2)) # [B,H,S,T] - - # 2) Rope term: q_rot @ k_rot^T - q_rot_only = q_states[..., -self.qk_rope_head_dim :] # [B,H,S,rope_D] - k_rot_only = pe_full.expand(B, self.num_heads, -1, -1) # [B,H,T,rope_D] - scores_rot = torch.matmul(q_rot_only.float(), k_rot_only.float().transpose(-1, -2)) # [B,H,S,T] - - scores = (scores_nope + scores_rot) * self.scaling - - # Indexer top-k (decode) - topk_idx = self.indexer( - hidden_states, - q_resid, - position_embeddings, - attention_mask, - past_key_values_index=past_key_values, - cache_position=cache_position, + return self._standard_attention( + hidden_states, position_embeddings, attention_mask, past_key_values, cache_position, **kwargs ) - # For decode single-step S==1 typically; build a [B,H,1,T] mask - keep_mask = torch.full_like(scores, float("-inf")) - if topk_idx.dim() == 3: - topk_idx = topk_idx.unsqueeze(1).expand(B, self.num_heads, S, -1) - keep_mask.scatter_(-1, topk_idx, 0.0) - scores = scores + keep_mask - probs = nn.functional.softmax(scores, dim=-1, dtype=torch.float32).type_as(hidden_states) # [B,H,S,T] + # Sparse attention implementation would go here + # This requires custom CUDA kernels for efficient top-k selection and indexing + return self._standard_attention( + hidden_states, position_embeddings, attention_mask, past_key_values, cache_position, **kwargs + ) - # Rebuild V for decode fast-path: v = (kv_comp @ w_v^T) - # kv_comp: [B,H,T,kv_rank], w_v: [H, V, kv_rank] - v_from_comp = torch.matmul(kv_comp, w_v.transpose(-1, -2)) # [B,H,T,V] - attn_output = torch.matmul(probs, v_from_comp) # [B,H,S,V] + def _standard_attention( + self, + hidden_states: torch.Tensor, + position_embeddings: tuple[torch.Tensor, torch.Tensor], + attention_mask: torch.Tensor | None, + past_key_values: Cache | None = None, + cache_position: torch.LongTensor | None = None, + **kwargs: Unpack[FlashAttentionKwargs], + ) -> tuple[torch.Tensor, torch.Tensor | None, tuple[torch.Tensor] | None]: + """Standard attention fallback (same as DeepSeek V3)""" + batch_size, seq_length = hidden_states.shape[:-1] + query_shape = (batch_size, seq_length, -1, self.qk_head_dim) + key_shape = (batch_size, seq_length, -1, self.qk_nope_head_dim + self.v_head_dim) - # Output projection - attn_output = attn_output.transpose(1, 2).reshape(B, S, -1).contiguous() # [B,S,H*V] - attn_output = self.o_proj(attn_output) # [B,S,hidden] - return attn_output, None, None + if self.q_lora_rank is None: + q_states = self.q_proj(hidden_states) + else: + q_states = self.q_b_proj(self.q_a_layernorm(self.q_a_proj(hidden_states))) + q_states = q_states.view(query_shape).transpose(1, 2) + q_pass, q_rot = torch.split(q_states, [self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1) + + compressed_kv = self.kv_a_proj_with_mqa(hidden_states) + k_pass, k_rot = torch.split(compressed_kv, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1) + + k_pass = self.kv_b_proj(self.kv_a_layernorm(k_pass)).view(key_shape).transpose(1, 2) + k_pass, value_states = torch.split(k_pass, [self.qk_nope_head_dim, self.v_head_dim], dim=-1) + + k_rot = k_rot.view(batch_size, 1, seq_length, self.qk_rope_head_dim) + + cos, sin = position_embeddings + if self.config.rope_interleave: + q_rot, k_rot = apply_rotary_pos_emb_interleave(q_rot, k_rot, cos, sin) + else: + q_rot, k_rot = apply_rotary_pos_emb(q_rot, k_rot, cos, sin) + k_rot = k_rot.expand(*k_pass.shape[:-1], -1) + + query_states = torch.cat((q_pass, q_rot), dim=-1) + key_states = torch.cat((k_pass, k_rot), dim=-1) + + if past_key_values is not None: + cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position} + key_states, value_states = past_key_values.update(key_states, value_states, self.layer_idx, cache_kwargs) + + if self.config._attn_implementation == "flash_attention_2" and self.qk_head_dim != self.v_head_dim: + value_states = F.pad(value_states, [0, self.qk_head_dim - self.v_head_dim]) + + attention_interface: Callable = eager_attention_forward + if self.config._attn_implementation != "eager": + attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation] + + attn_output, attn_weights = attention_interface( + self, + query_states, + key_states, + value_states, + attention_mask, + dropout=0.0 if not self.training else self.attention_dropout, + scaling=self.scaling, + **kwargs, + ) + + if self.config._attn_implementation == "flash_attention_2" and self.qk_head_dim != self.v_head_dim: + attn_output = attn_output[:, :, :, : self.v_head_dim] + + attn_output = attn_output.reshape(batch_size, seq_length, -1).contiguous() + attn_output = self.o_proj(attn_output) + return attn_output, attn_weights class GlmMoeDsaMLP(nn.Module): @@ -558,16 +488,15 @@ class GlmMoeDsaDecoderLayer(GradientCheckpointingLayer): def __init__(self, config: GlmMoeDsaConfig, layer_idx: int): super().__init__() self.hidden_size = config.hidden_size + self.self_attn = GlmMoeDsaAttention(config, layer_idx) - self.self_attn = GlmMoeDsaAttention(config=config, layer_idx=layer_idx) - - if layer_idx >= config.first_k_dense_replace: + if config.mlp_layer_types[layer_idx] == "sparse": self.mlp = GlmMoeDsaMoE(config) else: self.mlp = GlmMoeDsaMLP(config) - self.input_layernorm = GlmMoeDsaRMSNorm(config.hidden_size, eps=config.rms_norm_eps) - self.post_attention_layernorm = GlmMoeDsaRMSNorm(config.hidden_size, eps=config.rms_norm_eps) + self.input_layernorm = GlmMoeDsaRMSNorm(config.hidden_size, config.rms_norm_eps) + self.post_attention_layernorm = GlmMoeDsaRMSNorm(config.hidden_size, config.rms_norm_eps) def forward( self, diff --git a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py index 074553d641a3..fb90cbad9016 100644 --- a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py @@ -13,26 +13,31 @@ # limitations under the License. -import math +import warnings +from collections.abc import Callable import torch -from torch import nn +import torch.nn as nn +import torch.nn.functional as F from ...cache_utils import Cache +from ...modeling_flash_attention_utils import FlashAttentionKwargs +from ...modeling_utils import ALL_ATTENTION_FUNCTIONS from ...models.llama.modeling_llama import ( apply_rotary_pos_emb, ) +from ...processing_utils import Unpack from ...utils import logging -from ..deepseek_v2.modeling_deepseek_v2 import DeepseekV2Attention -from ..deepseek_v3.modeling_deepseek_v3 import apply_rotary_pos_emb_interleave +from ..deepseek_v3.modeling_deepseek_v3 import apply_rotary_pos_emb_interleave, yarn_get_mscale from ..glm4_moe.modeling_glm4_moe import ( - Glm4MoeDecoderLayer, Glm4MoeForCausalLM, Glm4MoeModel, Glm4MoePreTrainedModel, Glm4MoeRMSNorm, + eager_attention_forward, ) from ..glm4_moe_lite.configuration_glm4_moe_lite import Glm4MoeLiteConfig +from ..glm4_moe_lite.modeling_glm4_moe_lite import Glm4MoeLiteDecoderLayer logger = logging.get_logger(__name__) @@ -204,254 +209,179 @@ class GlmMoeDsaRMSNorm(Glm4MoeRMSNorm): pass -class GLmMoeDsaIndexer(nn.Module): - def __init__(self, config: "GlmMoeDsaConfig", index_layer_idx: int): - super().__init__() - self.config = config - self.layer_idx = index_layer_idx - - self.hidden_size: int = config.hidden_size - self.num_heads: int = config.index_n_heads - self.num_local_heads: int = config.index_n_heads # world_size handling can be added as needed - self.head_dim: int = config.index_head_dim - self.qk_rope_head_dim: int = config.qk_rope_head_dim - self.index_topk: int = config.index_topk - self.q_lora_rank: int = config.q_lora_rank - - self.q_b_proj = nn.Linear(self.q_lora_rank, self.num_heads * self.head_dim, bias=False) - self.k_proj = nn.Linear(self.hidden_size, self.head_dim, bias=False) - self.k_layernorm = nn.LayerNorm(self.head_dim) - self.weights_proj = nn.Linear(self.hidden_size, self.num_heads, dtype=torch.get_default_dtype(), bias=False) - self.softmax_scale = self.head_dim**-0.5 - - @torch.no_grad() - def forward( - self, - hidden_states: torch.Tensor, # [B, S, hidden] - q_resid: torch.Tensor, # [B, S, q_lora_rank] - position_embeddings: tuple[torch.Tensor, torch.Tensor], - attention_mask: torch.Tensor | None, - past_key_values_index: "Cache", - cache_position: torch.LongTensor | None, - ) -> torch.LongTensor: - B, S, _ = hidden_states.shape - cos, sin = position_embeddings - - # Queries - q_states = self.q_b_proj(q_resid) # [B, S, H*D] - q_states = q_states.view(B, S, self.num_heads, self.head_dim) # [B, S, H, D] - q_rot, q_pass = torch.split(q_states, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1) - q_rot = apply_rotary_pos_emb_interleave(q_rot, cos, sin) # [B, S, H, rope_D] - q_states = torch.cat([q_rot, q_pass], dim=-1) # [B, S, H, D] - - # Keys - k = self.k_layernorm(self.k_proj(hidden_states)) # [B, S, D] - k_rot, k_pass = torch.split(k, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1) - # MLA uses single-head rope stream, then expands later; keep [B, 1, S, rope_D] here - k_rot = k_rot.unsqueeze(1) # [B, 1, S, rope_D] - k_rot = apply_rotary_pos_emb_interleave(k_rot, cos, sin) # [B, 1, S, rope_D] - k_states = torch.cat( - [ - k_rot.expand(B, self.num_heads, S, -1), # expand rope - k_pass.view(B, 1, S, -1).expand(B, self.num_heads, S, -1), - ], - dim=-1, - ) # [B, H, S, D] - - # Quantize (per provided utilities) - # Update indexer cache (layer idx belongs to the attention layer using this indexer) - # We store as: keys = k_fp8 (as [B, 1, S, D] or [B, H, S, D]? We keep [B, 1, S, D] like original) - # For compactness, collapse heads to 1 for the indexer (you can keep H if your fp8_index expects it). - k_1h = k_states.mean(dim=1, keepdim=True) # [B, 1, S, D] (cheap head merge; adjust if needed) - k_cache = past_key_values_index.update(k_1h, self.layer_idx, cache_kwargs={"cache_position": cache_position}) - - # Weights per head - head_weights = self.weights_proj(hidden_states) * (self.num_heads**-0.5) # [B, S, H] - head_weights = head_weights.unsqueeze(-1) * self.softmax_scale # [B, S, H, *] - logits = torch.matmul(k_cache.unsqueeze(1), q_states.transpose(-1, -2)) # [B, M, N, H] - - # ReLU and sum over heads -> [B, M, N] - logits.clamp_min_(0) - index_scores = logits.sum(dim=-1) # [B, M, N] - - if attention_mask is not None: - index_scores = index_scores + attention_mask - - T = index_scores.shape[-1] - topk = min(self.index_topk, T) - topk_indices = index_scores.topk(topk, dim=-1).indices # [..., topk] - return topk_indices - - -class GlmMoeDsaAttention(DeepseekV2Attention): +class GlmMoeDsaAttention(nn.Module): """ DeepSeek V3.2 sparse attention mechanism with indexer. This implements the native sparse attention from [DeepSeek V3.2](https://huggingface.co/deepseek-ai/DeepSeek-V3.2) which uses an indexer to select top-k tokens for attention computation, making it more efficient for long sequences. + In GLM-5, the indexer RoPE uses neox_style = false. Therefore, we introduced the indexer_rope_interleave parameter: when indexer_rope_interleave is set to True, RoPE is computed using the same neox_style = false behavior as in the GLMMOEDSA model. This part has not yet been implemented in transformers. + Switch to the implementation from this [PR](https://github.com/huggingface/transformers/pull/41251) as soon as it’s merged. """ - def __init__(self, config, layer_idx): - super().__init__(config, layer_idx) - self.softmax_scale = self.qk_head_dim**-0.5 - if config.max_seq_len > config.original_seq_len: - mscale = 0.1 * config.mscale * math.log(config.rope_factor) + 1.0 - self.softmax_scale = self.softmax_scale * mscale * mscale + def __init__(self, config: GlmMoeDsaConfig, layer_idx: int): + super().__init__() + self.config = config + self.layer_idx = layer_idx + self.num_key_value_groups = config.num_attention_heads // config.num_key_value_heads + self.attention_dropout = config.attention_dropout + self.num_heads = config.num_attention_heads + + self.q_lora_rank = config.q_lora_rank + self.qk_rope_head_dim = config.qk_rope_head_dim + self.kv_lora_rank = config.kv_lora_rank + self.v_head_dim = config.v_head_dim + self.qk_nope_head_dim = config.qk_nope_head_dim + self.qk_head_dim = config.qk_head_dim + self.index_topk = config.index_topk + + self.is_causal = True + + # Query projection + if self.q_lora_rank is None: + self.q_proj = nn.Linear(config.hidden_size, self.num_heads * self.qk_head_dim, bias=False) + else: + self.q_a_proj = nn.Linear(config.hidden_size, config.q_lora_rank, bias=config.attention_bias) + self.q_a_layernorm = GlmMoeDsaRMSNorm(config.q_lora_rank) + self.q_b_proj = nn.Linear(config.q_lora_rank, self.num_heads * self.qk_head_dim, bias=False) + + # Key-Value projections + self.kv_a_proj_with_mqa = nn.Linear( + config.hidden_size, + self.kv_lora_rank + self.qk_rope_head_dim, + bias=config.attention_bias, + ) + self.kv_a_layernorm = GlmMoeDsaRMSNorm(self.kv_lora_rank) + self.kv_b_proj = nn.Linear( + self.kv_lora_rank, + self.num_heads * (self.qk_nope_head_dim + self.v_head_dim), + bias=False, + ) - self.indexer = GLmMoeDsaIndexer(config, layer_idx) + # Output projection + self.o_proj = nn.Linear( + self.num_heads * self.v_head_dim, + config.hidden_size, + bias=config.attention_bias, + ) + + # Indexer components for sparse attention + self.wq_b = nn.Linear(config.q_lora_rank, self.num_heads * self.qk_head_dim, bias=False) + self.wk = nn.Linear(config.hidden_size, self.qk_head_dim, bias=config.attention_bias) + self.k_norm = GlmMoeDsaRMSNorm(self.qk_head_dim) + self.weights_proj = nn.Linear(config.hidden_size, self.num_heads, bias=False) + + self.scaling = self.qk_head_dim ** (-0.5) + if self.config.rope_parameters.get("rope_type", "default") != "default": + mscale_all_dim = self.config.rope_parameters.get("mscale_all_dim", 0) + scaling_factor = self.config.rope_parameters["factor"] + if mscale_all_dim: + mscale = yarn_get_mscale(scaling_factor, mscale_all_dim) + self.scaling = self.scaling * mscale * mscale def forward( self, - hidden_states: torch.Tensor, # [B, S, hidden] - position_embeddings: tuple[torch.Tensor, torch.Tensor], # (cos, sin) + hidden_states: torch.Tensor, + position_embeddings: tuple[torch.Tensor, torch.Tensor], attention_mask: torch.Tensor | None, - past_key_values: Cache | None = None, # must be Cache with MlaLayer at `layer_idx` + past_key_values: Cache | None = None, cache_position: torch.LongTensor | None = None, - **kwargs, + **kwargs: Unpack[FlashAttentionKwargs], ) -> tuple[torch.Tensor, torch.Tensor | None, tuple[torch.Tensor] | None]: - B, S, _ = hidden_states.shape - cos, sin = position_embeddings + batch_size, seq_length = hidden_states.shape[:-1] + + # For training or when index_topk is not effective, fall back to standard attention + # This is a simplified implementation - in practice, you'd implement the full sparse indexer + if self.training or seq_length <= self.index_topk: + warnings.warn( + "DeepSeek V3.2 sparse attention is not fully implemented in this version. " + "Falling back to standard attention. For production use, please use vLLM or " + "other optimized inference engines.", + UserWarning, + ) + return self._standard_attention( + hidden_states, position_embeddings, attention_mask, past_key_values, cache_position, **kwargs + ) - # ----- Q path ----- - q_resid = self.q_a_layernorm(self.q_a_proj(hidden_states)) # [B, S, q_lora_rank] - q_states = self.q_b_proj(q_resid).view(B, S, self.num_heads, self.qk_head_dim) # [B, S, H, D] - # Split into pass/rot then apply RoPE on q_rot + # Sparse attention implementation would go here + # This requires custom CUDA kernels for efficient top-k selection and indexing + return self._standard_attention( + hidden_states, position_embeddings, attention_mask, past_key_values, cache_position, **kwargs + ) + + def _standard_attention( + self, + hidden_states: torch.Tensor, + position_embeddings: tuple[torch.Tensor, torch.Tensor], + attention_mask: torch.Tensor | None, + past_key_values: Cache | None = None, + cache_position: torch.LongTensor | None = None, + **kwargs: Unpack[FlashAttentionKwargs], + ) -> tuple[torch.Tensor, torch.Tensor | None, tuple[torch.Tensor] | None]: + """Standard attention fallback (same as DeepSeek V3)""" + batch_size, seq_length = hidden_states.shape[:-1] + query_shape = (batch_size, seq_length, -1, self.qk_head_dim) + key_shape = (batch_size, seq_length, -1, self.qk_nope_head_dim + self.v_head_dim) + + if self.q_lora_rank is None: + q_states = self.q_proj(hidden_states) + else: + q_states = self.q_b_proj(self.q_a_layernorm(self.q_a_proj(hidden_states))) + q_states = q_states.view(query_shape).transpose(1, 2) q_pass, q_rot = torch.split(q_states, [self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1) - q_rot = apply_rotary_pos_emb(q_rot, cos, sin) # [B, S, H, rope_D] - q_states = torch.cat([q_pass, q_rot], dim=-1) # [B, S, H, D] - - # Layout for matmul: [B, H, S, D] - q_states = q_states.transpose(1, 2).contiguous() # [B, H, S, D] - - # ----- KV path (compressed + rope stream) ----- - kv_all = self.kv_a_proj_with_mqa(hidden_states) # [B, S, kv_rank + rope_D] - kv_compressed, k_rot = torch.split(kv_all, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1) - kv_compressed = self.kv_a_layernorm(kv_compressed) # [B, S, kv_rank] - # Pre-project to K_pass and V - kv_proj = self.kv_b_proj(kv_compressed) # [B, S, H*(qk_nope + v)] - kv_proj = kv_proj.view(B, S, self.num_heads, self.qk_nope_head_dim + self.v_head_dim) - k_pass, v_states = torch.split( - kv_proj, [self.qk_nope_head_dim, self.v_head_dim], dim=-1 - ) # [B,S,H,nope], [B,S,H,V] - - # Rope on K side: keep a single-head rope stream like MLA, then expand - k_rot = k_rot.view(B, 1, S, self.qk_rope_head_dim) # [B, 1, S, rope_D] - k_rot = apply_rotary_pos_emb(k_rot, cos, sin) # [B, 1, S, rope_D] - - # Concatenate K = [K_pass, K_rot(expanded)] - k_states = torch.cat( - ( - k_pass.transpose(1, 2), # [B, H, S, nope_D] - k_rot.expand(B, self.num_heads, S, -1), - ), # [B, H, S, rope_D] - dim=-1, - ) # [B, H, S, D] - v_states = v_states.transpose(1, 2).contiguous() # [B, H, S, V] - - # ----- Cache update/usage ----- + + compressed_kv = self.kv_a_proj_with_mqa(hidden_states) + k_pass, k_rot = torch.split(compressed_kv, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1) + + k_pass = self.kv_b_proj(self.kv_a_layernorm(k_pass)).view(key_shape).transpose(1, 2) + k_pass, value_states = torch.split(k_pass, [self.qk_nope_head_dim, self.v_head_dim], dim=-1) + + k_rot = k_rot.view(batch_size, 1, seq_length, self.qk_rope_head_dim) + + cos, sin = position_embeddings + if self.config.rope_interleave: + q_rot, k_rot = apply_rotary_pos_emb_interleave(q_rot, k_rot, cos, sin) + else: + q_rot, k_rot = apply_rotary_pos_emb(q_rot, k_rot, cos, sin) + k_rot = k_rot.expand(*k_pass.shape[:-1], -1) + + query_states = torch.cat((q_pass, q_rot), dim=-1) + key_states = torch.cat((k_pass, k_rot), dim=-1) + if past_key_values is not None: - # Store compressed stream & rope stream (as in original MLA path) - # We cache `kv_compressed` under `keys` and `k_rot` under `values` in MlaLayer. - # Shapes must be [B, H, t, *] and [B, 1, t, rope_D]. - kv_comp_cache = kv_compressed.view(B, 1, S, self.kv_lora_rank).expand(B, self.num_heads, S, -1) - k_rot_cache = k_rot # [B, 1, S, rope_D] - cached_kv, cached_pe = past_key_values.update( - kv_comp_cache, k_rot_cache, layer_idx=self.layer_idx, cache_kwargs={"cache_position": cache_position} - ) - # Decode path makes use of cached projections; Prefill can use full K/V directly. - - # ----- Two paths (prefill vs decode) ----- - if attention_mask is not None: - # Prefill (full attention over local window): standard scaled dot-product with top-k pruning from indexer - - # Build scores: [B, H, S, S_total] - # K layout already [B, H, T, D] - scores = (q_states.float() @ k_states.float().transpose(-1, -2)) * self.scaling # [B, H, S, T] - - # Indexer top-k - if past_key_values is not None: - topk_idx = self.indexer( - hidden_states, - q_resid, - position_embeddings, - attention_mask, - past_key_values_index=past_key_values, # we reuse same Cache with IndexerLayer? (separate cache recommended) - cache_position=cache_position, - ) - # Build mask to keep only top-k per (B,S,head?) - # Expect topk_idx shape to broadcast to [B, H, S, T]. We scatter along last dim. - keep_mask = torch.full_like(scores, float("-inf")) - # If topk_idx is [B,S,topk], expand for heads: - if topk_idx.dim() == 3: - topk_idx = topk_idx.unsqueeze(1).expand(B, self.num_heads, S, -1) - keep_mask.scatter_(-1, topk_idx, 0.0) - scores = scores + keep_mask - - probs = nn.functional.softmax(scores, dim=-1, dtype=torch.float32).type_as(hidden_states) # [B, H, S, T] - attn_output = probs @ v_states # [B, H, S, V] - - elif past_key_values is not None: - # Decode: use cached compressed KV & rope stream to recompose attention scores efficiently - # Compose q_pass and q_rot pieces as in MLA math, but via matmul - # 1) Rebuild "nope" term via kv_b weights (dequant on the fly) - wkv_b = self.kv_b_proj.weight.view( - self.num_heads, self.qk_nope_head_dim + self.v_head_dim, self.kv_lora_rank - ) - w_k_nope = wkv_b[:, : self.qk_nope_head_dim, :] # [H, nope_D, kv_rank] - w_v = wkv_b[:, self.qk_nope_head_dim :, :] # [H, V, kv_rank] - - # q_pass: [B,H,S,nope_D]; cached_kv: [B,H,T,kv_rank] - q_pass = q_states[..., : self.qk_nope_head_dim] # [B,H,S,nope_D] - kv_comp = past_key_values[self.layer_idx][0] # keys -> [B,H,T,kv_rank] - pe_full = past_key_values[self.layer_idx][1] # values -> [B,1,T,rope_D] - # Project q_pass with w_k_nope: [B,H,S,kv_rank] - qk_nope = torch.matmul(q_pass, w_k_nope.transpose(-1, -2)) # [B,H,S,kv_rank] - # Scores_nope = qk_nope @ kv_comp^T - scores_nope = torch.matmul(qk_nope.float(), kv_comp.float().transpose(-1, -2)) # [B,H,S,T] - - # 2) Rope term: q_rot @ k_rot^T - q_rot_only = q_states[..., -self.qk_rope_head_dim :] # [B,H,S,rope_D] - k_rot_only = pe_full.expand(B, self.num_heads, -1, -1) # [B,H,T,rope_D] - scores_rot = torch.matmul(q_rot_only.float(), k_rot_only.float().transpose(-1, -2)) # [B,H,S,T] - - scores = (scores_nope + scores_rot) * self.scaling - - # Indexer top-k (decode) - topk_idx = self.indexer( - hidden_states, - q_resid, - position_embeddings, - attention_mask, - past_key_values_index=past_key_values, - cache_position=cache_position, - ) - # For decode single-step S==1 typically; build a [B,H,1,T] mask - keep_mask = torch.full_like(scores, float("-inf")) - if topk_idx.dim() == 3: - topk_idx = topk_idx.unsqueeze(1).expand(B, self.num_heads, S, -1) - keep_mask.scatter_(-1, topk_idx, 0.0) - scores = scores + keep_mask + cache_kwargs = {"sin": sin, "cos": cos, "cache_position": cache_position} + key_states, value_states = past_key_values.update(key_states, value_states, self.layer_idx, cache_kwargs) - probs = nn.functional.softmax(scores, dim=-1, dtype=torch.float32).type_as(hidden_states) # [B,H,S,T] + if self.config._attn_implementation == "flash_attention_2" and self.qk_head_dim != self.v_head_dim: + value_states = F.pad(value_states, [0, self.qk_head_dim - self.v_head_dim]) - # Rebuild V for decode fast-path: v = (kv_comp @ w_v^T) - # kv_comp: [B,H,T,kv_rank], w_v: [H, V, kv_rank] - v_from_comp = torch.matmul(kv_comp, w_v.transpose(-1, -2)) # [B,H,T,V] - attn_output = torch.matmul(probs, v_from_comp) # [B,H,S,V] + attention_interface: Callable = eager_attention_forward + if self.config._attn_implementation != "eager": + attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation] - # Output projection - attn_output = attn_output.transpose(1, 2).reshape(B, S, -1).contiguous() # [B,S,H*V] - attn_output = self.o_proj(attn_output) # [B,S,hidden] - return attn_output, None, None + attn_output, attn_weights = attention_interface( + self, + query_states, + key_states, + value_states, + attention_mask, + dropout=0.0 if not self.training else self.attention_dropout, + scaling=self.scaling, + **kwargs, + ) + if self.config._attn_implementation == "flash_attention_2" and self.qk_head_dim != self.v_head_dim: + attn_output = attn_output[:, :, :, : self.v_head_dim] -class GlmMoeDsaDecoderLayer(Glm4MoeDecoderLayer): - def __init__(self, config: GlmMoeDsaConfig, layer_idx: int): - super().__init__(config, layer_idx) + attn_output = attn_output.reshape(batch_size, seq_length, -1).contiguous() + attn_output = self.o_proj(attn_output) + return attn_output, attn_weights - self.self_attn = GlmMoeDsaAttention(config=config, layer_idx=layer_idx) + +class GlmMoeDsaDecoderLayer(Glm4MoeLiteDecoderLayer): + pass class GlmMoeDsaPreTrainedModel(Glm4MoePreTrainedModel): From 4165f8fd600cfaeb399a2dcbdf1c5c5d1d2c346a Mon Sep 17 00:00:00 2001 From: zRzRzRzRzRzRzR <2448370773@qq.com> Date: Sat, 7 Feb 2026 19:15:46 +0100 Subject: [PATCH 08/19] update --- .../glm_moe_dsa/configuration_glm_moe_dsa.py | 49 +++++++++++- .../models/glm_moe_dsa/modular_glm_moe_dsa.py | 75 +++++++++++++++++-- 2 files changed, 113 insertions(+), 11 deletions(-) diff --git a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py index f62d08ca50b5..940164abd5c6 100644 --- a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py @@ -120,10 +120,6 @@ class GlmMoeDsaConfig(PreTrainedConfig): \--k dense layers--/ index_topk (`int`, *optional*, defaults to 2048): Number of top tokens selected by the indexer for retrieval/attention in each step. - index_head_dim (`int`, *optional*, defaults to 128): - Hidden size (per-head dimension) of each indexer attention head. - index_n_heads (`int`, *optional*, defaults to 32): - Number of attention heads used by the indexer module. ```python >>> from transformers import Glm4MoeLiteModel, Glm4MoeLiteConfig @@ -219,6 +215,51 @@ def __init__( self.index_topk = index_topk self.index_head_dim = index_head_dim self.index_n_heads = index_n_heads + self.mlp_layer_types = mlp_layer_types + self.vocab_size = vocab_size + self.max_position_embeddings = max_position_embeddings + self.hidden_size = hidden_size + self.intermediate_size = intermediate_size + self.num_hidden_layers = num_hidden_layers + + # Default to MoE from the second layer and on + self.mlp_layer_types = mlp_layer_types + if self.mlp_layer_types is None: + self.mlp_layer_types = ["dense"] * self.first_k_dense_replace + ["sparse"] * ( + self.num_hidden_layers - self.first_k_dense_replace + ) + layer_type_validation(self.mlp_layer_types, self.num_hidden_layers, attention=False) + + self.moe_intermediate_size = moe_intermediate_size + self.num_attention_heads = num_attention_heads + self.n_shared_experts = n_shared_experts + self.n_routed_experts = n_routed_experts + self.routed_scaling_factor = routed_scaling_factor + self.kv_lora_rank = kv_lora_rank + self.q_lora_rank = q_lora_rank + self.qk_rope_head_dim = qk_rope_head_dim + self.v_head_dim = v_head_dim + self.qk_nope_head_dim = qk_nope_head_dim + self.qk_head_dim = qk_nope_head_dim + qk_rope_head_dim + self.head_dim = qk_rope_head_dim + self.n_group = n_group + self.topk_group = topk_group + self.num_experts_per_tok = num_experts_per_tok + self.norm_topk_prob = norm_topk_prob + self.rope_interleave = rope_interleave + self.num_key_value_heads = num_key_value_heads + self.hidden_act = hidden_act + self.initializer_range = initializer_range + self.rms_norm_eps = rms_norm_eps + self.pretraining_tp = pretraining_tp + self.use_cache = use_cache + self.attention_bias = attention_bias + self.attention_dropout = attention_dropout + self.rope_parameters = rope_parameters + self.pad_token_id = pad_token_id + self.bos_token_id = bos_token_id + self.eos_token_id = eos_token_id + self.tie_word_embeddings = tie_word_embeddings self.vocab_size = vocab_size self.max_position_embeddings = max_position_embeddings self.hidden_size = hidden_size diff --git a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py index fb90cbad9016..641fec5331df 100644 --- a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py @@ -21,7 +21,9 @@ import torch.nn.functional as F from ...cache_utils import Cache +from ...configuration_utils import layer_type_validation from ...modeling_flash_attention_utils import FlashAttentionKwargs +from ...modeling_rope_utils import RopeParameters from ...modeling_utils import ALL_ATTENTION_FUNCTIONS from ...models.llama.modeling_llama import ( apply_rotary_pos_emb, @@ -140,10 +142,6 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): \--k dense layers--/ index_topk (`int`, *optional*, defaults to 2048): Number of top tokens selected by the indexer for retrieval/attention in each step. - index_head_dim (`int`, *optional*, defaults to 128): - Hidden size (per-head dimension) of each indexer attention head. - index_n_heads (`int`, *optional*, defaults to 32): - Number of attention heads used by the indexer module. ```python >>> from transformers import Glm4MoeLiteModel, Glm4MoeLiteConfig @@ -157,13 +155,13 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): def __init__( self, + vocab_size: int | None = 154880, hidden_size: int | None = 6144, intermediate_size: int | None = 12288, moe_intermediate_size: int | None = 2048, num_hidden_layers: int | None = 78, num_attention_heads: int | None = 64, num_key_value_heads: int | None = 64, - first_k_dense_replace: int | None = 3, n_shared_experts: int | None = 1, n_routed_experts: int | None = 256, routed_scaling_factor: float | None = 2.5, @@ -172,12 +170,30 @@ def __init__( qk_rope_head_dim: int | None = 64, v_head_dim: int | None = 256, qk_nope_head_dim: int | None = 192, + n_group: int | None = 1, + topk_group: int | None = 1, num_experts_per_tok: int | None = 8, + norm_topk_prob: bool | None = True, + hidden_act: str | None = "silu", + max_position_embeddings: int | None = 202752, initializer_range: float | None = 0.02, + rms_norm_eps: int | None = 1e-5, + use_cache: bool | None = True, + pad_token_id: int | None = None, + bos_token_id: int | None = 0, + eos_token_id: int | None = 1, + pretraining_tp: int | None = 1, + tie_word_embeddings: bool | None = False, + rope_parameters: RopeParameters | dict[str, RopeParameters] | None = None, + rope_interleave: bool | None = True, + mlp_layer_types=None, + attention_bias: bool | None = False, + attention_dropout: float | None = 0.0, + first_k_dense_replace: int | None = 3, index_topk: int | None = 2048, index_head_dim: int | None = 128, index_n_heads: int | None = 32, - **super_kwargs, + **kwargs, ): self.hidden_size = hidden_size self.intermediate_size = intermediate_size @@ -201,8 +217,53 @@ def __init__( self.index_topk = index_topk self.index_head_dim = index_head_dim self.index_n_heads = index_n_heads + self.mlp_layer_types = mlp_layer_types + self.vocab_size = vocab_size + self.max_position_embeddings = max_position_embeddings + self.hidden_size = hidden_size + self.intermediate_size = intermediate_size + self.num_hidden_layers = num_hidden_layers - super().__init__(**super_kwargs) + # Default to MoE from the second layer and on + self.mlp_layer_types = mlp_layer_types + if self.mlp_layer_types is None: + self.mlp_layer_types = ["dense"] * self.first_k_dense_replace + ["sparse"] * ( + self.num_hidden_layers - self.first_k_dense_replace + ) + layer_type_validation(self.mlp_layer_types, self.num_hidden_layers, attention=False) + + self.moe_intermediate_size = moe_intermediate_size + self.num_attention_heads = num_attention_heads + self.n_shared_experts = n_shared_experts + self.n_routed_experts = n_routed_experts + self.routed_scaling_factor = routed_scaling_factor + self.kv_lora_rank = kv_lora_rank + self.q_lora_rank = q_lora_rank + self.qk_rope_head_dim = qk_rope_head_dim + self.v_head_dim = v_head_dim + self.qk_nope_head_dim = qk_nope_head_dim + self.qk_head_dim = qk_nope_head_dim + qk_rope_head_dim + self.head_dim = qk_rope_head_dim + self.n_group = n_group + self.topk_group = topk_group + self.num_experts_per_tok = num_experts_per_tok + self.norm_topk_prob = norm_topk_prob + self.rope_interleave = rope_interleave + self.num_key_value_heads = num_key_value_heads + self.hidden_act = hidden_act + self.initializer_range = initializer_range + self.rms_norm_eps = rms_norm_eps + self.pretraining_tp = pretraining_tp + self.use_cache = use_cache + self.attention_bias = attention_bias + self.attention_dropout = attention_dropout + self.rope_parameters = rope_parameters + self.pad_token_id = pad_token_id + self.bos_token_id = bos_token_id + self.eos_token_id = eos_token_id + self.tie_word_embeddings = tie_word_embeddings + + super().__init__(**kwargs) class GlmMoeDsaRMSNorm(Glm4MoeRMSNorm): From 59d1057c8dc7e5475b47bc6c6cb1a89d87f5938d Mon Sep 17 00:00:00 2001 From: zRzRzRzRzRzRzR <2448370773@qq.com> Date: Sun, 8 Feb 2026 07:43:44 +0100 Subject: [PATCH 09/19] 1 --- docs/source/en/model_doc/glm_moe_dsa.md | 2 +- .../models/glm_moe_dsa/configuration_glm_moe_dsa.py | 4 ---- src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py | 4 ---- 3 files changed, 1 insertion(+), 9 deletions(-) diff --git a/docs/source/en/model_doc/glm_moe_dsa.md b/docs/source/en/model_doc/glm_moe_dsa.md index cd863d1205c5..e720e0000f90 100644 --- a/docs/source/en/model_doc/glm_moe_dsa.md +++ b/docs/source/en/model_doc/glm_moe_dsa.md @@ -16,7 +16,7 @@ limitations under the License. ⚠️ Note that this file is in Markdown but contain specific syntax for our doc-builder (similar to MDX) that may not be rendered properly in your Markdown viewer. --> -*This model was released on {release_date} and added to Hugging Face Transformers on 2026-01-28.* +*This model was released on {release_date} and added to Hugging Face Transformers on 2026-02-08.* # GlmMoeDsa diff --git a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py index 940164abd5c6..f0cef67072fa 100644 --- a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py @@ -189,8 +189,6 @@ def __init__( attention_dropout: float | None = 0.0, first_k_dense_replace: int | None = 3, index_topk: int | None = 2048, - index_head_dim: int | None = 128, - index_n_heads: int | None = 32, **kwargs, ): self.hidden_size = hidden_size @@ -213,8 +211,6 @@ def __init__( self.num_key_value_heads = num_key_value_heads self.initializer_range = initializer_range self.index_topk = index_topk - self.index_head_dim = index_head_dim - self.index_n_heads = index_n_heads self.mlp_layer_types = mlp_layer_types self.vocab_size = vocab_size self.max_position_embeddings = max_position_embeddings diff --git a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py index 641fec5331df..d199015327d6 100644 --- a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py @@ -191,8 +191,6 @@ def __init__( attention_dropout: float | None = 0.0, first_k_dense_replace: int | None = 3, index_topk: int | None = 2048, - index_head_dim: int | None = 128, - index_n_heads: int | None = 32, **kwargs, ): self.hidden_size = hidden_size @@ -215,8 +213,6 @@ def __init__( self.num_key_value_heads = num_key_value_heads self.initializer_range = initializer_range self.index_topk = index_topk - self.index_head_dim = index_head_dim - self.index_n_heads = index_n_heads self.mlp_layer_types = mlp_layer_types self.vocab_size = vocab_size self.max_position_embeddings = max_position_embeddings From b1bdd1ae976f4ed0fbb71f057c92c36f66c22da9 Mon Sep 17 00:00:00 2001 From: zRzRzRzRzRzRzR <2448370773@qq.com> Date: Mon, 9 Feb 2026 09:38:44 +0100 Subject: [PATCH 10/19] update --- .../models/glm_moe_dsa/configuration_glm_moe_dsa.py | 9 +-------- .../models/glm_moe_dsa/modular_glm_moe_dsa.py | 9 +-------- 2 files changed, 2 insertions(+), 16 deletions(-) diff --git a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py index f0cef67072fa..14080b225fb5 100644 --- a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py @@ -115,9 +115,6 @@ class GlmMoeDsaConfig(PreTrainedConfig): Whether to use a bias in the query, key, value and output projection layers during self-attention. attention_dropout (`float`, *optional*, defaults to 0.0): The dropout ratio for the attention probabilities. - first_k_dense_replace (`int`, *optional*, defaults to 3): - Number of dense layers in shallow layers(embed->dense->dense->...->dense->moe->moe...->lm_head). - \--k dense layers--/ index_topk (`int`, *optional*, defaults to 2048): Number of top tokens selected by the indexer for retrieval/attention in each step. @@ -187,7 +184,6 @@ def __init__( mlp_layer_types=None, attention_bias: bool | None = False, attention_dropout: float | None = 0.0, - first_k_dense_replace: int | None = 3, index_topk: int | None = 2048, **kwargs, ): @@ -203,7 +199,6 @@ def __init__( self.q_lora_rank = q_lora_rank self.qk_rope_head_dim = qk_rope_head_dim self.v_head_dim = v_head_dim - self.first_k_dense_replace = first_k_dense_replace self.qk_nope_head_dim = qk_nope_head_dim self.qk_head_dim = qk_nope_head_dim + qk_rope_head_dim self.head_dim = qk_rope_head_dim @@ -221,9 +216,7 @@ def __init__( # Default to MoE from the second layer and on self.mlp_layer_types = mlp_layer_types if self.mlp_layer_types is None: - self.mlp_layer_types = ["dense"] * self.first_k_dense_replace + ["sparse"] * ( - self.num_hidden_layers - self.first_k_dense_replace - ) + self.mlp_layer_types = ["dense"] * 3 + ["sparse"] * (self.num_hidden_layers - 3) layer_type_validation(self.mlp_layer_types, self.num_hidden_layers, attention=False) self.moe_intermediate_size = moe_intermediate_size diff --git a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py index d199015327d6..d042d8e8f654 100644 --- a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py @@ -137,9 +137,6 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): Whether to use a bias in the query, key, value and output projection layers during self-attention. attention_dropout (`float`, *optional*, defaults to 0.0): The dropout ratio for the attention probabilities. - first_k_dense_replace (`int`, *optional*, defaults to 3): - Number of dense layers in shallow layers(embed->dense->dense->...->dense->moe->moe...->lm_head). - \--k dense layers--/ index_topk (`int`, *optional*, defaults to 2048): Number of top tokens selected by the indexer for retrieval/attention in each step. @@ -189,7 +186,6 @@ def __init__( mlp_layer_types=None, attention_bias: bool | None = False, attention_dropout: float | None = 0.0, - first_k_dense_replace: int | None = 3, index_topk: int | None = 2048, **kwargs, ): @@ -205,7 +201,6 @@ def __init__( self.q_lora_rank = q_lora_rank self.qk_rope_head_dim = qk_rope_head_dim self.v_head_dim = v_head_dim - self.first_k_dense_replace = first_k_dense_replace self.qk_nope_head_dim = qk_nope_head_dim self.qk_head_dim = qk_nope_head_dim + qk_rope_head_dim self.head_dim = qk_rope_head_dim @@ -223,9 +218,7 @@ def __init__( # Default to MoE from the second layer and on self.mlp_layer_types = mlp_layer_types if self.mlp_layer_types is None: - self.mlp_layer_types = ["dense"] * self.first_k_dense_replace + ["sparse"] * ( - self.num_hidden_layers - self.first_k_dense_replace - ) + self.mlp_layer_types = ["dense"] * 3 + ["sparse"] * (self.num_hidden_layers - 3) layer_type_validation(self.mlp_layer_types, self.num_hidden_layers, attention=False) self.moe_intermediate_size = moe_intermediate_size From 057c247b2da3a11b93e5f5843322d57ee08d08ba Mon Sep 17 00:00:00 2001 From: Cyril Vallez Date: Mon, 9 Feb 2026 11:31:59 +0100 Subject: [PATCH 11/19] fix attention and date --- .../models/glm_moe_dsa/configuration_glm_moe_dsa.py | 6 +++--- .../models/glm_moe_dsa/modeling_glm_moe_dsa.py | 12 ++++++------ .../models/glm_moe_dsa/modular_glm_moe_dsa.py | 6 +++--- .../models/glm_moe_dsa/test_modeling_glm_moe_dsa.py | 2 +- 4 files changed, 13 insertions(+), 13 deletions(-) diff --git a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py index 14080b225fb5..981842f124a3 100644 --- a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py @@ -132,9 +132,9 @@ class GlmMoeDsaConfig(PreTrainedConfig): keys_to_ignore_at_inference = ["past_key_values"] base_model_tp_plan = { "layers.*.self_attn.o_proj": "rowwise", - "layers.*.mlp.experts.gate_up_proj": "local_rowwise", - "layers.*.mlp.experts.down_proj": "local_rowwise", - "layers.*.mlp.experts": "gather", + "layers.*.mlp.experts.gate_up_proj": "packed_colwise", + "layers.*.mlp.experts.down_proj": "rowwise", + "layers.*.mlp.experts": "moe_tp_experts", "layers.*.mlp.gate_proj": "colwise", "layers.*.mlp.up_proj": "colwise", "layers.*.mlp.down_proj": "rowwise", diff --git a/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py index c47688bc1a0e..8da9ece18a90 100644 --- a/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py @@ -47,7 +47,7 @@ @use_kernel_forward_from_hub("RMSNorm") class GlmMoeDsaRMSNorm(nn.Module): - def __init__(self, hidden_size, eps=1e-6): + def __init__(self, hidden_size, eps: float = 1e-6) -> None: """ GlmMoeDsaRMSNorm is equivalent to T5LayerNorm """ @@ -55,7 +55,7 @@ def __init__(self, hidden_size, eps=1e-6): self.weight = nn.Parameter(torch.ones(hidden_size)) self.variance_epsilon = eps - def forward(self, hidden_states): + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: input_dtype = hidden_states.dtype hidden_states = hidden_states.to(torch.float32) variance = hidden_states.pow(2).mean(-1, keepdim=True) @@ -329,9 +329,9 @@ def _standard_attention( if self.config._attn_implementation == "flash_attention_2" and self.qk_head_dim != self.v_head_dim: value_states = F.pad(value_states, [0, self.qk_head_dim - self.v_head_dim]) - attention_interface: Callable = eager_attention_forward - if self.config._attn_implementation != "eager": - attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation] + attention_interface: Callable = ALL_ATTENTION_FUNCTIONS.get_interface( + self.config._attn_implementation, eager_attention_forward + ) attn_output, attn_weights = attention_interface( self, @@ -715,7 +715,7 @@ def forward( @auto_docstring class GlmMoeDsaForCausalLM(GlmMoeDsaPreTrainedModel, GenerationMixin): _tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"} - _tp_plan = {"lm_head": "colwise_rep"} + _tp_plan = {"lm_head": "colwise_gather_output"} _pp_plan = {"lm_head": (["hidden_states"], ["logits"])} def __init__(self, config): diff --git a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py index d042d8e8f654..bb54f9e37d58 100644 --- a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py @@ -407,9 +407,9 @@ def _standard_attention( if self.config._attn_implementation == "flash_attention_2" and self.qk_head_dim != self.v_head_dim: value_states = F.pad(value_states, [0, self.qk_head_dim - self.v_head_dim]) - attention_interface: Callable = eager_attention_forward - if self.config._attn_implementation != "eager": - attention_interface = ALL_ATTENTION_FUNCTIONS[self.config._attn_implementation] + attention_interface: Callable = ALL_ATTENTION_FUNCTIONS.get_interface( + self.config._attn_implementation, eager_attention_forward + ) attn_output, attn_weights = attention_interface( self, diff --git a/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py b/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py index bfcc63ec9268..69476101bcb6 100644 --- a/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py +++ b/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py @@ -1,4 +1,4 @@ -# Copyright 2025 the HuggingFace Team. All rights reserved. +# Copyright 2026 the HuggingFace Team. All rights reserved. # # Licensed under the Apache License, Version 2.0 (the "License"); # you may not use this file except in compliance with the License. From 1d96f00ba7f6b0b73391a54384dd39bc5665be79 Mon Sep 17 00:00:00 2001 From: Cyril Vallez Date: Mon, 9 Feb 2026 11:48:29 +0100 Subject: [PATCH 12/19] remove pretraining_tp and improve tests --- .../models/glm_moe_dsa/modular_glm_moe_dsa.py | 13 +++---------- .../models/glm_moe_dsa/test_modeling_glm_moe_dsa.py | 7 +++++-- 2 files changed, 8 insertions(+), 12 deletions(-) diff --git a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py index bb54f9e37d58..f3a4c9d84e44 100644 --- a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py @@ -118,11 +118,6 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): Beginning of stream token id. eos_token_id (`int`, *optional*, defaults to 1): End of stream token id. - pretraining_tp (`int`, *optional*, defaults to 1): - Experimental feature. Tensor parallelism rank used during pretraining. Please refer to [this - document](https://huggingface.co/docs/transformers/parallelism) to understand more about it. This value is - necessary to ensure exact reproducibility of the pretraining results. Please refer to [this - issue](https://github.com/pytorch/pytorch/issues/76232). tie_word_embeddings (`bool`, *optional*, defaults to `False`): Whether to tie weight embeddings rope_parameters (`RopeParameters`, *optional*): @@ -179,7 +174,6 @@ def __init__( pad_token_id: int | None = None, bos_token_id: int | None = 0, eos_token_id: int | None = 1, - pretraining_tp: int | None = 1, tie_word_embeddings: bool | None = False, rope_parameters: RopeParameters | dict[str, RopeParameters] | None = None, rope_interleave: bool | None = True, @@ -242,7 +236,6 @@ def __init__( self.hidden_act = hidden_act self.initializer_range = initializer_range self.rms_norm_eps = rms_norm_eps - self.pretraining_tp = pretraining_tp self.use_cache = use_cache self.attention_bias = attention_bias self.attention_dropout = attention_dropout @@ -266,9 +259,9 @@ class GlmMoeDsaAttention(nn.Module): This implements the native sparse attention from [DeepSeek V3.2](https://huggingface.co/deepseek-ai/DeepSeek-V3.2) which uses an indexer to select top-k tokens for attention computation, making it more efficient for long sequences. - In GLM-5, the indexer RoPE uses neox_style = false. Therefore, we introduced the indexer_rope_interleave parameter: when indexer_rope_interleave is set to True, RoPE is computed using the same neox_style = false behavior as in the GLMMOEDSA model. This part has not yet been implemented in transformers. - - Switch to the implementation from this [PR](https://github.com/huggingface/transformers/pull/41251) as soon as it’s merged. + In GLM-5, the indexer RoPE uses neox_style = false. Therefore, we introduced the indexer_rope_interleave parameter: + when indexer_rope_interleave is set to True, RoPE is computed using the same neox_style = false behavior as in the + GlmMoeDsa model. This part has not yet been implemented in transformers. """ def __init__(self, config: GlmMoeDsaConfig, layer_idx: int): diff --git a/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py b/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py index 69476101bcb6..6e98fdfb5828 100644 --- a/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py +++ b/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py @@ -11,7 +11,7 @@ # WITHOUT WARRANTIES OR CONDITIONS OF ANY KIND, either express or implied. # See the License for the specific language governing permissions and # limitations under the License. -"""Testing suite for the PyTorch GLM-4.5, GLM-4.6, GLM-4.7 model.""" +"""Testing suite for the PyTorch GlmMoeDsa model.""" import unittest @@ -48,14 +48,17 @@ def __init__( qk_nope_head_dim=64, qk_rope_head_dim=64, v_head_dim=128, + num_hidden_layers=2, + mlp_layer_types=["sparse", "mlp"], # sparse is MoE here... Not sure why it was not called moe... ): - super().__init__(parent=parent) + super().__init__(parent=parent, num_hidden_layers=num_hidden_layers) self.n_routed_experts = n_routed_experts self.kv_lora_rank = kv_lora_rank self.q_lora_rank = q_lora_rank self.qk_nope_head_dim = qk_nope_head_dim self.qk_rope_head_dim = qk_rope_head_dim self.v_head_dim = v_head_dim + self.mlp_layer_types = mlp_layer_types @require_torch From b42d474221dd4b32db7b654fd2a35d020dca4062 Mon Sep 17 00:00:00 2001 From: Cyril Vallez Date: Mon, 9 Feb 2026 11:50:15 +0100 Subject: [PATCH 13/19] style --- .../models/glm_moe_dsa/configuration_glm_moe_dsa.py | 7 ------- .../models/glm_moe_dsa/modeling_glm_moe_dsa.py | 6 +++--- tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py | 2 +- 3 files changed, 4 insertions(+), 11 deletions(-) diff --git a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py index 981842f124a3..9a6125aab5ee 100644 --- a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py @@ -96,11 +96,6 @@ class GlmMoeDsaConfig(PreTrainedConfig): Beginning of stream token id. eos_token_id (`int`, *optional*, defaults to 1): End of stream token id. - pretraining_tp (`int`, *optional*, defaults to 1): - Experimental feature. Tensor parallelism rank used during pretraining. Please refer to [this - document](https://huggingface.co/docs/transformers/parallelism) to understand more about it. This value is - necessary to ensure exact reproducibility of the pretraining results. Please refer to [this - issue](https://github.com/pytorch/pytorch/issues/76232). tie_word_embeddings (`bool`, *optional*, defaults to `False`): Whether to tie weight embeddings rope_parameters (`RopeParameters`, *optional*): @@ -177,7 +172,6 @@ def __init__( pad_token_id: int | None = None, bos_token_id: int | None = 0, eos_token_id: int | None = 1, - pretraining_tp: int | None = 1, tie_word_embeddings: bool | None = False, rope_parameters: RopeParameters | dict[str, RopeParameters] | None = None, rope_interleave: bool | None = True, @@ -240,7 +234,6 @@ def __init__( self.hidden_act = hidden_act self.initializer_range = initializer_range self.rms_norm_eps = rms_norm_eps - self.pretraining_tp = pretraining_tp self.use_cache = use_cache self.attention_bias = attention_bias self.attention_dropout = attention_dropout diff --git a/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py index 8da9ece18a90..b33eca495f75 100644 --- a/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py @@ -188,9 +188,9 @@ class GlmMoeDsaAttention(nn.Module): This implements the native sparse attention from [DeepSeek V3.2](https://huggingface.co/deepseek-ai/DeepSeek-V3.2) which uses an indexer to select top-k tokens for attention computation, making it more efficient for long sequences. - In GLM-5, the indexer RoPE uses neox_style = false. Therefore, we introduced the indexer_rope_interleave parameter: when indexer_rope_interleave is set to True, RoPE is computed using the same neox_style = false behavior as in the GLMMOEDSA model. This part has not yet been implemented in transformers. - - Switch to the implementation from this [PR](https://github.com/huggingface/transformers/pull/41251) as soon as it’s merged. + In GLM-5, the indexer RoPE uses neox_style = false. Therefore, we introduced the indexer_rope_interleave parameter: + when indexer_rope_interleave is set to True, RoPE is computed using the same neox_style = false behavior as in the + GlmMoeDsa model. This part has not yet been implemented in transformers. """ def __init__(self, config: GlmMoeDsaConfig, layer_idx: int): diff --git a/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py b/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py index 6e98fdfb5828..f58232fb6f14 100644 --- a/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py +++ b/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py @@ -49,7 +49,7 @@ def __init__( qk_rope_head_dim=64, v_head_dim=128, num_hidden_layers=2, - mlp_layer_types=["sparse", "mlp"], # sparse is MoE here... Not sure why it was not called moe... + mlp_layer_types=["sparse", "dense"], ): super().__init__(parent=parent, num_hidden_layers=num_hidden_layers) self.n_routed_experts = n_routed_experts From 3728b91398d2c04f4ef2c08fcdbb9c5df734997a Mon Sep 17 00:00:00 2001 From: Cyril Vallez Date: Mon, 9 Feb 2026 11:53:30 +0100 Subject: [PATCH 14/19] remove pretraining_tp --- src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py | 1 - src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py | 1 + 2 files changed, 1 insertion(+), 1 deletion(-) diff --git a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py index 9a6125aab5ee..5116e84a0c5e 100644 --- a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py @@ -275,7 +275,6 @@ def __init__( self.hidden_act = hidden_act self.initializer_range = initializer_range self.rms_norm_eps = rms_norm_eps - self.pretraining_tp = pretraining_tp self.use_cache = use_cache self.attention_bias = attention_bias self.attention_dropout = attention_dropout diff --git a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py index f3a4c9d84e44..9003db0aa50d 100644 --- a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py @@ -246,6 +246,7 @@ def __init__( self.tie_word_embeddings = tie_word_embeddings super().__init__(**kwargs) + del self.pretraining_tp class GlmMoeDsaRMSNorm(Glm4MoeRMSNorm): From 6387f2daef3fb13b74f6203f8e1415ccb10517e9 Mon Sep 17 00:00:00 2001 From: Cyril Vallez Date: Mon, 9 Feb 2026 12:02:52 +0100 Subject: [PATCH 15/19] fix config --- .../models/glm_moe_dsa/configuration_glm_moe_dsa.py | 6 ++++-- src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py | 6 ++++-- 2 files changed, 8 insertions(+), 4 deletions(-) diff --git a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py index 5116e84a0c5e..eab56dd5a7e2 100644 --- a/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/configuration_glm_moe_dsa.py @@ -207,10 +207,12 @@ def __init__( self.intermediate_size = intermediate_size self.num_hidden_layers = num_hidden_layers - # Default to MoE from the second layer and on + # Default to MoE from the fourth layer and on self.mlp_layer_types = mlp_layer_types if self.mlp_layer_types is None: - self.mlp_layer_types = ["dense"] * 3 + ["sparse"] * (self.num_hidden_layers - 3) + self.mlp_layer_types = ["dense"] * min(3, self.num_hidden_layers) + ["sparse"] * ( + self.num_hidden_layers - 3 + ) layer_type_validation(self.mlp_layer_types, self.num_hidden_layers, attention=False) self.moe_intermediate_size = moe_intermediate_size diff --git a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py index 9003db0aa50d..d2fcf09c533f 100644 --- a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py @@ -209,10 +209,12 @@ def __init__( self.intermediate_size = intermediate_size self.num_hidden_layers = num_hidden_layers - # Default to MoE from the second layer and on + # Default to MoE from the fourth layer and on self.mlp_layer_types = mlp_layer_types if self.mlp_layer_types is None: - self.mlp_layer_types = ["dense"] * 3 + ["sparse"] * (self.num_hidden_layers - 3) + self.mlp_layer_types = ["dense"] * min(3, self.num_hidden_layers) + ["sparse"] * ( + self.num_hidden_layers - 3 + ) layer_type_validation(self.mlp_layer_types, self.num_hidden_layers, attention=False) self.moe_intermediate_size = moe_intermediate_size From 9706ab153af05df24b9d7362323571dbf133cabc Mon Sep 17 00:00:00 2001 From: Cyril Vallez Date: Mon, 9 Feb 2026 12:07:30 +0100 Subject: [PATCH 16/19] fix compile --- .../glm_moe_dsa/modeling_glm_moe_dsa.py | 19 +++++++++++-------- .../models/glm_moe_dsa/modular_glm_moe_dsa.py | 14 +++++++------- 2 files changed, 18 insertions(+), 15 deletions(-) diff --git a/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py index b33eca495f75..4105eaec328a 100644 --- a/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/modeling_glm_moe_dsa.py @@ -20,7 +20,6 @@ import math -import warnings from collections.abc import Callable from typing import Optional @@ -40,11 +39,15 @@ from ...modeling_rope_utils import ROPE_INIT_FUNCTIONS, dynamic_rope_update from ...modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel from ...processing_utils import Unpack -from ...utils import TransformersKwargs, auto_docstring, can_return_tuple, is_grouped_mm_available +from ...utils import TransformersKwargs, auto_docstring, can_return_tuple, is_grouped_mm_available, logging from ...utils.generic import check_model_inputs, maybe_autocast +from ...utils.import_utils import is_tracing from .configuration_glm_moe_dsa import GlmMoeDsaConfig +logger = logging.get_logger(__name__) + + @use_kernel_forward_from_hub("RMSNorm") class GlmMoeDsaRMSNorm(nn.Module): def __init__(self, hidden_size, eps: float = 1e-6) -> None: @@ -267,12 +270,12 @@ def forward( # For training or when index_topk is not effective, fall back to standard attention # This is a simplified implementation - in practice, you'd implement the full sparse indexer if self.training or seq_length <= self.index_topk: - warnings.warn( - "DeepSeek V3.2 sparse attention is not fully implemented in this version. " - "Falling back to standard attention. For production use, please use vLLM or " - "other optimized inference engines.", - UserWarning, - ) + if not is_tracing(hidden_states): + logger.warning_once( + "DeepSeek V3.2 sparse attention is not fully implemented in this version. " + "Falling back to standard attention. For production use, please use vLLM or " + "other optimized inference engines.", + ) return self._standard_attention( hidden_states, position_embeddings, attention_mask, past_key_values, cache_position, **kwargs ) diff --git a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py index d2fcf09c533f..9fb17188fa6f 100644 --- a/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py +++ b/src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py @@ -13,7 +13,6 @@ # limitations under the License. -import warnings from collections.abc import Callable import torch @@ -30,6 +29,7 @@ ) from ...processing_utils import Unpack from ...utils import logging +from ...utils.import_utils import is_tracing from ..deepseek_v3.modeling_deepseek_v3 import apply_rotary_pos_emb_interleave, yarn_get_mscale from ..glm4_moe.modeling_glm4_moe import ( Glm4MoeForCausalLM, @@ -341,12 +341,12 @@ def forward( # For training or when index_topk is not effective, fall back to standard attention # This is a simplified implementation - in practice, you'd implement the full sparse indexer if self.training or seq_length <= self.index_topk: - warnings.warn( - "DeepSeek V3.2 sparse attention is not fully implemented in this version. " - "Falling back to standard attention. For production use, please use vLLM or " - "other optimized inference engines.", - UserWarning, - ) + if not is_tracing(hidden_states): + logger.warning_once( + "DeepSeek V3.2 sparse attention is not fully implemented in this version. " + "Falling back to standard attention. For production use, please use vLLM or " + "other optimized inference engines.", + ) return self._standard_attention( hidden_states, position_embeddings, attention_mask, past_key_values, cache_position, **kwargs ) From 2134ffda4905fdcaed4589a65a01cc062b9f6ccf Mon Sep 17 00:00:00 2001 From: Cyril Vallez Date: Mon, 9 Feb 2026 12:20:47 +0100 Subject: [PATCH 17/19] remove wrong integration test --- .../glm_moe_dsa/test_modeling_glm_moe_dsa.py | 48 +------------------ 1 file changed, 1 insertion(+), 47 deletions(-) diff --git a/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py b/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py index f58232fb6f14..3c408ef9b604 100644 --- a/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py +++ b/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py @@ -17,7 +17,6 @@ import pytest import torch -from packaging import version from transformers import Cache, is_torch_available from transformers.testing_utils import ( @@ -88,49 +87,4 @@ def _check_past_key_values_for_generate(self, batch_size, past_key_values, seq_l @require_torch_accelerator @slow class GlmMoeDsaIntegrationTest(unittest.TestCase): - def tearDown(self): - # See LlamaIntegrationTest.tearDown(). Can be removed once LlamaIntegrationTest.tearDown() is removed. - cleanup(torch_device, gc_collect=False) - - @slow - @require_torch_accelerator - @pytest.mark.torch_compile_test - def test_compile_static_cache(self): - # `torch==2.2` will throw an error on this test (as in other compilation tests), but torch==2.1.2 and torch>2.2 - # work as intended. See https://github.com/pytorch/pytorch/issues/121943 - if version.parse(torch.__version__) < version.parse("2.3.0"): - self.skipTest(reason="This test requires torch >= 2.3 to run.") - - NUM_TOKENS_TO_GENERATE = 40 - EXPECTED_TEXT_COMPLETION = [ - 'hello, world!\'\'\')\nprint(\'hello, world!\')\nprint("hello, world!")\nprint("hello, world!")\nprint("hello, world!")\nprint("hello, world!")\nprint("hello, world!")\n', - "tell me the story of the first Thanksgiving. commonly known as the Pilgrims, arrived in the autumn of 1620. They were seeking religious freedom and a new life in the Plymouth Colony. Their first", - ] - - prompts = ["[gMASK]hello", "[gMASK]tell me"] - tokenizer = AutoTokenizer.from_pretrained("zai-org/GLM-4.7-Flash") - model = GlmMoeDsaForCausalLM.from_pretrained( - "zai-org/GLM-4.7-Flash", device_map=torch_device, dtype=torch.bfloat16 - ) - inputs = tokenizer(prompts, return_tensors="pt", padding=True).to(model.device) - - # Dynamic Cache - generated_ids = model.generate(**inputs, max_new_tokens=NUM_TOKENS_TO_GENERATE, do_sample=False) - dynamic_text = tokenizer.batch_decode(generated_ids, skip_special_tokens=True) - self.assertEqual(EXPECTED_TEXT_COMPLETION, dynamic_text) - - # Static Cache - generated_ids = model.generate( - **inputs, max_new_tokens=NUM_TOKENS_TO_GENERATE, do_sample=False, cache_implementation="static" - ) - static_text = tokenizer.batch_decode(generated_ids, skip_special_tokens=True) - self.assertEqual(EXPECTED_TEXT_COMPLETION, static_text) - - # Static Cache + compile - model._cache = None # clear cache object, initialized when we pass `cache_implementation="static"` - model.forward = torch.compile(model.forward, mode="reduce-overhead", fullgraph=True) - generated_ids = model.generate( - **inputs, max_new_tokens=NUM_TOKENS_TO_GENERATE, do_sample=False, cache_implementation="static" - ) - static_compiled_text = tokenizer.batch_decode(generated_ids, skip_special_tokens=True) - self.assertEqual(EXPECTED_TEXT_COMPLETION, static_compiled_text) + pass From 88d3c266431a4d28aaf9d069663708c3da5045c9 Mon Sep 17 00:00:00 2001 From: Cyril Vallez Date: Mon, 9 Feb 2026 12:26:24 +0100 Subject: [PATCH 18/19] fix --- tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py | 7 +------ 1 file changed, 1 insertion(+), 6 deletions(-) diff --git a/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py b/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py index 3c408ef9b604..3ba79a7c96db 100644 --- a/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py +++ b/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py @@ -15,23 +15,18 @@ import unittest -import pytest -import torch - from transformers import Cache, is_torch_available from transformers.testing_utils import ( - cleanup, require_torch, require_torch_accelerator, slow, - torch_device, ) from ...causal_lm_tester import CausalLMModelTest, CausalLMModelTester if is_torch_available(): - from transformers import AutoTokenizer, GlmMoeDsaForCausalLM, GlmMoeDsaModel + from transformers import GlmMoeDsaModel class GlmMoeDsaModelTester(CausalLMModelTester): From c436a18a74b3cc4efc8fd7a9e70ff865af509b22 Mon Sep 17 00:00:00 2001 From: Cyril Vallez Date: Mon, 9 Feb 2026 12:29:21 +0100 Subject: [PATCH 19/19] better --- tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py | 3 ++- 1 file changed, 2 insertions(+), 1 deletion(-) diff --git a/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py b/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py index 3ba79a7c96db..c3a26d62d392 100644 --- a/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py +++ b/tests/models/glm_moe_dsa/test_modeling_glm_moe_dsa.py @@ -26,12 +26,13 @@ if is_torch_available(): - from transformers import GlmMoeDsaModel + from transformers import GlmMoeDsaForCausalLM, GlmMoeDsaModel class GlmMoeDsaModelTester(CausalLMModelTester): if is_torch_available(): base_model_class = GlmMoeDsaModel + causal_lm_class = GlmMoeDsaForCausalLM def __init__( self,