From 67a9994b5e45c4583e30549eb42f2c5f0b8d2fd9 Mon Sep 17 00:00:00 2001 From: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> Date: Tue, 2 Jun 2026 08:00:34 +0000 Subject: [PATCH 1/8] Fix error attributes mapping Signed-off-by: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> --- 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, 2 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 8f11f42794b3..06e170dfd00e 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 @@ -80,7 +80,6 @@ class GlmMoeDsaConfig(PreTrainedConfig): } attribute_map = { "num_local_experts": "n_routed_experts", - "head_dim": "qk_rope_head_dim", } vocab_size: int = 154880 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 d909bb97e704..ce5d6808484a 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 @@ -140,7 +140,6 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): indexer_types: list[str] | None = None attribute_map = { "num_local_experts": "n_routed_experts", - "head_dim": "qk_rope_head_dim", } def __post_init__(self, **kwargs): From 3f02fef4e282a17e2a4228a9c257209ad70fac55 Mon Sep 17 00:00:00 2001 From: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> Date: Tue, 2 Jun 2026 10:45:31 +0000 Subject: [PATCH 2/8] Fix tests Signed-off-by: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> --- .../glm_moe_dsa/modeling_glm_moe_dsa.py | 14 ++++++------- .../models/glm_moe_dsa/modular_glm_moe_dsa.py | 21 +++++++++++++++++++ .../glm_moe_dsa/test_modeling_glm_moe_dsa.py | 9 ++++++++ 3 files changed, 36 insertions(+), 8 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 736dcdce32c3..d5bc5e4aa08b 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 @@ -19,7 +19,6 @@ # limitations under the License. from collections.abc import Callable -from typing import Optional import torch import torch.nn as nn @@ -697,7 +696,7 @@ def __init__(self, config: GlmMoeDsaConfig, device=None): @staticmethod def compute_default_rope_parameters( config: GlmMoeDsaConfig | None = None, - device: Optional["torch.device"] = None, + device=None, seq_len: int | None = None, ) -> tuple["torch.Tensor", float]: """ @@ -714,15 +713,14 @@ def compute_default_rope_parameters( 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) + head_dim = config.qk_rope_head_dim + attention_factor = 1.0 - attention_factor = 1.0 # Unused in this type of RoPE + if head_dim == 0: + return torch.empty(0, device=device), attention_factor - # Compute the inverse frequencies inv_freq = 1.0 / ( - base ** (torch.arange(0, dim, 2, dtype=torch.int64).to(device=device, dtype=torch.float) / dim) + base ** (torch.arange(0, head_dim, 2, dtype=torch.int64).to(device=device, dtype=torch.float) / head_dim) ) return inv_freq, attention_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 ce5d6808484a..7937be10e606 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 @@ -34,6 +34,7 @@ Glm4MoeModel, Glm4MoePreTrainedModel, Glm4MoeRMSNorm, + Glm4MoeRotaryEmbedding, ) from ..glm4_moe_lite.configuration_glm4_moe_lite import Glm4MoeLiteConfig from ..glm4_moe_lite.modeling_glm4_moe_lite import ( @@ -539,6 +540,26 @@ class GlmMoeDsaPreTrainedModel(Glm4MoePreTrainedModel): _compatible_flash_implementations = ["kernels-community/flash-mla"] +class GlmMoeDsaRotaryEmbedding(Glm4MoeRotaryEmbedding): + @staticmethod + def compute_default_rope_parameters( + config: GlmMoeDsaConfig | None = None, + device=None, + seq_len: int | None = None, + ): + base = config.rope_parameters["rope_theta"] + head_dim = config.qk_rope_head_dim + attention_factor = 1.0 + + if head_dim == 0: + return torch.empty(0, device=device), attention_factor + + inv_freq = 1.0 / ( + base ** (torch.arange(0, head_dim, 2, dtype=torch.int64).to(device=device, dtype=torch.float) / head_dim) + ) + return inv_freq, attention_factor + + class GlmMoeDsaModel(Glm4MoeModel): def forward( 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 9a1e26800eec..8b3a734af01a 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 @@ -156,6 +156,15 @@ def test_generate_compilation_all_outputs(self): def test_generate_with_static_cache(self): pass + @unittest.skip("GLM-MoE-DSA uses qk_rope_head_dim; generic rope scaling tests assume config.head_dim") + def test_model_rope_scaling_frequencies(self): + pass + + @parameterized.expand([("linear",), ("dynamic",), ("yarn",)]) + @unittest.skip("GLM-MoE-DSA uses qk_rope_head_dim; generic rope scaling tests assume config.head_dim") + def test_model_rope_scaling_from_config(self, scaling_type): + pass + @require_torch_accelerator @slow From 23d789576a5d947f2c18a1707cc80ac35dbd4f4f Mon Sep 17 00:00:00 2001 From: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> Date: Wed, 3 Jun 2026 08:14:40 +0000 Subject: [PATCH 3/8] Fix: Support interleaved RoPE for MLA Signed-off-by: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> --- .../glm_moe_dsa/configuration_glm_moe_dsa.py | 57 +++++--- .../glm_moe_dsa/modeling_glm_moe_dsa.py | 100 +++++++++++--- .../models/glm_moe_dsa/modular_glm_moe_dsa.py | 124 ++++++++++++------ .../glm_moe_dsa/test_modeling_glm_moe_dsa.py | 7 + utils/check_config_attributes.py | 1 + 5 files changed, 211 insertions(+), 78 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 06e170dfd00e..6e8b797436c5 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 @@ -31,16 +31,28 @@ class GlmMoeDsaConfig(PreTrainedConfig): r""" n_group (`int`, *optional*, defaults to 1): Number of groups for routed experts. + rope_interleave (`bool`, *optional*, defaults to `True`): + Whether main MLA rotary embeddings use interleaved pair layout. mlp_layer_types (`list`, *optional*): - MLP type pattern for each layer (`"dense"` or `"sparse"`). Defaults to 3 dense + rest sparse. + MLP type pattern for each layer (`"dense"` or `"sparse"`). Defaults to `3` dense layers and then every `moe_layer_freq`-th layer sparse. + moe_layer_freq (`int`, *optional*, defaults to 1): + Frequency for sparse MoE layers. index_topk (`int`, *optional*, defaults to 2048): Number of top tokens selected by the indexer for sparse attention. index_head_dim (`int`, *optional*, defaults to 128): Head dimension for the indexer projections (DSA). index_n_heads (`int | None`, *optional*, defaults to 32): Number of heads for the indexer projections (DSA). + index_topk_freq (`int`, *optional*, defaults to 1): + Frequency for full indexer recomputation when `index_topk_pattern` is not provided. + index_topk_pattern (`str | list[str]`, *optional*): + Explicit full/shared indexer pattern using `"F"`/`"S"` or `"full"`/`"shared"` values. + index_skip_topk_offset (`int`, *optional*, defaults to 2): + Offset used with `index_topk_freq` to decide which layers recompute top-k indices. + indexer_rope_interleave (`bool`, *optional*, defaults to `False`): + Whether DSA indexer rotary embeddings use interleaved pair layout. indexer_types (`list[str]`, *optional*): - Indexer mode for each layer (`"full"` or `"shared"`). Defaults to first layer full, then every `index_topk_freq`-th layer full, rest shared. + Indexer mode for each layer (`"full"` or `"shared"`). Defaults to the pattern derived from `index_topk_freq` and `index_skip_topk_offset`. ```python >>> from transformers import GlmMoeDsaConfig, GlmMoeDsaModel @@ -60,24 +72,27 @@ class GlmMoeDsaConfig(PreTrainedConfig): base_model_tp_plan = { "layers.*.self_attn.q_b_proj": "colwise", - "layers.*.self_attn.kv_a_proj_with_mqa": "mla_kv_a_proj", "layers.*.self_attn.kv_b_proj": "colwise", - "layers.*.self_attn.o_proj": "rowwise", - "layers.*.mlp.experts.gate_up_proj": "packed_colwise", - "layers.*.mlp.experts.down_proj": "rowwise", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.self_attn.o_proj": "rowwise_allreduce", + "layers.*.mlp.experts": "moe_experts_allreduce", "layers.*.mlp.shared_experts.gate_proj": "colwise", "layers.*.mlp.shared_experts.up_proj": "colwise", - "layers.*.mlp.shared_experts.down_proj": "rowwise", + "layers.*.mlp.shared_experts.down_proj": "rowwise_allreduce", "layers.*.mlp.gate_proj": "colwise", "layers.*.mlp.up_proj": "colwise", - "layers.*.mlp.down_proj": "rowwise", + "layers.*.mlp.down_proj": "rowwise_allreduce", } base_model_pp_plan = { "embed_tokens": (["input_ids"], ["inputs_embeds"]), "layers": (["hidden_states", "attention_mask"], ["hidden_states"]), "norm": (["hidden_states"], ["hidden_states"]), } + + base_model_fsdp_plan = { + "embed_tokens": "free_full_weight", + "layers.*": "free_full_weight", + "norm": "keep_full_weight", + } attribute_map = { "num_local_experts": "n_routed_experts", } @@ -110,37 +125,41 @@ class GlmMoeDsaConfig(PreTrainedConfig): pad_token_id: int | None = None bos_token_id: int | None = 0 eos_token_id: int | list[int] | None = 1 + pretraining_tp: int = 1 tie_word_embeddings: bool = False rope_parameters: RopeParameters | dict | None = None + rope_interleave: bool = True mlp_layer_types: list[str] | None = None attention_bias: bool = False attention_dropout: float | int = 0.0 index_topk: int = 2048 index_head_dim: int = 128 index_n_heads: int = 32 + moe_layer_freq: int = 1 + index_topk_freq: int = 1 + index_topk_pattern: str | list[str] | None = None + index_skip_topk_offset: int = 2 + indexer_rope_interleave: bool = False indexer_types: list[str] | None = None def __post_init__(self, **kwargs): self.qk_head_dim = self.qk_nope_head_dim + self.qk_rope_head_dim - - # MLP layer types: first 3 dense, rest sparse if self.mlp_layer_types is None: - self.mlp_layer_types = ["dense"] * min(3, self.num_hidden_layers) + ["sparse"] * ( - self.num_hidden_layers - 3 - ) + self.mlp_layer_types = [ + "sparse" if i >= 3 and i % self.moe_layer_freq == 0 else "dense" for i in range(self.num_hidden_layers) + ] - # Indexer layer types if self.indexer_types is None: - pattern = kwargs.pop("index_topk_pattern", None) - freq = kwargs.pop("index_topk_freq", 1) + pattern = self.index_topk_pattern if pattern is not None: self.indexer_types = ( [{"F": "full", "S": "shared"}[c] for c in pattern] if isinstance(pattern, str) else list(pattern) ) else: - # First layer full, then every freq-th layer full, rest shared + freq = max(self.index_topk_freq, 1) + offset = self.index_skip_topk_offset self.indexer_types = [ - "full" if (max(i - 1, 0) % freq) == 0 else "shared" for i in range(self.num_hidden_layers) + "full" if (max(i - offset + 1, 0) % freq) == 0 else "shared" for i in range(self.num_hidden_layers) ] super().__post_init__(**kwargs) 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 d5bc5e4aa08b..fa86845ee055 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 @@ -77,9 +77,9 @@ def apply_rotary_pos_emb( unsqueeze_dim: int = 1, ) -> torch.Tensor: """ - Applies Rotary Position Embedding to a single tensor. + Applies (non-interleaved, NeoX/Llama style) Rotary Position Embedding to a single tensor. - This is the transformers equivalent of DeepSeek V3.2's `apply_rotary_emb(x, freqs_cis, interleaved)`. + This is the transformers equivalent of DeepSeek V3.2's `apply_rotary_emb(x, freqs_cis, interleaved=False)`. Instead of using complex-number `freqs_cis`, we use pre-split `(cos, sin)` tensors from RotaryEmbedding. Args: @@ -94,11 +94,21 @@ def apply_rotary_pos_emb( """ cos = cos.unsqueeze(unsqueeze_dim) sin = sin.unsqueeze(unsqueeze_dim) + return (x * cos) + (rotate_half(x) * sin) - # Split-half (NeoX/Llama style): (x[:d/2], x[d/2:]) - # This matches llama's apply_rotary_pos_emb logic. - x_rotated = (x * cos) + (rotate_half(x) * sin) - return x_rotated + +def apply_rotary_pos_emb_interleave_single( + x: torch.Tensor, + cos: torch.Tensor, + sin: torch.Tensor, + unsqueeze_dim: int = 1, +) -> torch.Tensor: + """Interleaved (GPT-J style) RoPE applied to a single tensor (the indexer's q/k stream).""" + cos = cos.unsqueeze(unsqueeze_dim) + sin = sin.unsqueeze(unsqueeze_dim) + *leading_dims, head_dim = x.shape + x = x.view(*leading_dims, head_dim // 2, 2).transpose(-1, -2).reshape(*leading_dims, head_dim) + return (x * cos) + (rotate_half(x) * sin) class GlmMoeDsaIndexer(nn.Module): @@ -106,8 +116,8 @@ class GlmMoeDsaIndexer(nn.Module): DeepSeek Sparse Attention (DSA) indexer for selecting top-k tokens. The Indexer has its own lightweight projections (wq_b, wk) separate from the - main MLA attention. It uses non-interleaved (NeoX/Llama) RoPE, unlike the main attention - which uses interleaved RoPE. + main MLA attention. RoPE layout (interleaved vs non-interleaved) is controlled + independently from the main attention by `config.indexer_rope_interleave`. **Cache strategy**: The Indexer manages its own key cache (`_cached_keys`) separately from the DynamicCache used by MLA attention, since DynamicCache is sized for exactly @@ -140,6 +150,12 @@ def __init__(self, config: "GlmMoeDsaConfig", layer_idx: int): # Indexer maintains its own key cache (not in DynamicCache, which is sized for attention layers only) self.register_buffer("_cached_keys", None, persistent=False) + def _apply_indexer_rope(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: + """Apply the indexer's RoPE (interleaved or not) to a `[B, S, H, rope_D]` tensor.""" + if self.config.indexer_rope_interleave: + return apply_rotary_pos_emb_interleave_single(x, cos, sin, unsqueeze_dim=2) + return apply_rotary_pos_emb(x, cos, sin, unsqueeze_dim=2) + @torch.no_grad() def forward( self, @@ -177,13 +193,13 @@ def forward( q = self.wq_b(q_resid) # [B, S, H*D] q = q.view(batch_size, seq_len, self.n_heads, self.head_dim) # [B, S, H, D] q_pe, q_nope = torch.split(q, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1) - q_pe = apply_rotary_pos_emb(q_pe, cos, sin, unsqueeze_dim=2) # [B, S, H, rope_D] + q_pe = self._apply_indexer_rope(q_pe, cos, sin) # [B, S, H, rope_D] q = torch.cat([q_pe, q_nope], dim=-1) # [B, S, H, D] # === Keys === k = self.k_norm(self.wk(hidden_states)) # [B, S, D] k_pe, k_nope = torch.split(k, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1) - k_pe = apply_rotary_pos_emb(k_pe.unsqueeze(2), cos, sin, unsqueeze_dim=2).squeeze(2) # [B, S, rope_D] + k_pe = self._apply_indexer_rope(k_pe.unsqueeze(2), cos, sin).squeeze(2) # [B, S, rope_D] k = torch.cat([k_pe, k_nope], dim=-1) # [B, S, D] # === Key cache (managed by the indexer, not DynamicCache) === @@ -228,6 +244,44 @@ def forward( 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. + + 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 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, @@ -281,6 +335,9 @@ class GlmMoeDsaAttention(nn.Module): and SDPA backends. The reference's compressed-cache decode path (which avoids the kv_b_proj expansion at decode time) is a future optimization that would require a dedicated MLA cache class. + **DSA layer sharing**: full layers run the indexer; shared layers reuse the previous full + layer's top-k indices (`skip_topk` / `next_skip_topk`, derived from `config.indexer_types`). + **FP8 compatibility**: all weight accesses use standard nn.Linear forward calls (never raw `.weight` access), so FP8-quantized checkpoints work transparently. """ @@ -332,15 +389,15 @@ def __init__(self, config: GlmMoeDsaConfig, layer_idx: int): self.scaling = self.qk_head_dim ** (-0.5) - self.indexer = GlmMoeDsaIndexer(config, layer_idx) - # Refer: https://arxiv.org/abs/2603.12201 for more details. - # skip_topk: when True, this layer will skip computation and reuse previous layer's topk indices. - # next_skip_topk: when True, the next layer will skip computation and reuse this layer's topk indices. + # skip_topk: when True, this layer reuses the previous full indexer layer's top-k indices. + # next_skip_topk: when True, the next layer reuses this layer's top-k indices. + # Shared layers have no indexer of their own. self.skip_topk = config.indexer_types[layer_idx] == "shared" self.next_skip_topk = ( config.indexer_types[layer_idx + 1] == "shared" if layer_idx < len(config.indexer_types) - 1 else False ) + self.indexer = None if self.skip_topk else GlmMoeDsaIndexer(config, layer_idx) def forward( self, @@ -364,7 +421,6 @@ def forward( query_states = query_states.view(batch_size, seq_length, -1, self.qk_head_dim).transpose(1, 2) # Split nope/rope, apply RoPE, recombine — layout: [B, H, S, D] q_nope, q_pe = torch.split(query_states, [self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1) - q_pe = apply_rotary_pos_emb(q_pe, cos, sin, unsqueeze_dim=1) # BHSD format # ===== KV path ===== compressed_kv = self.kv_a_proj_with_mqa(hidden_states) # [B, S, kv_rank + rope_D] @@ -378,9 +434,13 @@ def forward( k_nope = k_nope.transpose(1, 2) # [B, H, S, nope_D] value_states = value_states.transpose(1, 2) # [B, H, S, v_D] - # RoPE on k_pe (single-head rope stream) + # RoPE on q_pe / k_pe (single-head rope stream for k) k_pe = k_pe.view(batch_size, 1, seq_length, self.qk_rope_head_dim) # [B, 1, S, rope_D] - k_pe = apply_rotary_pos_emb(k_pe, cos, sin, unsqueeze_dim=1) # BHSD format + if self.config.rope_interleave: + q_pe, k_pe = apply_rotary_pos_emb_interleave(q_pe, k_pe, cos, sin) + else: + q_pe = apply_rotary_pos_emb(q_pe, cos, sin, unsqueeze_dim=1) + k_pe = apply_rotary_pos_emb(k_pe, cos, sin, unsqueeze_dim=1) k_pe = k_pe.expand(-1, k_nope.shape[1], -1, -1) # [B, H, S, rope_D] # Assemble full Q and K @@ -394,6 +454,8 @@ def forward( # ===== Indexer (DSA sparse mask) ===== # attention_mask is [B, 1, S, T] (4D) for eager and (2D) otherwise but indexer works with [B, S, T] (3D) if not self.skip_topk or prev_topk_indices is None: + if self.indexer is None: + raise ValueError("Shared DSA layers require top-k indices from a previous full indexer layer.") indexer_mask = ( attention_mask[:, 0, :, :] if attention_mask is not None and attention_mask.dim() == 4 @@ -819,8 +881,10 @@ def forward( @auto_docstring class GlmMoeDsaForCausalLM(GlmMoeDsaPreTrainedModel, GenerationMixin): _tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"} - _tp_plan = {"lm_head": "colwise_gather_output"} + _tp_plan = {"lm_head": "colwise_allgather"} + _sp_plan = {"lm_head": "colwise_loss_parallel"} _pp_plan = {"lm_head": (["hidden_states"], ["logits"])} + _fsdp_plan = {"lm_head": "keep_full_weight"} def __init__(self, config): 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 7937be10e606..19866cc79c5d 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 @@ -29,6 +29,7 @@ from ...processing_utils import Unpack from ...utils import TransformersKwargs, auto_docstring, logging from ...utils.generic import is_flash_attention_requested +from ..deepseek_v3.modeling_deepseek_v3 import apply_rotary_pos_emb_interleave from ..glm4_moe.modeling_glm4_moe import ( Glm4MoeForCausalLM, Glm4MoeModel, @@ -53,9 +54,9 @@ def apply_rotary_pos_emb( unsqueeze_dim: int = 1, ) -> torch.Tensor: """ - Applies Rotary Position Embedding to a single tensor. + Applies (non-interleaved, NeoX/Llama style) Rotary Position Embedding to a single tensor. - This is the transformers equivalent of DeepSeek V3.2's `apply_rotary_emb(x, freqs_cis, interleaved)`. + This is the transformers equivalent of DeepSeek V3.2's `apply_rotary_emb(x, freqs_cis, interleaved=False)`. Instead of using complex-number `freqs_cis`, we use pre-split `(cos, sin)` tensors from RotaryEmbedding. Args: @@ -70,11 +71,21 @@ def apply_rotary_pos_emb( """ cos = cos.unsqueeze(unsqueeze_dim) sin = sin.unsqueeze(unsqueeze_dim) + return (x * cos) + (rotate_half(x) * sin) - # Split-half (NeoX/Llama style): (x[:d/2], x[d/2:]) - # This matches llama's apply_rotary_pos_emb logic. - x_rotated = (x * cos) + (rotate_half(x) * sin) - return x_rotated + +def apply_rotary_pos_emb_interleave_single( + x: torch.Tensor, + cos: torch.Tensor, + sin: torch.Tensor, + unsqueeze_dim: int = 1, +) -> torch.Tensor: + """Interleaved (GPT-J style) RoPE applied to a single tensor (the indexer's q/k stream).""" + cos = cos.unsqueeze(unsqueeze_dim) + sin = sin.unsqueeze(unsqueeze_dim) + *leading_dims, head_dim = x.shape + x = x.view(*leading_dims, head_dim // 2, 2).transpose(-1, -2).reshape(*leading_dims, head_dim) + return (x * cos) + (rotate_half(x) * sin) @auto_docstring(checkpoint="zai-org/GLM-5") @@ -83,16 +94,28 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): r""" n_group (`int`, *optional*, defaults to 1): Number of groups for routed experts. + rope_interleave (`bool`, *optional*, defaults to `True`): + Whether main MLA rotary embeddings use interleaved pair layout. mlp_layer_types (`list`, *optional*): - MLP type pattern for each layer (`"dense"` or `"sparse"`). Defaults to 3 dense + rest sparse. + MLP type pattern for each layer (`"dense"` or `"sparse"`). Defaults to `3` dense layers and then every `moe_layer_freq`-th layer sparse. + moe_layer_freq (`int`, *optional*, defaults to 1): + Frequency for sparse MoE layers. index_topk (`int`, *optional*, defaults to 2048): Number of top tokens selected by the indexer for sparse attention. index_head_dim (`int`, *optional*, defaults to 128): Head dimension for the indexer projections (DSA). index_n_heads (`int | None`, *optional*, defaults to 32): Number of heads for the indexer projections (DSA). + index_topk_freq (`int`, *optional*, defaults to 1): + Frequency for full indexer recomputation when `index_topk_pattern` is not provided. + index_topk_pattern (`str | list[str]`, *optional*): + Explicit full/shared indexer pattern using `"F"`/`"S"` or `"full"`/`"shared"` values. + index_skip_topk_offset (`int`, *optional*, defaults to 2): + Offset used with `index_topk_freq` to decide which layers recompute top-k indices. + indexer_rope_interleave (`bool`, *optional*, defaults to `False`): + Whether DSA indexer rotary embeddings use interleaved pair layout. indexer_types (`list[str]`, *optional*): - Indexer mode for each layer (`"full"` or `"shared"`). Defaults to first layer full, then every `index_topk_freq`-th layer full, rest shared. + Indexer mode for each layer (`"full"` or `"shared"`). Defaults to the pattern derived from `index_topk_freq` and `index_skip_topk_offset`. ```python >>> from transformers import GlmMoeDsaConfig, GlmMoeDsaModel @@ -109,18 +132,24 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): base_model_tp_plan = { "layers.*.self_attn.q_b_proj": "colwise", - "layers.*.self_attn.kv_a_proj_with_mqa": "mla_kv_a_proj", "layers.*.self_attn.kv_b_proj": "colwise", - "layers.*.self_attn.o_proj": "rowwise", - "layers.*.mlp.experts.gate_up_proj": "packed_colwise", - "layers.*.mlp.experts.down_proj": "rowwise", - "layers.*.mlp.experts": "moe_tp_experts", + "layers.*.self_attn.o_proj": "rowwise_allreduce", + "layers.*.mlp.experts": "moe_experts_allreduce", "layers.*.mlp.shared_experts.gate_proj": "colwise", "layers.*.mlp.shared_experts.up_proj": "colwise", - "layers.*.mlp.shared_experts.down_proj": "rowwise", + "layers.*.mlp.shared_experts.down_proj": "rowwise_allreduce", "layers.*.mlp.gate_proj": "colwise", "layers.*.mlp.up_proj": "colwise", - "layers.*.mlp.down_proj": "rowwise", + "layers.*.mlp.down_proj": "rowwise_allreduce", + } + + base_model_fsdp_plan = { + "embed_tokens": "free_full_weight", + "layers.*": "free_full_weight", + "norm": "keep_full_weight", + } + attribute_map = { + "num_local_experts": "n_routed_experts", } hidden_size: int = 6144 @@ -136,34 +165,33 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): index_topk: int = 2048 index_head_dim: int = 128 index_n_heads: int = 32 - pretraining_tp = AttributeError() - rope_interleave = AttributeError() + moe_layer_freq: int = 1 + pretraining_tp: int = 1 + rope_interleave: bool = True + index_topk_freq: int = 1 + index_topk_pattern: str | list[str] | None = None + index_skip_topk_offset: int = 2 + indexer_rope_interleave: bool = False indexer_types: list[str] | None = None - attribute_map = { - "num_local_experts": "n_routed_experts", - } def __post_init__(self, **kwargs): self.qk_head_dim = self.qk_nope_head_dim + self.qk_rope_head_dim - - # MLP layer types: first 3 dense, rest sparse if self.mlp_layer_types is None: - self.mlp_layer_types = ["dense"] * min(3, self.num_hidden_layers) + ["sparse"] * ( - self.num_hidden_layers - 3 - ) + self.mlp_layer_types = [ + "sparse" if i >= 3 and i % self.moe_layer_freq == 0 else "dense" for i in range(self.num_hidden_layers) + ] - # Indexer layer types if self.indexer_types is None: - pattern = kwargs.pop("index_topk_pattern", None) - freq = kwargs.pop("index_topk_freq", 1) + pattern = self.index_topk_pattern if pattern is not None: self.indexer_types = ( [{"F": "full", "S": "shared"}[c] for c in pattern] if isinstance(pattern, str) else list(pattern) ) else: - # First layer full, then every freq-th layer full, rest shared + freq = max(self.index_topk_freq, 1) + offset = self.index_skip_topk_offset self.indexer_types = [ - "full" if (max(i - 1, 0) % freq) == 0 else "shared" for i in range(self.num_hidden_layers) + "full" if (max(i - offset + 1, 0) % freq) == 0 else "shared" for i in range(self.num_hidden_layers) ] PreTrainedConfig.__post_init__(self, **kwargs) @@ -177,8 +205,8 @@ class GlmMoeDsaIndexer(nn.Module): DeepSeek Sparse Attention (DSA) indexer for selecting top-k tokens. The Indexer has its own lightweight projections (wq_b, wk) separate from the - main MLA attention. It uses non-interleaved (NeoX/Llama) RoPE, unlike the main attention - which uses interleaved RoPE. + main MLA attention. RoPE layout (interleaved vs non-interleaved) is controlled + independently from the main attention by `config.indexer_rope_interleave`. **Cache strategy**: The Indexer manages its own key cache (`_cached_keys`) separately from the DynamicCache used by MLA attention, since DynamicCache is sized for exactly @@ -211,6 +239,12 @@ def __init__(self, config: "GlmMoeDsaConfig", layer_idx: int): # Indexer maintains its own key cache (not in DynamicCache, which is sized for attention layers only) self.register_buffer("_cached_keys", None, persistent=False) + def _apply_indexer_rope(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: + """Apply the indexer's RoPE (interleaved or not) to a `[B, S, H, rope_D]` tensor.""" + if self.config.indexer_rope_interleave: + return apply_rotary_pos_emb_interleave_single(x, cos, sin, unsqueeze_dim=2) + return apply_rotary_pos_emb(x, cos, sin, unsqueeze_dim=2) + @torch.no_grad() def forward( self, @@ -248,13 +282,13 @@ def forward( q = self.wq_b(q_resid) # [B, S, H*D] q = q.view(batch_size, seq_len, self.n_heads, self.head_dim) # [B, S, H, D] q_pe, q_nope = torch.split(q, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1) - q_pe = apply_rotary_pos_emb(q_pe, cos, sin, unsqueeze_dim=2) # [B, S, H, rope_D] + q_pe = self._apply_indexer_rope(q_pe, cos, sin) # [B, S, H, rope_D] q = torch.cat([q_pe, q_nope], dim=-1) # [B, S, H, D] # === Keys === k = self.k_norm(self.wk(hidden_states)) # [B, S, D] k_pe, k_nope = torch.split(k, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1) - k_pe = apply_rotary_pos_emb(k_pe.unsqueeze(2), cos, sin, unsqueeze_dim=2).squeeze(2) # [B, S, rope_D] + k_pe = self._apply_indexer_rope(k_pe.unsqueeze(2), cos, sin).squeeze(2) # [B, S, rope_D] k = torch.cat([k_pe, k_nope], dim=-1) # [B, S, D] # === Key cache (managed by the indexer, not DynamicCache) === @@ -315,6 +349,9 @@ class GlmMoeDsaAttention(nn.Module): and SDPA backends. The reference's compressed-cache decode path (which avoids the kv_b_proj expansion at decode time) is a future optimization that would require a dedicated MLA cache class. + **DSA layer sharing**: full layers run the indexer; shared layers reuse the previous full + layer's top-k indices (`skip_topk` / `next_skip_topk`, derived from `config.indexer_types`). + **FP8 compatibility**: all weight accesses use standard nn.Linear forward calls (never raw `.weight` access), so FP8-quantized checkpoints work transparently. """ @@ -366,15 +403,15 @@ def __init__(self, config: GlmMoeDsaConfig, layer_idx: int): self.scaling = self.qk_head_dim ** (-0.5) - self.indexer = GlmMoeDsaIndexer(config, layer_idx) - # Refer: https://arxiv.org/abs/2603.12201 for more details. - # skip_topk: when True, this layer will skip computation and reuse previous layer's topk indices. - # next_skip_topk: when True, the next layer will skip computation and reuse this layer's topk indices. + # skip_topk: when True, this layer reuses the previous full indexer layer's top-k indices. + # next_skip_topk: when True, the next layer reuses this layer's top-k indices. + # Shared layers have no indexer of their own. self.skip_topk = config.indexer_types[layer_idx] == "shared" self.next_skip_topk = ( config.indexer_types[layer_idx + 1] == "shared" if layer_idx < len(config.indexer_types) - 1 else False ) + self.indexer = None if self.skip_topk else GlmMoeDsaIndexer(config, layer_idx) def forward( self, @@ -398,7 +435,6 @@ def forward( query_states = query_states.view(batch_size, seq_length, -1, self.qk_head_dim).transpose(1, 2) # Split nope/rope, apply RoPE, recombine — layout: [B, H, S, D] q_nope, q_pe = torch.split(query_states, [self.qk_nope_head_dim, self.qk_rope_head_dim], dim=-1) - q_pe = apply_rotary_pos_emb(q_pe, cos, sin, unsqueeze_dim=1) # BHSD format # ===== KV path ===== compressed_kv = self.kv_a_proj_with_mqa(hidden_states) # [B, S, kv_rank + rope_D] @@ -412,9 +448,13 @@ def forward( k_nope = k_nope.transpose(1, 2) # [B, H, S, nope_D] value_states = value_states.transpose(1, 2) # [B, H, S, v_D] - # RoPE on k_pe (single-head rope stream) + # RoPE on q_pe / k_pe (single-head rope stream for k) k_pe = k_pe.view(batch_size, 1, seq_length, self.qk_rope_head_dim) # [B, 1, S, rope_D] - k_pe = apply_rotary_pos_emb(k_pe, cos, sin, unsqueeze_dim=1) # BHSD format + if self.config.rope_interleave: + q_pe, k_pe = apply_rotary_pos_emb_interleave(q_pe, k_pe, cos, sin) + else: + q_pe = apply_rotary_pos_emb(q_pe, cos, sin, unsqueeze_dim=1) + k_pe = apply_rotary_pos_emb(k_pe, cos, sin, unsqueeze_dim=1) k_pe = k_pe.expand(-1, k_nope.shape[1], -1, -1) # [B, H, S, rope_D] # Assemble full Q and K @@ -428,6 +468,8 @@ def forward( # ===== Indexer (DSA sparse mask) ===== # attention_mask is [B, 1, S, T] (4D) for eager and (2D) otherwise but indexer works with [B, S, T] (3D) if not self.skip_topk or prev_topk_indices is None: + if self.indexer is None: + raise ValueError("Shared DSA layers require top-k indices from a previous full indexer layer.") indexer_mask = ( attention_mask[:, 0, :, :] if attention_mask is not None and attention_mask.dim() == 4 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 8b3a734af01a..e52c9aec4f13 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 @@ -103,6 +103,13 @@ def test_default_mlp_layer_types(self): config.mlp_layer_types, ["dense", "dense", "dense", "sparse", "sparse", "sparse", "sparse", "sparse"] ) + def test_indexer_types_respect_skip_topk_offset(self): + config = GlmMoeDsaConfig(num_hidden_layers=8, index_topk_freq=4, index_skip_topk_offset=3) + self.assertEqual( + config.indexer_types, + ["full", "full", "full", "shared", "shared", "shared", "full", "shared"], + ) + @parameterized.expand(TEST_EAGER_MATCHES_SDPA_INFERENCE_PARAMETERIZATION) @unittest.skip("Won't fix: Blip2 + T5 backbone needs custom input preparation for this test") def test_eager_matches_sdpa_inference(self, *args): diff --git a/utils/check_config_attributes.py b/utils/check_config_attributes.py index 8cd992d541ff..fc86d51c21e4 100644 --- a/utils/check_config_attributes.py +++ b/utils/check_config_attributes.py @@ -109,6 +109,7 @@ "Cohere2MoeConfig": ["rope_scaling", "sliding_window_pattern"], "CsmConfig": ["tie_codebooks_embeddings"], "DeepseekV2Config": ["norm_topk_prob"], + "GlmMoeDsaConfig": ["index_skip_topk_offset", "index_topk_freq", "index_topk_pattern", "moe_layer_freq"], "DeepseekV4Config": [ # All BC / config-compat surface that the modeling code never reads but # checkpoints in the wild expose (so we keep accepting them in `__init__`): From 8eab52b50196864b13a81e491b10f4167b98bcd5 Mon Sep 17 00:00:00 2001 From: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> Date: Wed, 3 Jun 2026 09:52:55 +0000 Subject: [PATCH 4/8] Fix tests Signed-off-by: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> --- .../glm_moe_dsa/configuration_glm_moe_dsa.py | 23 +++++++++++-------- .../glm_moe_dsa/modeling_glm_moe_dsa.py | 4 +--- .../models/glm_moe_dsa/modular_glm_moe_dsa.py | 11 +++++---- 3 files changed, 21 insertions(+), 17 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 6e8b797436c5..82f17dc4b21b 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 @@ -72,27 +72,24 @@ class GlmMoeDsaConfig(PreTrainedConfig): base_model_tp_plan = { "layers.*.self_attn.q_b_proj": "colwise", + "layers.*.self_attn.kv_a_proj_with_mqa": "mla_kv_a_proj", "layers.*.self_attn.kv_b_proj": "colwise", - "layers.*.self_attn.o_proj": "rowwise_allreduce", - "layers.*.mlp.experts": "moe_experts_allreduce", + "layers.*.self_attn.o_proj": "rowwise", + "layers.*.mlp.experts.gate_up_proj": "packed_colwise", + "layers.*.mlp.experts.down_proj": "rowwise", + "layers.*.mlp.experts": "moe_tp_experts", "layers.*.mlp.shared_experts.gate_proj": "colwise", "layers.*.mlp.shared_experts.up_proj": "colwise", - "layers.*.mlp.shared_experts.down_proj": "rowwise_allreduce", + "layers.*.mlp.shared_experts.down_proj": "rowwise", "layers.*.mlp.gate_proj": "colwise", "layers.*.mlp.up_proj": "colwise", - "layers.*.mlp.down_proj": "rowwise_allreduce", + "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"]), } - - base_model_fsdp_plan = { - "embed_tokens": "free_full_weight", - "layers.*": "free_full_weight", - "norm": "keep_full_weight", - } attribute_map = { "num_local_experts": "n_routed_experts", } @@ -132,6 +129,12 @@ class GlmMoeDsaConfig(PreTrainedConfig): mlp_layer_types: list[str] | None = None attention_bias: bool = False attention_dropout: float | int = 0.0 + + base_model_fsdp_plan = { + "embed_tokens": "free_full_weight", + "layers.*": "free_full_weight", + "norm": "keep_full_weight", + } index_topk: int = 2048 index_head_dim: int = 128 index_n_heads: int = 32 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 fa86845ee055..6cd8f810a76a 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 @@ -881,10 +881,8 @@ def forward( @auto_docstring class GlmMoeDsaForCausalLM(GlmMoeDsaPreTrainedModel, GenerationMixin): _tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"} - _tp_plan = {"lm_head": "colwise_allgather"} - _sp_plan = {"lm_head": "colwise_loss_parallel"} + _tp_plan = {"lm_head": "colwise_gather_output"} _pp_plan = {"lm_head": (["hidden_states"], ["logits"])} - _fsdp_plan = {"lm_head": "keep_full_weight"} def __init__(self, config): 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 19866cc79c5d..a83b5a4561c0 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 @@ -132,15 +132,18 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): base_model_tp_plan = { "layers.*.self_attn.q_b_proj": "colwise", + "layers.*.self_attn.kv_a_proj_with_mqa": "mla_kv_a_proj", "layers.*.self_attn.kv_b_proj": "colwise", - "layers.*.self_attn.o_proj": "rowwise_allreduce", - "layers.*.mlp.experts": "moe_experts_allreduce", + "layers.*.self_attn.o_proj": "rowwise", + "layers.*.mlp.experts.gate_up_proj": "packed_colwise", + "layers.*.mlp.experts.down_proj": "rowwise", + "layers.*.mlp.experts": "moe_tp_experts", "layers.*.mlp.shared_experts.gate_proj": "colwise", "layers.*.mlp.shared_experts.up_proj": "colwise", - "layers.*.mlp.shared_experts.down_proj": "rowwise_allreduce", + "layers.*.mlp.shared_experts.down_proj": "rowwise", "layers.*.mlp.gate_proj": "colwise", "layers.*.mlp.up_proj": "colwise", - "layers.*.mlp.down_proj": "rowwise_allreduce", + "layers.*.mlp.down_proj": "rowwise", } base_model_fsdp_plan = { From aa8928c5b98b7940dd6622128b6681b85da43ec5 Mon Sep 17 00:00:00 2001 From: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> Date: Thu, 4 Jun 2026 11:47:49 +0000 Subject: [PATCH 5/8] Remove redundant fields Signed-off-by: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> --- 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, 2 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 82f17dc4b21b..f9242f4651b7 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 @@ -122,7 +122,6 @@ class GlmMoeDsaConfig(PreTrainedConfig): pad_token_id: int | None = None bos_token_id: int | None = 0 eos_token_id: int | list[int] | None = 1 - pretraining_tp: int = 1 tie_word_embeddings: bool = False rope_parameters: RopeParameters | dict | None = None rope_interleave: bool = True 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 a83b5a4561c0..948a7ab5f974 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 @@ -169,7 +169,6 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): index_head_dim: int = 128 index_n_heads: int = 32 moe_layer_freq: int = 1 - pretraining_tp: int = 1 rope_interleave: bool = True index_topk_freq: int = 1 index_topk_pattern: str | list[str] | None = None From 45f7dcde598f5a0a9b5cbb7d4bcefea14b435c50 Mon Sep 17 00:00:00 2001 From: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> Date: Thu, 4 Jun 2026 11:57:49 +0000 Subject: [PATCH 6/8] Fix tests Signed-off-by: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> --- src/transformers/models/glm_moe_dsa/modular_glm_moe_dsa.py | 1 + 1 file changed, 1 insertion(+) 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 948a7ab5f974..b14c1a7f1b7d 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 @@ -175,6 +175,7 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): index_skip_topk_offset: int = 2 indexer_rope_interleave: bool = False indexer_types: list[str] | None = None + pretraining_tp = AttributeError() def __post_init__(self, **kwargs): self.qk_head_dim = self.qk_nope_head_dim + self.qk_rope_head_dim From 7da480ed82dc3a1bb9e613c45636e3050c50f798 Mon Sep 17 00:00:00 2001 From: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> Date: Thu, 4 Jun 2026 14:30:29 +0000 Subject: [PATCH 7/8] Remove portentially codepath Signed-off-by: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> --- .../glm_moe_dsa/configuration_glm_moe_dsa.py | 9 +-- .../glm_moe_dsa/modeling_glm_moe_dsa.py | 81 +++---------------- .../models/glm_moe_dsa/modular_glm_moe_dsa.py | 54 ++++--------- utils/check_config_attributes.py | 8 +- 4 files changed, 38 insertions(+), 114 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 f9242f4651b7..8dacc73eb8e4 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 @@ -31,8 +31,6 @@ class GlmMoeDsaConfig(PreTrainedConfig): r""" n_group (`int`, *optional*, defaults to 1): Number of groups for routed experts. - rope_interleave (`bool`, *optional*, defaults to `True`): - Whether main MLA rotary embeddings use interleaved pair layout. mlp_layer_types (`list`, *optional*): MLP type pattern for each layer (`"dense"` or `"sparse"`). Defaults to `3` dense layers and then every `moe_layer_freq`-th layer sparse. moe_layer_freq (`int`, *optional*, defaults to 1): @@ -49,8 +47,8 @@ class GlmMoeDsaConfig(PreTrainedConfig): Explicit full/shared indexer pattern using `"F"`/`"S"` or `"full"`/`"shared"` values. index_skip_topk_offset (`int`, *optional*, defaults to 2): Offset used with `index_topk_freq` to decide which layers recompute top-k indices. - indexer_rope_interleave (`bool`, *optional*, defaults to `False`): - Whether DSA indexer rotary embeddings use interleaved pair layout. + indexer_rope_interleave (`bool`, *optional*, defaults to `True`): + DSA indexer rotary embeddings always use interleaved pair layout. indexer_types (`list[str]`, *optional*): Indexer mode for each layer (`"full"` or `"shared"`). Defaults to the pattern derived from `index_topk_freq` and `index_skip_topk_offset`. @@ -124,7 +122,6 @@ class GlmMoeDsaConfig(PreTrainedConfig): eos_token_id: int | list[int] | None = 1 tie_word_embeddings: bool = False rope_parameters: RopeParameters | dict | None = None - rope_interleave: bool = True mlp_layer_types: list[str] | None = None attention_bias: bool = False attention_dropout: float | int = 0.0 @@ -141,7 +138,7 @@ class GlmMoeDsaConfig(PreTrainedConfig): index_topk_freq: int = 1 index_topk_pattern: str | list[str] | None = None index_skip_topk_offset: int = 2 - indexer_rope_interleave: bool = False + indexer_rope_interleave: bool = True indexer_types: list[str] | None = None def __post_init__(self, **kwargs): 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 6cd8f810a76a..d6a7d5a16859 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 @@ -77,10 +77,12 @@ def apply_rotary_pos_emb( unsqueeze_dim: int = 1, ) -> torch.Tensor: """ - Applies (non-interleaved, NeoX/Llama style) Rotary Position Embedding to a single tensor. + Applies Rotary Position Embedding to a single tensor (query, key, or the indexer's q/k stream). - This is the transformers equivalent of DeepSeek V3.2's `apply_rotary_emb(x, freqs_cis, interleaved=False)`. - Instead of using complex-number `freqs_cis`, we use pre-split `(cos, sin)` tensors from RotaryEmbedding. + This is the transformers equivalent of DeepSeek V3.2's `apply_rotary_emb(x, freqs_cis, interleaved=True)`. + Instead of complex-number `freqs_cis`, we use pre-split `(cos, sin)` tensors from RotaryEmbedding. + Rotary pairs are always interpreted as adjacent elements (GPT-J style), so we first de-interleave `x` into + halves before applying the standard `rotate_half`. Args: x (`torch.Tensor`): Input tensor of shape `[..., head_dim]`. @@ -94,18 +96,6 @@ def apply_rotary_pos_emb( """ cos = cos.unsqueeze(unsqueeze_dim) sin = sin.unsqueeze(unsqueeze_dim) - return (x * cos) + (rotate_half(x) * sin) - - -def apply_rotary_pos_emb_interleave_single( - x: torch.Tensor, - cos: torch.Tensor, - sin: torch.Tensor, - unsqueeze_dim: int = 1, -) -> torch.Tensor: - """Interleaved (GPT-J style) RoPE applied to a single tensor (the indexer's q/k stream).""" - cos = cos.unsqueeze(unsqueeze_dim) - sin = sin.unsqueeze(unsqueeze_dim) *leading_dims, head_dim = x.shape x = x.view(*leading_dims, head_dim // 2, 2).transpose(-1, -2).reshape(*leading_dims, head_dim) return (x * cos) + (rotate_half(x) * sin) @@ -116,8 +106,8 @@ class GlmMoeDsaIndexer(nn.Module): DeepSeek Sparse Attention (DSA) indexer for selecting top-k tokens. The Indexer has its own lightweight projections (wq_b, wk) separate from the - main MLA attention. RoPE layout (interleaved vs non-interleaved) is controlled - independently from the main attention by `config.indexer_rope_interleave`. + main MLA attention. RoPE uses the same interleaved pair layout as main MLA + attention. **Cache strategy**: The Indexer manages its own key cache (`_cached_keys`) separately from the DynamicCache used by MLA attention, since DynamicCache is sized for exactly @@ -150,12 +140,6 @@ def __init__(self, config: "GlmMoeDsaConfig", layer_idx: int): # Indexer maintains its own key cache (not in DynamicCache, which is sized for attention layers only) self.register_buffer("_cached_keys", None, persistent=False) - def _apply_indexer_rope(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: - """Apply the indexer's RoPE (interleaved or not) to a `[B, S, H, rope_D]` tensor.""" - if self.config.indexer_rope_interleave: - return apply_rotary_pos_emb_interleave_single(x, cos, sin, unsqueeze_dim=2) - return apply_rotary_pos_emb(x, cos, sin, unsqueeze_dim=2) - @torch.no_grad() def forward( self, @@ -193,13 +177,13 @@ def forward( q = self.wq_b(q_resid) # [B, S, H*D] q = q.view(batch_size, seq_len, self.n_heads, self.head_dim) # [B, S, H, D] q_pe, q_nope = torch.split(q, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1) - q_pe = self._apply_indexer_rope(q_pe, cos, sin) # [B, S, H, rope_D] + q_pe = apply_rotary_pos_emb(q_pe, cos, sin, unsqueeze_dim=2) # [B, S, H, rope_D] q = torch.cat([q_pe, q_nope], dim=-1) # [B, S, H, D] # === Keys === k = self.k_norm(self.wk(hidden_states)) # [B, S, D] k_pe, k_nope = torch.split(k, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1) - k_pe = self._apply_indexer_rope(k_pe.unsqueeze(2), cos, sin).squeeze(2) # [B, S, rope_D] + k_pe = apply_rotary_pos_emb(k_pe.unsqueeze(2), cos, sin, unsqueeze_dim=2).squeeze(2) # [B, S, rope_D] k = torch.cat([k_pe, k_nope], dim=-1) # [B, S, D] # === Key cache (managed by the indexer, not DynamicCache) === @@ -244,44 +228,6 @@ def forward( 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. - - 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 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, @@ -434,13 +380,10 @@ def forward( k_nope = k_nope.transpose(1, 2) # [B, H, S, nope_D] value_states = value_states.transpose(1, 2) # [B, H, S, v_D] - # RoPE on q_pe / k_pe (single-head rope stream for k) + # RoPE on q_pe / k_pe (single-head rope stream for k), using interleaved pair layout. k_pe = k_pe.view(batch_size, 1, seq_length, self.qk_rope_head_dim) # [B, 1, S, rope_D] - if self.config.rope_interleave: - q_pe, k_pe = apply_rotary_pos_emb_interleave(q_pe, k_pe, cos, sin) - else: - q_pe = apply_rotary_pos_emb(q_pe, cos, sin, unsqueeze_dim=1) - k_pe = apply_rotary_pos_emb(k_pe, cos, sin, unsqueeze_dim=1) + q_pe = apply_rotary_pos_emb(q_pe, cos, sin, unsqueeze_dim=1) + k_pe = apply_rotary_pos_emb(k_pe, cos, sin, unsqueeze_dim=1) k_pe = k_pe.expand(-1, k_nope.shape[1], -1, -1) # [B, H, S, rope_D] # Assemble full Q and K 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 b14c1a7f1b7d..4e2174142754 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 @@ -29,7 +29,6 @@ from ...processing_utils import Unpack from ...utils import TransformersKwargs, auto_docstring, logging from ...utils.generic import is_flash_attention_requested -from ..deepseek_v3.modeling_deepseek_v3 import apply_rotary_pos_emb_interleave from ..glm4_moe.modeling_glm4_moe import ( Glm4MoeForCausalLM, Glm4MoeModel, @@ -54,10 +53,12 @@ def apply_rotary_pos_emb( unsqueeze_dim: int = 1, ) -> torch.Tensor: """ - Applies (non-interleaved, NeoX/Llama style) Rotary Position Embedding to a single tensor. + Applies Rotary Position Embedding to a single tensor (query, key, or the indexer's q/k stream). - This is the transformers equivalent of DeepSeek V3.2's `apply_rotary_emb(x, freqs_cis, interleaved=False)`. - Instead of using complex-number `freqs_cis`, we use pre-split `(cos, sin)` tensors from RotaryEmbedding. + This is the transformers equivalent of DeepSeek V3.2's `apply_rotary_emb(x, freqs_cis, interleaved=True)`. + Instead of complex-number `freqs_cis`, we use pre-split `(cos, sin)` tensors from RotaryEmbedding. + Rotary pairs are always interpreted as adjacent elements (GPT-J style), so we first de-interleave `x` into + halves before applying the standard `rotate_half`. Args: x (`torch.Tensor`): Input tensor of shape `[..., head_dim]`. @@ -71,18 +72,6 @@ def apply_rotary_pos_emb( """ cos = cos.unsqueeze(unsqueeze_dim) sin = sin.unsqueeze(unsqueeze_dim) - return (x * cos) + (rotate_half(x) * sin) - - -def apply_rotary_pos_emb_interleave_single( - x: torch.Tensor, - cos: torch.Tensor, - sin: torch.Tensor, - unsqueeze_dim: int = 1, -) -> torch.Tensor: - """Interleaved (GPT-J style) RoPE applied to a single tensor (the indexer's q/k stream).""" - cos = cos.unsqueeze(unsqueeze_dim) - sin = sin.unsqueeze(unsqueeze_dim) *leading_dims, head_dim = x.shape x = x.view(*leading_dims, head_dim // 2, 2).transpose(-1, -2).reshape(*leading_dims, head_dim) return (x * cos) + (rotate_half(x) * sin) @@ -94,8 +83,6 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): r""" n_group (`int`, *optional*, defaults to 1): Number of groups for routed experts. - rope_interleave (`bool`, *optional*, defaults to `True`): - Whether main MLA rotary embeddings use interleaved pair layout. mlp_layer_types (`list`, *optional*): MLP type pattern for each layer (`"dense"` or `"sparse"`). Defaults to `3` dense layers and then every `moe_layer_freq`-th layer sparse. moe_layer_freq (`int`, *optional*, defaults to 1): @@ -112,8 +99,8 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): Explicit full/shared indexer pattern using `"F"`/`"S"` or `"full"`/`"shared"` values. index_skip_topk_offset (`int`, *optional*, defaults to 2): Offset used with `index_topk_freq` to decide which layers recompute top-k indices. - indexer_rope_interleave (`bool`, *optional*, defaults to `False`): - Whether DSA indexer rotary embeddings use interleaved pair layout. + indexer_rope_interleave (`bool`, *optional*, defaults to `True`): + DSA indexer rotary embeddings always use interleaved pair layout. indexer_types (`list[str]`, *optional*): Indexer mode for each layer (`"full"` or `"shared"`). Defaults to the pattern derived from `index_topk_freq` and `index_skip_topk_offset`. @@ -169,13 +156,13 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): index_head_dim: int = 128 index_n_heads: int = 32 moe_layer_freq: int = 1 - rope_interleave: bool = True index_topk_freq: int = 1 index_topk_pattern: str | list[str] | None = None index_skip_topk_offset: int = 2 - indexer_rope_interleave: bool = False + indexer_rope_interleave: bool = True indexer_types: list[str] | None = None pretraining_tp = AttributeError() + rope_interleave = AttributeError() def __post_init__(self, **kwargs): self.qk_head_dim = self.qk_nope_head_dim + self.qk_rope_head_dim @@ -208,8 +195,8 @@ class GlmMoeDsaIndexer(nn.Module): DeepSeek Sparse Attention (DSA) indexer for selecting top-k tokens. The Indexer has its own lightweight projections (wq_b, wk) separate from the - main MLA attention. RoPE layout (interleaved vs non-interleaved) is controlled - independently from the main attention by `config.indexer_rope_interleave`. + main MLA attention. RoPE uses the same interleaved pair layout as main MLA + attention. **Cache strategy**: The Indexer manages its own key cache (`_cached_keys`) separately from the DynamicCache used by MLA attention, since DynamicCache is sized for exactly @@ -242,12 +229,6 @@ def __init__(self, config: "GlmMoeDsaConfig", layer_idx: int): # Indexer maintains its own key cache (not in DynamicCache, which is sized for attention layers only) self.register_buffer("_cached_keys", None, persistent=False) - def _apply_indexer_rope(self, x: torch.Tensor, cos: torch.Tensor, sin: torch.Tensor) -> torch.Tensor: - """Apply the indexer's RoPE (interleaved or not) to a `[B, S, H, rope_D]` tensor.""" - if self.config.indexer_rope_interleave: - return apply_rotary_pos_emb_interleave_single(x, cos, sin, unsqueeze_dim=2) - return apply_rotary_pos_emb(x, cos, sin, unsqueeze_dim=2) - @torch.no_grad() def forward( self, @@ -285,13 +266,13 @@ def forward( q = self.wq_b(q_resid) # [B, S, H*D] q = q.view(batch_size, seq_len, self.n_heads, self.head_dim) # [B, S, H, D] q_pe, q_nope = torch.split(q, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1) - q_pe = self._apply_indexer_rope(q_pe, cos, sin) # [B, S, H, rope_D] + q_pe = apply_rotary_pos_emb(q_pe, cos, sin, unsqueeze_dim=2) # [B, S, H, rope_D] q = torch.cat([q_pe, q_nope], dim=-1) # [B, S, H, D] # === Keys === k = self.k_norm(self.wk(hidden_states)) # [B, S, D] k_pe, k_nope = torch.split(k, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1) - k_pe = self._apply_indexer_rope(k_pe.unsqueeze(2), cos, sin).squeeze(2) # [B, S, rope_D] + k_pe = apply_rotary_pos_emb(k_pe.unsqueeze(2), cos, sin, unsqueeze_dim=2).squeeze(2) # [B, S, rope_D] k = torch.cat([k_pe, k_nope], dim=-1) # [B, S, D] # === Key cache (managed by the indexer, not DynamicCache) === @@ -451,13 +432,10 @@ def forward( k_nope = k_nope.transpose(1, 2) # [B, H, S, nope_D] value_states = value_states.transpose(1, 2) # [B, H, S, v_D] - # RoPE on q_pe / k_pe (single-head rope stream for k) + # RoPE on q_pe / k_pe (single-head rope stream for k), using interleaved pair layout. k_pe = k_pe.view(batch_size, 1, seq_length, self.qk_rope_head_dim) # [B, 1, S, rope_D] - if self.config.rope_interleave: - q_pe, k_pe = apply_rotary_pos_emb_interleave(q_pe, k_pe, cos, sin) - else: - q_pe = apply_rotary_pos_emb(q_pe, cos, sin, unsqueeze_dim=1) - k_pe = apply_rotary_pos_emb(k_pe, cos, sin, unsqueeze_dim=1) + q_pe = apply_rotary_pos_emb(q_pe, cos, sin, unsqueeze_dim=1) + k_pe = apply_rotary_pos_emb(k_pe, cos, sin, unsqueeze_dim=1) k_pe = k_pe.expand(-1, k_nope.shape[1], -1, -1) # [B, H, S, rope_D] # Assemble full Q and K diff --git a/utils/check_config_attributes.py b/utils/check_config_attributes.py index fc86d51c21e4..1f817aa469ea 100644 --- a/utils/check_config_attributes.py +++ b/utils/check_config_attributes.py @@ -109,7 +109,13 @@ "Cohere2MoeConfig": ["rope_scaling", "sliding_window_pattern"], "CsmConfig": ["tie_codebooks_embeddings"], "DeepseekV2Config": ["norm_topk_prob"], - "GlmMoeDsaConfig": ["index_skip_topk_offset", "index_topk_freq", "index_topk_pattern", "moe_layer_freq"], + "GlmMoeDsaConfig": [ + "index_skip_topk_offset", + "index_topk_freq", + "index_topk_pattern", + "moe_layer_freq", + "indexer_rope_interleave", + ], "DeepseekV4Config": [ # All BC / config-compat surface that the modeling code never reads but # checkpoints in the wild expose (so we keep accepting them in `__init__`): From f225757a351b5b19db24d2ea93258f46b3b67e22 Mon Sep 17 00:00:00 2001 From: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> Date: Fri, 5 Jun 2026 08:29:37 +0000 Subject: [PATCH 8/8] Optimize code Signed-off-by: Shijin Zhang <75300765+Dovis01@users.noreply.github.com> --- .../glm_moe_dsa/configuration_glm_moe_dsa.py | 31 ++++--------------- .../models/glm_moe_dsa/modular_glm_moe_dsa.py | 29 +++-------------- utils/check_config_attributes.py | 7 ----- 3 files changed, 11 insertions(+), 56 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 8dacc73eb8e4..c96a274e1a4f 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 @@ -33,22 +33,12 @@ class GlmMoeDsaConfig(PreTrainedConfig): Number of groups for routed experts. mlp_layer_types (`list`, *optional*): MLP type pattern for each layer (`"dense"` or `"sparse"`). Defaults to `3` dense layers and then every `moe_layer_freq`-th layer sparse. - moe_layer_freq (`int`, *optional*, defaults to 1): - Frequency for sparse MoE layers. index_topk (`int`, *optional*, defaults to 2048): Number of top tokens selected by the indexer for sparse attention. index_head_dim (`int`, *optional*, defaults to 128): Head dimension for the indexer projections (DSA). index_n_heads (`int | None`, *optional*, defaults to 32): Number of heads for the indexer projections (DSA). - index_topk_freq (`int`, *optional*, defaults to 1): - Frequency for full indexer recomputation when `index_topk_pattern` is not provided. - index_topk_pattern (`str | list[str]`, *optional*): - Explicit full/shared indexer pattern using `"F"`/`"S"` or `"full"`/`"shared"` values. - index_skip_topk_offset (`int`, *optional*, defaults to 2): - Offset used with `index_topk_freq` to decide which layers recompute top-k indices. - indexer_rope_interleave (`bool`, *optional*, defaults to `True`): - DSA indexer rotary embeddings always use interleaved pair layout. indexer_types (`list[str]`, *optional*): Indexer mode for each layer (`"full"` or `"shared"`). Defaults to the pattern derived from `index_topk_freq` and `index_skip_topk_offset`. @@ -88,6 +78,7 @@ class GlmMoeDsaConfig(PreTrainedConfig): "layers": (["hidden_states", "attention_mask"], ["hidden_states"]), "norm": (["hidden_states"], ["hidden_states"]), } + attribute_map = { "num_local_experts": "n_routed_experts", } @@ -125,38 +116,28 @@ class GlmMoeDsaConfig(PreTrainedConfig): mlp_layer_types: list[str] | None = None attention_bias: bool = False attention_dropout: float | int = 0.0 - - base_model_fsdp_plan = { - "embed_tokens": "free_full_weight", - "layers.*": "free_full_weight", - "norm": "keep_full_weight", - } index_topk: int = 2048 index_head_dim: int = 128 index_n_heads: int = 32 - moe_layer_freq: int = 1 - index_topk_freq: int = 1 - index_topk_pattern: str | list[str] | None = None - index_skip_topk_offset: int = 2 - indexer_rope_interleave: bool = True indexer_types: list[str] | None = None def __post_init__(self, **kwargs): self.qk_head_dim = self.qk_nope_head_dim + self.qk_rope_head_dim if self.mlp_layer_types is None: + moe_layer_freq = kwargs.get("moe_layer_freq", 1) self.mlp_layer_types = [ - "sparse" if i >= 3 and i % self.moe_layer_freq == 0 else "dense" for i in range(self.num_hidden_layers) + "sparse" if i >= 3 and i % moe_layer_freq == 0 else "dense" for i in range(self.num_hidden_layers) ] if self.indexer_types is None: - pattern = self.index_topk_pattern + pattern = kwargs.get("index_topk_pattern") if pattern is not None: self.indexer_types = ( [{"F": "full", "S": "shared"}[c] for c in pattern] if isinstance(pattern, str) else list(pattern) ) else: - freq = max(self.index_topk_freq, 1) - offset = self.index_skip_topk_offset + freq = max(kwargs.get("index_topk_freq", 1), 1) + offset = kwargs.get("index_skip_topk_offset", 2) self.indexer_types = [ "full" if (max(i - offset + 1, 0) % freq) == 0 else "shared" for i in range(self.num_hidden_layers) ] 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 4e2174142754..28d8dd6e17f6 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 @@ -85,22 +85,12 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): Number of groups for routed experts. mlp_layer_types (`list`, *optional*): MLP type pattern for each layer (`"dense"` or `"sparse"`). Defaults to `3` dense layers and then every `moe_layer_freq`-th layer sparse. - moe_layer_freq (`int`, *optional*, defaults to 1): - Frequency for sparse MoE layers. index_topk (`int`, *optional*, defaults to 2048): Number of top tokens selected by the indexer for sparse attention. index_head_dim (`int`, *optional*, defaults to 128): Head dimension for the indexer projections (DSA). index_n_heads (`int | None`, *optional*, defaults to 32): Number of heads for the indexer projections (DSA). - index_topk_freq (`int`, *optional*, defaults to 1): - Frequency for full indexer recomputation when `index_topk_pattern` is not provided. - index_topk_pattern (`str | list[str]`, *optional*): - Explicit full/shared indexer pattern using `"F"`/`"S"` or `"full"`/`"shared"` values. - index_skip_topk_offset (`int`, *optional*, defaults to 2): - Offset used with `index_topk_freq` to decide which layers recompute top-k indices. - indexer_rope_interleave (`bool`, *optional*, defaults to `True`): - DSA indexer rotary embeddings always use interleaved pair layout. indexer_types (`list[str]`, *optional*): Indexer mode for each layer (`"full"` or `"shared"`). Defaults to the pattern derived from `index_topk_freq` and `index_skip_topk_offset`. @@ -133,11 +123,6 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): "layers.*.mlp.down_proj": "rowwise", } - base_model_fsdp_plan = { - "embed_tokens": "free_full_weight", - "layers.*": "free_full_weight", - "norm": "keep_full_weight", - } attribute_map = { "num_local_experts": "n_routed_experts", } @@ -155,11 +140,6 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): index_topk: int = 2048 index_head_dim: int = 128 index_n_heads: int = 32 - moe_layer_freq: int = 1 - index_topk_freq: int = 1 - index_topk_pattern: str | list[str] | None = None - index_skip_topk_offset: int = 2 - indexer_rope_interleave: bool = True indexer_types: list[str] | None = None pretraining_tp = AttributeError() rope_interleave = AttributeError() @@ -167,19 +147,20 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig): def __post_init__(self, **kwargs): self.qk_head_dim = self.qk_nope_head_dim + self.qk_rope_head_dim if self.mlp_layer_types is None: + moe_layer_freq = kwargs.get("moe_layer_freq", 1) self.mlp_layer_types = [ - "sparse" if i >= 3 and i % self.moe_layer_freq == 0 else "dense" for i in range(self.num_hidden_layers) + "sparse" if i >= 3 and i % moe_layer_freq == 0 else "dense" for i in range(self.num_hidden_layers) ] if self.indexer_types is None: - pattern = self.index_topk_pattern + pattern = kwargs.get("index_topk_pattern") if pattern is not None: self.indexer_types = ( [{"F": "full", "S": "shared"}[c] for c in pattern] if isinstance(pattern, str) else list(pattern) ) else: - freq = max(self.index_topk_freq, 1) - offset = self.index_skip_topk_offset + freq = max(kwargs.get("index_topk_freq", 1), 1) + offset = kwargs.get("index_skip_topk_offset", 2) self.indexer_types = [ "full" if (max(i - offset + 1, 0) % freq) == 0 else "shared" for i in range(self.num_hidden_layers) ] diff --git a/utils/check_config_attributes.py b/utils/check_config_attributes.py index 1f817aa469ea..8cd992d541ff 100644 --- a/utils/check_config_attributes.py +++ b/utils/check_config_attributes.py @@ -109,13 +109,6 @@ "Cohere2MoeConfig": ["rope_scaling", "sliding_window_pattern"], "CsmConfig": ["tie_codebooks_embeddings"], "DeepseekV2Config": ["norm_topk_prob"], - "GlmMoeDsaConfig": [ - "index_skip_topk_offset", - "index_topk_freq", - "index_topk_pattern", - "moe_layer_freq", - "indexer_rope_interleave", - ], "DeepseekV4Config": [ # All BC / config-compat surface that the modeling code never reads but # checkpoints in the wild expose (so we keep accepting them in `__init__`):