diff --git a/docs/source/en/_toctree.yml b/docs/source/en/_toctree.yml
index 4675e847f6b1..3c776a0c8d3c 100644
--- a/docs/source/en/_toctree.yml
+++ b/docs/source/en/_toctree.yml
@@ -569,6 +569,8 @@
title: DeepSeek-V2
- local: model_doc/deepseek_v3
title: DeepSeek-V3
+ - local: model_doc/deepseek_v32
+ title: DeepSeek-V3.2
- local: model_doc/deepseek_v4
title: DeepSeek-V4
- local: model_doc/dialogpt
diff --git a/docs/source/en/model_doc/deepseek_v32.md b/docs/source/en/model_doc/deepseek_v32.md
new file mode 100644
index 000000000000..b499e22bd4be
--- /dev/null
+++ b/docs/source/en/model_doc/deepseek_v32.md
@@ -0,0 +1,100 @@
+
+*This model was published in HF papers on 2025-12-02 and contributed to Hugging Face Transformers on 2026-06-10.*
+
+
+
+# DeepSeek-V3.2
+
+## Overview
+
+[DeepSeek-V3.2-Exp](https://huggingface.co/deepseek-ai/DeepSeek-V3.2-Exp) is an experimental release from DeepSeek-AI that introduces **DeepSeek Sparse Attention (DSA)**, a trainable, fine-grained sparse attention mechanism designed to improve training and inference efficiency in long-context scenarios. It is built directly on top of [DeepSeek-V3.1-Terminus](https://huggingface.co/deepseek-ai/DeepSeek-V3.1-Terminus): the model keeps the same 685B-parameter Mixture-of-Experts (MoE) backbone and Multi-head Latent Attention (MLA), and is obtained through continued training that adds the sparse-attention indexer while deliberately aligning the training distribution with V3.1-Terminus so the two models can be compared head-to-head.
+
+The work was later extended in the [DeepSeek-V3.2 technical report](https://huggingface.co/papers/2512.02556), *DeepSeek-V3.2: Pushing the Frontier of Open Large Language Models*, which pairs DSA with a scalable reinforcement-learning framework and reports gold-medal level results on competition math (IMO) and competitive programming (IOI) benchmarks.
+
+The abstract from the DeepSeek-V3.2-Exp release is the following:
+
+*We introduce DeepSeek-V3.2-Exp, an experimental version of our model that incorporates DeepSeek Sparse Attention (DSA) to explore and validate optimizations for training and inference efficiency in long-context scenarios. DeepSeek Sparse Attention achieves fine-grained sparse attention for the first time with minimal impact on model output quality. Built upon DeepSeek-V3.1-Terminus, DeepSeek-V3.2-Exp delivers substantially improved efficiency in both training and inference, especially in long-context settings, while maintaining virtually identical benchmark performance.*
+
+### DeepSeek Sparse Attention (DSA)
+
+DSA reduces the quadratic cost of attention over long sequences by attending only to a selected subset of past tokens. It has two components:
+
+1. **Lightning indexer.** A lightweight, low-head-count scoring module computes an *index score* between each query and every preceding key. In the reference implementation it runs in FP8 with a Hadamard (`rotate_activation`) transform; because the transform is orthogonal (`Hq·Hk = q·k`) and FP8 is only a precision optimization, the transformers port computes the same scores directly in bf16/fp32, keeping the indexer cheap relative to the main attention.
+2. **Fine-grained token selection.** For each query the indexer keeps the top-`index_topk` (2048 by default) tokens, and main MLA attention is then computed only over those tokens via an additive mask. This turns the per-query attention cost from `O(L)` to `O(index_topk)` for long sequences when using `flash_mla`, which is not supported yet 😉.
+
+The indexer keeps its own small per-token key cache (single-head, `index_head_dim`) alongside the main K/V cache. In transformers this lives on a dedicated cache layer — [`DynamicIndexedLayer`] for growing caches and [`StaticIndexedLayer`] for static / `torch.compile` caches — and is updated through `past_key_values.update_indexer()`.
+
+In DeepSeek-V3.2 **every layer runs its own indexer** — there is no cross-layer top-k sharing.
+
+> [!NOTE]
+> **The MLA query LoRA path (`q_lora_rank`) is required.** The indexer scores queries from the low-rank query latent `q_a_layernorm(q_a_proj(x))` (its `wq_b` projection is sized by `q_lora_rank`), so the model always uses the LoRA query path and `q_lora_rank` must be set — the released checkpoint uses `1536`. The optional non-LoRA `q_proj` path that [DeepSeek-V3](./deepseek_v3) exposes for `q_lora_rank=None` is **not supported** here: without the query latent there is nothing for the indexer to consume.
+
+## Usage examples
+
+DeepSeek-V3.2-Exp is distributed as an FP8 checkpoint. The indexer projections are kept out of FP8 quantization, since the checkpoint stores them in bf16/fp32:
+
+```python
+from transformers import FineGrainedFP8Config, AutoModelForCausalLM, AutoTokenizer
+import torch
+
+model_name = "deepseek-ai/DeepSeek-V3.2-Exp"
+quantization_config = FineGrainedFP8Config(
+ modules_to_not_convert=["model.layers.*.mlp.gate.*", "*.self_attn.indexer.weights_proj.*"],
+ weight_block_size=(128, 128),
+)
+model = AutoModelForCausalLM.from_pretrained(
+ model_name,
+ torch_dtype="auto",
+ device_map="auto",
+ quantization_config=quantization_config,
+)
+tokenizer = AutoTokenizer.from_pretrained(model_name)
+
+inputs = tokenizer("What are we having for dinner?", return_tensors="pt").to(model.device)
+outputs = model.generate(**inputs, max_new_tokens=20)
+print(tokenizer.decode(outputs[0], skip_special_tokens=True))
+```
+
+The original code can be found [here](https://github.com/deepseek-ai/DeepSeek-V3.2-Exp).
+
+## DeepseekV32Config
+
+[[autodoc]] DeepseekV32Config
+
+## DeepseekV32PreTrainedModel
+
+[[autodoc]] DeepseekV32PreTrainedModel
+ - forward
+
+## DeepseekV32Model
+
+[[autodoc]] DeepseekV32Model
+ - forward
+
+## DeepseekV32ForCausalLM
+
+[[autodoc]] DeepseekV32ForCausalLM
diff --git a/docs/source/en/model_doc/glm_moe_dsa.md b/docs/source/en/model_doc/glm_moe_dsa.md
index 49442cee5c1a..96e799a90107 100644
--- a/docs/source/en/model_doc/glm_moe_dsa.md
+++ b/docs/source/en/model_doc/glm_moe_dsa.md
@@ -66,6 +66,9 @@ print(tokenizer.decode(output[0], skip_special_tokens=True))
+> [!NOTE]
+> **The MLA query LoRA path (`q_lora_rank`) is required.** Like DeepSeek-V3.2, the DSA indexer scores queries from the low-rank query latent `q_a_layernorm(q_a_proj(x))` (its `wq_b` projection is sized by `q_lora_rank`), so the model always uses the LoRA query path and `q_lora_rank` must be set — the released checkpoint uses `2048`. The optional non-LoRA `q_proj` path that [DeepSeek-V3](./deepseek_v3) exposes for `q_lora_rank=None` is **not supported** here: without the query latent there is nothing for the indexer to consume.
+
## GlmMoeDsaConfig
[[autodoc]] GlmMoeDsaConfig
diff --git a/src/transformers/__init__.py b/src/transformers/__init__.py
index 20758d5421ae..722f9fcff0c2 100755
--- a/src/transformers/__init__.py
+++ b/src/transformers/__init__.py
@@ -367,7 +367,9 @@
_import_structure["cache_utils"] = [
"CacheLayerMixin",
"DynamicLayer",
+ "DynamicIndexedLayer",
"StaticLayer",
+ "StaticIndexedLayer",
"StaticSlidingWindowLayer",
"QuantoQuantizedLayer",
"HQQQuantizedLayer",
@@ -487,12 +489,14 @@
from .backbone_utils import BackboneConfigMixin, BackboneMixin
from .cache_utils import Cache as Cache
from .cache_utils import DynamicCache as DynamicCache
+ from .cache_utils import DynamicIndexedLayer as DynamicIndexedLayer
from .cache_utils import DynamicLayer as DynamicLayer
from .cache_utils import EncoderDecoderCache as EncoderDecoderCache
from .cache_utils import HQQQuantizedLayer as HQQQuantizedLayer
from .cache_utils import QuantizedCache as QuantizedCache
from .cache_utils import QuantoQuantizedLayer as QuantoQuantizedLayer
from .cache_utils import StaticCache as StaticCache
+ from .cache_utils import StaticIndexedLayer as StaticIndexedLayer
from .cache_utils import StaticLayer as StaticLayer
from .cache_utils import StaticSlidingWindowLayer as StaticSlidingWindowLayer
from .configuration_utils import PreTrainedConfig as PreTrainedConfig
diff --git a/src/transformers/cache_utils.py b/src/transformers/cache_utils.py
index 1f4f4b9c957d..4801b37a8558 100644
--- a/src/transformers/cache_utils.py
+++ b/src/transformers/cache_utils.py
@@ -274,6 +274,82 @@ def crop(self, max_length: int) -> None:
self.cumulative_length = self.keys.shape[-2]
+class DynamicIndexedLayer(DynamicLayer):
+ """
+ A cache layer that extends `DynamicLayer` with an extra indexer key cache for Dynamic Sparse Attention (DSA)
+ models (e.g. GLM MoE DSA, DeepSeek V32).
+
+ The main K/V cache stores tensors of shape `[batch_size, num_heads, seq_len, head_dim]` (inherited).
+ The indexer key cache stores a tensor of shape `[batch_size, seq_len, index_head_dim]` (3D, single-head).
+ """
+
+ # Auto-registers in ``LAYER_TYPE_CACHE_MAPPING`` so ``DynamicCache`` dispatches DSA layers here.
+ layer_type = "deepseek_sparse_attention"
+
+ def __init__(self, config: PreTrainedConfig | None = None):
+ super().__init__(config)
+ self.indexer_keys: torch.Tensor | None = None
+ self.is_indexer_initialized: bool = False
+
+ def lazy_initialization_indexer(self, indexer_key_states: torch.Tensor) -> None:
+ self.indexer_dtype, self.indexer_device = indexer_key_states.dtype, indexer_key_states.device
+ self.indexer_keys = torch.tensor([], dtype=self.indexer_dtype, device=self.indexer_device)
+ self.is_indexer_initialized = True
+
+ def update_indexer(self, indexer_key_states: torch.Tensor) -> torch.Tensor:
+ """
+ Update the indexer key cache by concatenation, and return the full indexer keys.
+
+ Args:
+ indexer_key_states (`torch.Tensor`): New indexer keys, shape `[batch_size, seq_len, index_head_dim]`.
+
+ Returns:
+ `torch.Tensor`: The full cached indexer keys, shape `[batch_size, total_len, index_head_dim]`.
+ """
+ if not self.is_indexer_initialized:
+ self.lazy_initialization_indexer(indexer_key_states)
+ self.indexer_keys = torch.cat([self.indexer_keys, indexer_key_states], dim=1)
+ return self.indexer_keys
+
+ def offload(self):
+ super().offload()
+ if self.is_indexer_initialized:
+ self.indexer_keys = self.indexer_keys.to("cpu", non_blocking=True)
+
+ def prefetch(self):
+ super().prefetch()
+ if self.is_indexer_initialized and self.indexer_keys.device != self.device:
+ self.indexer_keys = self.indexer_keys.to(self.device, non_blocking=True)
+
+ def reset(self) -> None:
+ super().reset()
+ if self.is_indexer_initialized:
+ self.indexer_keys.zero_()
+
+ def reorder_cache(self, beam_idx: torch.LongTensor) -> None:
+ super().reorder_cache(beam_idx)
+ if self.is_indexer_initialized and self.indexer_keys.numel() > 0:
+ self.indexer_keys = self.indexer_keys.index_select(0, beam_idx.to(self.indexer_keys.device))
+
+ def crop(self, max_length: int) -> None:
+ super().crop(max_length)
+ if not self.is_indexer_initialized or self.indexer_keys.numel() == 0:
+ return
+ effective = max_length if max_length >= 0 else self.indexer_keys.shape[1] - abs(max_length)
+ if self.indexer_keys.shape[1] > effective:
+ self.indexer_keys = self.indexer_keys[:, :effective, :]
+
+ def batch_repeat_interleave(self, repeats: int) -> None:
+ super().batch_repeat_interleave(repeats)
+ if self.is_indexer_initialized and self.indexer_keys.numel() > 0:
+ self.indexer_keys = self.indexer_keys.repeat_interleave(repeats, dim=0)
+
+ def batch_select_indices(self, indices: torch.Tensor) -> None:
+ super().batch_select_indices(indices)
+ if self.is_indexer_initialized and self.indexer_keys.numel() > 0:
+ self.indexer_keys = self.indexer_keys[indices, ...]
+
+
class StaticLayer(CacheLayerMixin):
"""
A static cache layer that stores the key and value states as static tensors of shape `[batch_size, num_heads, max_cache_len), head_dim]`.
@@ -511,6 +587,73 @@ def reset(self):
self.cumulative_length_int = 0
+class StaticIndexedLayer(StaticLayer):
+ """
+ A `StaticLayer` with an additional statically-allocated indexer key cache for Dynamic Sparse
+ Attention (DSA) models (e.g. GLM MoE DSA, DeepSeek V32). This is the static, `torch.compile`-friendly
+ counterpart of `DynamicIndexedLayer`: the indexer key buffer is preallocated once and mutated in-place.
+
+ The main K/V cache is inherited from `StaticLayer` (`[batch_size, num_heads, max_cache_len, head_dim]`).
+ The indexer key cache stores a tensor of shape `[batch_size, max_cache_len, index_head_dim]` (3D, single-head).
+ """
+
+ def __init__(self, max_cache_len: int):
+ super().__init__(max_cache_len=max_cache_len)
+ self.indexer_keys: torch.Tensor | None = None
+ self.is_indexer_initialized: bool = False
+ # The indexer update runs independently of (and after) the main K/V `update` in the attention
+ # forward, so it tracks its own cumulative length rather than reusing `self.cumulative_length`.
+ self.indexer_cumulative_length = torch.tensor([0], dtype=int)
+
+ def lazy_initialization_indexer(self, indexer_key_states: torch.Tensor) -> None:
+ self.indexer_dtype, self.indexer_device = indexer_key_states.dtype, indexer_key_states.device
+ max_batch_size, _, index_head_dim = indexer_key_states.shape
+ self.indexer_keys = torch.zeros(
+ (max_batch_size, self.max_cache_len, index_head_dim),
+ dtype=self.indexer_dtype,
+ device=self.indexer_device,
+ )
+ self.indexer_cumulative_length = self.indexer_cumulative_length.to(self.indexer_device)
+ # Tag as static addresses for cudagraphs / compile, mirroring the main K/V buffers.
+ if not is_torchdynamo_compiling():
+ torch._dynamo.mark_static_address(self.indexer_keys)
+ torch._dynamo.mark_static_address(self.indexer_cumulative_length)
+ self.is_indexer_initialized = True
+
+ def update_indexer(self, indexer_key_states: torch.Tensor) -> torch.Tensor:
+ """
+ Update the indexer key cache in-place at the current positions, and return the full static buffer.
+
+ Args:
+ indexer_key_states (`torch.Tensor`): New indexer keys, shape `[batch_size, seq_len, index_head_dim]`.
+
+ Returns:
+ `torch.Tensor`: The full static indexer key cache, shape `[batch_size, max_cache_len, index_head_dim]`.
+ Unfilled positions are masked out downstream by the indexer's attention mask, exactly as the
+ main `StaticLayer` returns its full preallocated K/V.
+ """
+ if not self.is_indexer_initialized:
+ self.lazy_initialization_indexer(indexer_key_states)
+
+ seq_len = indexer_key_states.shape[1]
+ cache_position = torch.arange(seq_len, device=self.indexer_device) + self.indexer_cumulative_length
+ # In-place to preserve the static data pointer (required for cudagraphs).
+ self.indexer_cumulative_length.add_(seq_len)
+ try:
+ self.indexer_keys.index_copy_(1, cache_position, indexer_key_states)
+ except NotImplementedError:
+ # Fallback for devices like MPS where index_copy_ might not be supported.
+ self.indexer_keys[:, cache_position] = indexer_key_states
+
+ return self.indexer_keys
+
+ def reset(self) -> None:
+ super().reset()
+ if self.is_indexer_initialized:
+ self.indexer_keys.zero_()
+ self.indexer_cumulative_length.zero_()
+
+
class QuantizedLayer(DynamicLayer):
"""
A quantized layer similar to what is described in the [KIVI: A Tuning-Free Asymmetric 2bit Quantization for KV Cache paper](https://huggingface.co/papers/2402.02750).
@@ -1037,6 +1180,27 @@ def update_recurrent_state(self, recurrent_states: torch.Tensor, layer_idx: int,
recurrent_states = self.layers[layer_idx].update_recurrent_state(recurrent_states, **kwargs)
return recurrent_states
+ def update_indexer(self, indexer_key_states: torch.Tensor, layer_idx: int) -> torch.Tensor:
+ """
+ Updates the indexer key cache for layer `layer_idx`.
+
+ Parameters:
+ indexer_key_states (`torch.Tensor`):
+ The new indexer key states to cache, shape `[batch_size, seq_len, index_head_dim]`.
+ layer_idx (`int`):
+ The index of the layer to cache the states for.
+
+ Return:
+ `torch.Tensor`: The updated indexer key states (full cache).
+ """
+ if not hasattr(self.layers[layer_idx], "update_indexer"):
+ raise ValueError(
+ f"Cannot call `update_indexer` on layer {layer_idx} which is a "
+ f"{type(self.layers[layer_idx]).__name__}; it has no indexer key cache "
+ f"(expected a `DynamicIndexedLayer` or `StaticIndexedLayer`)."
+ )
+ return self.layers[layer_idx].update_indexer(indexer_key_states)
+
def early_initialization(
self,
batch_size: int,
@@ -1414,6 +1578,9 @@ def __init__(
# LinearAttention layers are static by essence - using `"moe"` as well is a trick, see the comment about it on DynamicCache
elif layer_type in ("mamba", "conv", "linear_attention", "moe"):
layer = LinearAttentionLayer()
+ elif layer_type == "deepseek_sparse_attention":
+ # Static / compile-friendly indexed layer (preallocated indexer key cache).
+ layer = StaticIndexedLayer(max_cache_len=max_cache_len)
else:
layer = StaticLayer(max_cache_len=max_cache_len)
layers.append(layer)
diff --git a/src/transformers/configuration_utils.py b/src/transformers/configuration_utils.py
index 89d39625ac19..f902994eb71b 100755
--- a/src/transformers/configuration_utils.py
+++ b/src/transformers/configuration_utils.py
@@ -73,6 +73,7 @@
"dense",
"hybrid", # for layers that have both mamba and attention in zamba and zamba2
"moe", # for nemotron_h, which uses either attention, mamba or moe
+ "deepseek_sparse_attention", # for models with DSA indexer (GLM MoE DSA, DeepSeek V32)
)
diff --git a/src/transformers/conversion_mapping.py b/src/transformers/conversion_mapping.py
index 34d129762586..fe1a66b7a210 100755
--- a/src/transformers/conversion_mapping.py
+++ b/src/transformers/conversion_mapping.py
@@ -44,6 +44,7 @@
"afmoe": "qwen2_moe",
"deepseek_v2": "qwen2_moe",
"deepseek_v3": "qwen2_moe",
+ "deepseek_v32": "qwen2_moe",
"dots1": "qwen2_moe",
"ernie4_5_moe": "qwen2_moe",
"glm4_moe": "qwen2_moe",
diff --git a/src/transformers/masking_utils.py b/src/transformers/masking_utils.py
index 32ef0eb97128..4a0022abac5f 100644
--- a/src/transformers/masking_utils.py
+++ b/src/transformers/masking_utils.py
@@ -1458,6 +1458,7 @@ def create_chunked_causal_mask(
"chunked_attention": create_chunked_causal_mask,
"compressed_sparse_attention": create_sliding_window_causal_mask,
"heavily_compressed_attention": create_sliding_window_causal_mask,
+ "deepseek_sparse_attention": create_causal_mask,
}
@@ -1514,10 +1515,14 @@ def create_masks_for_generate(
"block_sequence_ids": block_sequence_ids,
}
- # If the attribute exist, we need several masks
+ # If the attribute exist, we need several masks - unless every layer shares the same type, in which
+ # case we return a single mask.
if hasattr(effective_config, "layer_types"):
+ layer_patterns = set(effective_config.layer_types)
+ if len(layer_patterns) == 1:
+ return LAYER_PATTERN_TO_MASK_FUNCTION_MAPPING[next(iter(layer_patterns))](**mask_kwargs)
causal_masks = {}
- for layer_pattern in set(effective_config.layer_types):
+ for layer_pattern in layer_patterns:
causal_masks[layer_pattern] = LAYER_PATTERN_TO_MASK_FUNCTION_MAPPING[layer_pattern](**mask_kwargs)
return causal_masks
# In this case, all layers are sliding
diff --git a/src/transformers/models/__init__.py b/src/transformers/models/__init__.py
index 757ed0f02a2f..e5d880b9a46b 100644
--- a/src/transformers/models/__init__.py
+++ b/src/transformers/models/__init__.py
@@ -97,6 +97,7 @@
from .deepseek_v2 import *
from .deepseek_v3 import *
from .deepseek_v4 import *
+ from .deepseek_v32 import *
from .deepseek_vl import *
from .deepseek_vl_hybrid import *
from .deformable_detr import *
diff --git a/src/transformers/models/auto/auto_mappings.py b/src/transformers/models/auto/auto_mappings.py
index 5cd442387c06..7b088c8fe7f4 100644
--- a/src/transformers/models/auto/auto_mappings.py
+++ b/src/transformers/models/auto/auto_mappings.py
@@ -129,6 +129,7 @@
("deepseek_ocr2_vision", "DeepseekOcr2VisionConfig"),
("deepseek_v2", "DeepseekV2Config"),
("deepseek_v3", "DeepseekV3Config"),
+ ("deepseek_v32", "DeepseekV32Config"),
("deepseek_v4", "DeepseekV4Config"),
("deepseek_vl", "DeepseekVLConfig"),
("deepseek_vl_hybrid", "DeepseekVLHybridConfig"),
diff --git a/src/transformers/models/auto/modeling_auto.py b/src/transformers/models/auto/modeling_auto.py
index af5e01d83342..f8501f3e59ea 100644
--- a/src/transformers/models/auto/modeling_auto.py
+++ b/src/transformers/models/auto/modeling_auto.py
@@ -116,6 +116,7 @@ class _BaseModelWithGenerate(PreTrainedModel, GenerationMixin):
("deepseek_ocr2", "DeepseekOcr2Model"),
("deepseek_v2", "DeepseekV2Model"),
("deepseek_v3", "DeepseekV3Model"),
+ ("deepseek_v32", "DeepseekV32Model"),
("deepseek_v4", "DeepseekV4Model"),
("deepseek_vl", "DeepseekVLModel"),
("deepseek_vl_hybrid", "DeepseekVLHybridModel"),
@@ -667,6 +668,7 @@ class _BaseModelWithGenerate(PreTrainedModel, GenerationMixin):
("dbrx", "DbrxForCausalLM"),
("deepseek_v2", "DeepseekV2ForCausalLM"),
("deepseek_v3", "DeepseekV3ForCausalLM"),
+ ("deepseek_v32", "DeepseekV32ForCausalLM"),
("deepseek_v4", "DeepseekV4ForCausalLM"),
("diffllama", "DiffLlamaForCausalLM"),
("doge", "DogeForCausalLM"),
diff --git a/src/transformers/models/auto/tokenization_auto.py b/src/transformers/models/auto/tokenization_auto.py
index 03ea80134745..4175ea39d5b7 100644
--- a/src/transformers/models/auto/tokenization_auto.py
+++ b/src/transformers/models/auto/tokenization_auto.py
@@ -355,6 +355,7 @@
"chatlm",
"deepseek_v2",
"deepseek_v3",
+ "deepseek_v32",
"deepseek_v4",
"deepseek_vl",
"deepseek_vl_hybrid",
diff --git a/src/transformers/models/deepseek_v3/modeling_deepseek_v3.py b/src/transformers/models/deepseek_v3/modeling_deepseek_v3.py
index fe3acd9aeddd..9ca55dad5580 100644
--- a/src/transformers/models/deepseek_v3/modeling_deepseek_v3.py
+++ b/src/transformers/models/deepseek_v3/modeling_deepseek_v3.py
@@ -319,9 +319,12 @@ def eager_attention_forward(
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.
+ Applies interleaved Rotary Position Embedding to the query and key tensors.
+
+ DeepSeek lays the rotary dimensions out in interleaved pairs `(x0, x1), (x2, x3), ...`, each rotated by a
+ single frequency. We compute that rotation directly on the even/odd slices instead of de-interleaving with a
+ `view`/`transpose`/`reshape`; the output is bit-identical to the de-interleaved `rotate_half` formulation while
+ avoiding the extra contiguous copy.
Args:
q (`torch.Tensor`): The query tensor.
@@ -341,17 +344,15 @@ def apply_rotary_pos_emb_interleave(q, k, cos, sin, position_ids=None, unsqueeze
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)
+ # `cos`/`sin` are `cat(freqs, freqs)`; the first half holds the per-pair angle.
+ cos = cos[..., : cos.shape[-1] // 2].unsqueeze(unsqueeze_dim)
+ sin = sin[..., : sin.shape[-1] // 2].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)
+ q1, q2 = q[..., 0::2], q[..., 1::2]
+ k1, k2 = k[..., 0::2], k[..., 1::2]
- 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)
+ q_embed = torch.cat([q1 * cos - q2 * sin, q2 * cos + q1 * sin], dim=-1)
+ k_embed = torch.cat([k1 * cos - k2 * sin, k2 * cos + k1 * sin], dim=-1)
return q_embed, k_embed
diff --git a/src/transformers/models/deepseek_v3/modular_deepseek_v3.py b/src/transformers/models/deepseek_v3/modular_deepseek_v3.py
index 2bf7d347e85d..9117f48eff50 100644
--- a/src/transformers/models/deepseek_v3/modular_deepseek_v3.py
+++ b/src/transformers/models/deepseek_v3/modular_deepseek_v3.py
@@ -22,7 +22,6 @@
LlamaRotaryEmbedding,
apply_rotary_pos_emb,
eager_attention_forward,
- rotate_half,
)
from ..mixtral.modeling_mixtral import MixtralExperts
from ..qwen2_moe.modeling_qwen2_moe import Qwen2MoeMLP
@@ -46,9 +45,12 @@ class DeepseekV3MLP(Qwen2MoeMLP):
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.
+ Applies interleaved Rotary Position Embedding to the query and key tensors.
+
+ DeepSeek lays the rotary dimensions out in interleaved pairs `(x0, x1), (x2, x3), ...`, each rotated by a
+ single frequency. We compute that rotation directly on the even/odd slices instead of de-interleaving with a
+ `view`/`transpose`/`reshape`; the output is bit-identical to the de-interleaved `rotate_half` formulation while
+ avoiding the extra contiguous copy.
Args:
q (`torch.Tensor`): The query tensor.
@@ -68,17 +70,15 @@ def apply_rotary_pos_emb_interleave(q, k, cos, sin, position_ids=None, unsqueeze
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)
+ # `cos`/`sin` are `cat(freqs, freqs)`; the first half holds the per-pair angle.
+ cos = cos[..., : cos.shape[-1] // 2].unsqueeze(unsqueeze_dim)
+ sin = sin[..., : sin.shape[-1] // 2].unsqueeze(unsqueeze_dim)
- b, h, s, d = k.shape
- k = k.view(b, h, s, d // 2, 2).transpose(4, 3).reshape(b, h, s, d)
+ q1, q2 = q[..., 0::2], q[..., 1::2]
+ k1, k2 = k[..., 0::2], k[..., 1::2]
- q_embed = (q * cos) + (rotate_half(q) * sin)
- k_embed = (k * cos) + (rotate_half(k) * sin)
+ q_embed = torch.cat([q1 * cos - q2 * sin, q2 * cos + q1 * sin], dim=-1)
+ k_embed = torch.cat([k1 * cos - k2 * sin, k2 * cos + k1 * sin], dim=-1)
return q_embed, k_embed
diff --git a/src/transformers/models/deepseek_v32/__init__.py b/src/transformers/models/deepseek_v32/__init__.py
new file mode 100644
index 000000000000..562bdf58d4ee
--- /dev/null
+++ b/src/transformers/models/deepseek_v32/__init__.py
@@ -0,0 +1,28 @@
+# 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.
+# 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_deepseek_v32 import *
+ from .modeling_deepseek_v32 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/deepseek_v32/configuration_deepseek_v32.py b/src/transformers/models/deepseek_v32/configuration_deepseek_v32.py
new file mode 100644
index 000000000000..5ef27459d13e
--- /dev/null
+++ b/src/transformers/models/deepseek_v32/configuration_deepseek_v32.py
@@ -0,0 +1,143 @@
+# 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨
+# This file was automatically generated from src/transformers/models/deepseek_v32/modular_deepseek_v32.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_deepseek_v32.py file directly. One of our CI enforces this.
+# 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨
+# 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.
+# 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 huggingface_hub.dataclasses import strict
+
+from ...configuration_utils import PreTrainedConfig
+from ...modeling_rope_utils import RotaryEmbeddingConfigMixin
+from ...utils import auto_docstring
+
+
+@auto_docstring(checkpoint="deepseek-ai/DeepSeek-V3.2-Exp")
+@strict
+class DeepseekV32Config(PreTrainedConfig, RotaryEmbeddingConfigMixin):
+ r"""
+ n_group (`int`, *optional*, defaults to 1):
+ Number of groups for routed experts.
+ mlp_layer_types (`list`, *optional*):
+ MLP type pattern for each layer (`"dense"` or `"sparse"`). Defaults to 3 dense + rest sparse.
+ 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`, *optional*, defaults to 64):
+ Number of heads for the indexer projections (DSA).
+ first_k_dense_replace (`int`, *optional*, defaults to 3):
+ Number of leading layers that use a dense MLP; the rest use the MoE block.
+
+ ```python
+ >>> from transformers import DeepseekV32Config, DeepseekV32Model
+
+ >>> # Initializing a DeepSeek-V3.2 configuration
+ >>> configuration = DeepseekV32Config()
+
+ >>> # Initializing a model from the configuration
+ >>> model = DeepseekV32Model(configuration)
+
+ >>> # Accessing the model configuration
+ >>> configuration = model.config
+ ```"""
+
+ model_type = "deepseek_v32"
+ keys_to_ignore_at_inference = ["past_key_values"]
+
+ 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.*.mlp.shared_experts.gate_proj": "colwise",
+ "layers.*.mlp.shared_experts.up_proj": "colwise",
+ "layers.*.mlp.shared_experts.down_proj": "rowwise",
+ "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": "num_experts"}
+
+ vocab_size: int = 129280
+ hidden_size: int = 7168
+ intermediate_size: int = 18432
+ moe_intermediate_size: int = 2048
+ num_hidden_layers: int = 61
+ num_attention_heads: int = 128
+ num_key_value_heads: int = 128
+ n_shared_experts: int = 1
+ n_routed_experts: int = 256
+ routed_scaling_factor: float = 2.5
+ kv_lora_rank: int = 512
+ q_lora_rank: int = 1536
+ qk_rope_head_dim: int = 64
+ v_head_dim: int = 128
+ qk_nope_head_dim: int = 128
+ n_group: int = 8
+ topk_group: int = 4
+ num_experts_per_tok: int = 8
+ norm_topk_prob: bool = True
+ hidden_act: str = "silu"
+ max_position_embeddings: int = 163840
+ initializer_range: float = 0.02
+ rms_norm_eps: float = 1e-6
+ use_cache: bool = True
+ pad_token_id: int | None = None
+ bos_token_id: int | None = 0
+ eos_token_id: int | list[int] | None = 1
+ tie_word_embeddings: bool = False
+ rope_parameters: dict | None = None
+ 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 = 64
+ mlp_bias: bool = False
+ num_experts: int = 256
+ head_dim: int = 64
+ first_k_dense_replace: int = 3
+ layer_types: list[str] | None = None
+
+ def __post_init__(self, **kwargs):
+ self.qk_head_dim = self.qk_nope_head_dim + self.qk_rope_head_dim
+ # RoPE applies only to the rope slice, so point `head_dim` at it: the inherited (Llama) rotary
+ # embedding reads `config.head_dim` and then computes the right frequencies with no override needed.
+ self.head_dim = self.qk_rope_head_dim
+ # MLP layer types: the first `first_k_dense_replace` layers are dense, the rest are MoE.
+ if self.mlp_layer_types is None:
+ n_dense = min(self.first_k_dense_replace, self.num_hidden_layers)
+ self.mlp_layer_types = ["dense"] * n_dense + ["sparse"] * (self.num_hidden_layers - n_dense)
+ # Every layer is DSA — drives cache-class dispatch.
+ if self.layer_types is None:
+ self.layer_types = ["deepseek_sparse_attention"] * self.num_hidden_layers
+ # Default to MoE from the second layer and on
+ if self.mlp_layer_types is None:
+ self.mlp_layer_types = ["dense"] + ["sparse"] * (self.num_hidden_layers - 1)
+ self.qk_head_dim = self.qk_nope_head_dim + self.qk_rope_head_dim
+ super().__post_init__(**kwargs)
+
+
+__all__ = ["DeepseekV32Config"]
diff --git a/src/transformers/models/deepseek_v32/modeling_deepseek_v32.py b/src/transformers/models/deepseek_v32/modeling_deepseek_v32.py
new file mode 100644
index 000000000000..46f6f6c74762
--- /dev/null
+++ b/src/transformers/models/deepseek_v32/modeling_deepseek_v32.py
@@ -0,0 +1,836 @@
+# 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨
+# This file was automatically generated from src/transformers/models/deepseek_v32/modular_deepseek_v32.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_deepseek_v32.py file directly. One of our CI enforces this.
+# 🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨🚨
+# 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.
+# 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
+from collections.abc import Callable
+from typing import Optional
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+
+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
+from ...utils.generic import maybe_autocast, merge_with_config_defaults
+from ...utils.output_capturing import capture_outputs
+from .configuration_deepseek_v32 import DeepseekV32Config
+
+
+@use_kernel_forward_from_hub("RMSNorm")
+class DeepseekV32RMSNorm(nn.Module):
+ def __init__(self, hidden_size, eps: float = 1e-6) -> None:
+ """
+ DeepseekV32RMSNorm is equivalent to T5LayerNorm
+ """
+ super().__init__()
+ self.weight = nn.Parameter(torch.ones(hidden_size))
+ self.variance_epsilon = eps
+
+ 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)
+ 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}"
+
+
+class DeepseekV32RotaryEmbedding(nn.Module):
+ inv_freq: torch.Tensor # fix linting for `register_buffer`
+
+ def __init__(self, config: DeepseekV32Config, 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: DeepseekV32Config | 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"]
+ dim = getattr(config, "head_dim", None) or config.hidden_size // config.num_attention_heads
+
+ 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)
+
+
+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
+
+
+class DeepseekV32Indexer(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,
+ and returns the additive top-k sparse mask directly (`0` at the selected tokens, `-inf` elsewhere);
+ the raw top-k indices are only ever scattered into that mask, so they are not surfaced.
+
+ **Cache strategy**: the indexer key cache lives on the per-layer `DynamicIndexedLayer` (or the
+ `StaticIndexedLayer` for static caches) inside the shared cache, accessed via
+ `past_key_values.update_indexer()`.
+ """
+
+ def __init__(self, config: "DeepseekV32Config", layer_idx: int):
+ super().__init__()
+ self.config = config
+ self.layer_idx = layer_idx
+
+ self.hidden_size: int = config.hidden_size
+ self.n_heads: int = config.index_n_heads
+ 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.wq_b = nn.Linear(self.q_lora_rank, self.n_heads * self.head_dim, bias=False)
+ self.wk = nn.Linear(self.hidden_size, self.head_dim, bias=False)
+ self.k_norm = nn.LayerNorm(self.head_dim, eps=1e-6)
+ self.weights_proj = nn.Linear(self.hidden_size, self.n_heads, bias=False)
+ self.softmax_scale = self.head_dim**-0.5
+
+ @torch.no_grad()
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ q_resid: torch.Tensor,
+ position_embeddings: tuple[torch.Tensor, torch.Tensor],
+ attention_mask: torch.Tensor | None,
+ position_ids: torch.Tensor,
+ past_key_values: Cache | None = None,
+ ) -> torch.Tensor:
+ """
+ Selects the top-k tokens per query for DeepSeek Sparse Attention (DSA).
+
+ This is the bf16 equivalent of the reference Indexer which uses `rotate_activation` (Hadamard transform)
+ and `fp8_index` (FP8 quantized scoring kernel). Since the Hadamard transform is orthogonal (dot products
+ are preserved: Hq·Hk = q·k), and FP8 quantization is a precision optimization, we skip both and compute
+ scores directly in bf16/fp32.
+
+ The scoring logic computes:
+ index_score[b,s,t] = Σ_h (weight[b,s,h] · softmax_scale · q[b,s,h,:] · k[b,t,:])
+
+ Args:
+ hidden_states: Input hidden states `[B, S, hidden_size]`.
+ q_resid: Query residual from `q_a_layernorm(q_a_proj(x))`, shape `[B, S, q_lora_rank]`.
+ position_embeddings: `(cos, sin)` from RotaryEmbedding.
+ attention_mask: Causal mask, broadcastable to `[B, S, T]`.
+ past_key_values: Cache object containing the indexer key cache for this layer.
+
+ Returns:
+ `torch.Tensor`: the `int32` top-k token indices of shape `[B, S, topk]`. The eager / SDPA paths
+ turn these into an additive sparse mask; the `flash-mla` kernel consumes them directly.
+ """
+ batch_size, seq_len, _ = hidden_states.shape
+ cos, sin = position_embeddings
+ 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_rot, q_pass = torch.split(q, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1)
+
+ k = self.k_norm(self.wk(hidden_states)).unsqueeze(2) # [B, S, 1, D]
+ k_rot, k_pass = torch.split(k, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1)
+
+ # The indexer uses NON-interleaved (half-split) RoPE — unlike the main MLA attention
+ q_rot, k_rot = apply_rotary_pos_emb(q_rot, k_rot, cos, sin, unsqueeze_dim=2)
+ q = torch.cat([q_rot, q_pass], dim=-1) # [B, S, H, D]
+ k = torch.cat([k_rot, k_pass], dim=-1).squeeze(2) # [B, S, D]
+
+ if past_key_values is not None:
+ k = past_key_values.update_indexer(k, self.layer_idx)
+
+ scores = torch.matmul(q.float(), k.transpose(-1, -2).float().unsqueeze(1)) * self.softmax_scale
+ scores = F.relu(scores)
+
+ # Weight per head and sum across heads: [B, S, 1, H] @ [B, S, H, T] → [B, S, T]
+ weights = self.weights_proj(hidden_states.to(self.weights_proj.weight.dtype)).float() * (self.n_heads**-0.5)
+ index_scores = torch.matmul(weights.unsqueeze(-2), scores).squeeze(-2)
+
+ # Causality needs to be taken into account when computing scores so padding tokens don't affect computation
+ if attention_mask is not None:
+ index_scores = index_scores + attention_mask
+ else:
+ key_positions = torch.arange(index_scores.shape[-1], device=index_scores.device)
+ causal = key_positions[None, None, :] > position_ids[:, :, None] # [B, S, T]
+ index_scores = index_scores.masked_fill(causal, float("-inf"))
+
+ topk = min(self.index_topk, index_scores.shape[-1])
+ return index_scores.topk(topk, dim=-1).indices.to(torch.int32) # [B, S, topk]
+
+
+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:
+ attn_weights = attn_weights + attention_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"""
+ Applies interleaved Rotary Position Embedding to the query and key tensors.
+
+ DeepSeek lays the rotary dimensions out in interleaved pairs `(x0, x1), (x2, x3), ...`, each rotated by a
+ single frequency. We compute that rotation directly on the even/odd slices instead of de-interleaving with a
+ `view`/`transpose`/`reshape`; the output is bit-identical to the de-interleaved `rotate_half` formulation while
+ avoiding the extra contiguous copy.
+
+ 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`/`sin` are `cat(freqs, freqs)`; the first half holds the per-pair angle.
+ cos = cos[..., : cos.shape[-1] // 2].unsqueeze(unsqueeze_dim)
+ sin = sin[..., : sin.shape[-1] // 2].unsqueeze(unsqueeze_dim)
+
+ q1, q2 = q[..., 0::2], q[..., 1::2]
+ k1, k2 = k[..., 0::2], k[..., 1::2]
+
+ q_embed = torch.cat([q1 * cos - q2 * sin, q2 * cos + q1 * sin], dim=-1)
+ k_embed = torch.cat([k1 * cos - k2 * sin, k2 * cos + k1 * sin], dim=-1)
+ 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 DeepseekV32Attention(nn.Module):
+ """
+ DeepSeek-V3 MLA, with a DSA indexer whose top-k sparse mask is folded into the attention mask.
+ Qlora rank formulation is dropped as it is never used in released models.
+ """
+
+ def __init__(self, config: DeepseekV32Config, 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.is_causal = True
+ 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 = DeepseekV32RMSNorm(config.q_lora_rank)
+ self.q_b_proj = nn.Linear(config.q_lora_rank, self.num_heads * self.qk_head_dim, bias=False)
+
+ 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 = DeepseekV32RMSNorm(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.o_proj = nn.Linear(
+ self.num_heads * self.v_head_dim,
+ config.hidden_size,
+ bias=config.attention_bias,
+ )
+
+ 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 = DeepseekV32Indexer(config, layer_idx)
+
+ 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,
+ position_ids: torch.Tensor | None = None,
+ **kwargs: Unpack[FlashAttentionKwargs],
+ ) -> tuple[torch.Tensor, torch.Tensor | None]:
+ 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)
+
+ q_resid = self.q_a_layernorm(self.q_a_proj(hidden_states))
+ q_states = self.q_b_proj(q_resid).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
+ q_rot, k_rot = apply_rotary_pos_emb_interleave(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:
+ key_states, value_states = past_key_values.update(key_states, value_states, self.layer_idx)
+
+ # The indexer scores against a 3D `[B, S, T]` mask; the attention mask is 4D `[B, 1, S, T]`.
+ indexer_mask = attention_mask[:, 0, :, :] if attention_mask is not None else None
+ topk_indices = self.indexer(
+ hidden_states, q_resid, position_embeddings, indexer_mask, position_ids, past_key_values=past_key_values
+ ) # [B, S, topk]
+
+ sparse_indices = None
+ if self.config._attn_implementation in ("eager", "sdpa"):
+ # Boolean mask: `True` at keys *not* selected by the indexer (to be masked out).
+ index_mask = (
+ topk_indices.new_ones((batch_size, seq_length, key_states.shape[2]), dtype=torch.bool)
+ .scatter(-1, topk_indices.long(), False)
+ .unsqueeze(1)
+ ) # [B, 1, S, T]; True = masked
+ if attention_mask is None:
+ key_positions = torch.arange(key_states.shape[2], device=hidden_states.device)
+ index_mask = index_mask | (key_positions[None, None, None, :] > position_ids[:, None, :, None])
+ attention_mask = hidden_states.new_zeros((batch_size, 1, seq_length, key_states.shape[2]))
+ attention_mask = attention_mask.masked_fill(index_mask, torch.finfo(hidden_states.dtype).min)
+ else:
+ sparse_indices = topk_indices
+
+ attention_interface: Callable = ALL_ATTENTION_FUNCTIONS.get_interface(
+ self.config._attn_implementation, eager_attention_forward
+ )
+ 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,
+ indices=sparse_indices,
+ **kwargs,
+ )
+
+ attn_output = attn_output.reshape(batch_size, seq_length, -1).contiguous()
+ attn_output = self.o_proj(attn_output)
+ return attn_output, attn_weights
+
+
+class DeepseekV32MLP(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 DeepseekV32TopkRouter(nn.Module):
+ def __init__(self, config: DeepseekV32Config):
+ 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 DeepseekV32NaiveMoe(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 DeepseekV32MoE(nn.Module):
+ """
+ A mixed expert module containing shared experts.
+ """
+
+ def __init__(self, config):
+ super().__init__()
+ self.config = config
+ self.experts = DeepseekV32NaiveMoe(config)
+ self.gate = DeepseekV32TopkRouter(config)
+ self.shared_experts = DeepseekV32MLP(
+ 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(), float("-inf"))
+ 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 DeepseekV32DecoderLayer(GradientCheckpointingLayer):
+ def __init__(self, config: DeepseekV32Config, layer_idx: int):
+ super().__init__()
+ self.hidden_size = config.hidden_size
+ self.self_attn = DeepseekV32Attention(config, layer_idx)
+
+ if config.mlp_layer_types[layer_idx] == "sparse":
+ self.mlp = DeepseekV32MoE(config)
+ else:
+ self.mlp = DeepseekV32MLP(config)
+
+ self.input_layernorm = DeepseekV32RMSNorm(config.hidden_size, config.rms_norm_eps)
+ self.post_attention_layernorm = DeepseekV32RMSNorm(config.hidden_size, 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,
+ 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,
+ 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 DeepseekV32PreTrainedModel(PreTrainedModel):
+ config: DeepseekV32Config
+ base_model_prefix = "model"
+ supports_gradient_checkpointing = True
+ _no_split_modules = ["DeepseekV32DecoderLayer"]
+ _skip_keys_device_placement = ["past_key_values"]
+ _supports_flash_attn = False # flash-mla kernels need a bit more work in the way we enable them!
+ _supports_sdpa = True
+ _supports_flex_attn = False
+
+ _can_compile_fullgraph = True
+ _supports_attention_backend = True
+ _can_record_outputs = {
+ "hidden_states": DeepseekV32DecoderLayer,
+ "attentions": DeepseekV32Attention,
+ }
+ _keep_in_fp32_modules_strict = ["e_score_correction_bias"]
+ _keys_to_ignore_on_load_unexpected = [r"model\.layers\.61.*"]
+ _keep_in_fp32_modules = ["indexer.weights_proj"]
+
+ @torch.no_grad()
+ def _init_weights(self, module):
+ super()._init_weights(module)
+ if isinstance(module, DeepseekV32TopkRouter):
+ init.normal_(module.weight, mean=0.0, std=self.config.initializer_range)
+ init.zeros_(module.e_score_correction_bias)
+ elif isinstance(module, DeepseekV32NaiveMoe):
+ 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)
+
+
+@auto_docstring
+class DeepseekV32Model(DeepseekV32PreTrainedModel):
+ def __init__(self, config: DeepseekV32Config):
+ 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(
+ [DeepseekV32DecoderLayer(config, layer_idx) for layer_idx in range(config.num_hidden_layers)]
+ )
+ self.norm = DeepseekV32RMSNorm(config.hidden_size, eps=config.rms_norm_eps)
+ self.rotary_emb = DeepseekV32RotaryEmbedding(config=config)
+ self.gradient_checkpointing = False
+
+ # Initialize weights and apply final processing
+ self.post_init()
+
+ @merge_with_config_defaults
+ @capture_outputs
+ @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,
+ 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 position_ids is None:
+ past_seen_tokens = past_key_values.get_seq_length() if past_key_values is not None else 0
+ position_ids = torch.arange(inputs_embeds.shape[1], device=inputs_embeds.device) + past_seen_tokens
+ position_ids = position_ids.unsqueeze(0)
+
+ causal_mask = create_causal_mask(
+ config=self.config,
+ inputs_embeds=inputs_embeds,
+ attention_mask=attention_mask,
+ 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,
+ **kwargs,
+ )
+
+ hidden_states = self.norm(hidden_states)
+ return BaseModelOutputWithPast(
+ last_hidden_state=hidden_states,
+ past_key_values=past_key_values,
+ )
+
+
+@auto_docstring
+class DeepseekV32ForCausalLM(DeepseekV32PreTrainedModel, GenerationMixin):
+ _tied_weights_keys = {"lm_head.weight": "model.embed_tokens.weight"}
+ _tp_plan = {"lm_head": "colwise_gather_output"}
+ _pp_plan = {"lm_head": (["hidden_states"], ["logits"])}
+
+ def __init__(self, config):
+ super().__init__(config)
+ self.model = DeepseekV32Model(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,
+ logits_to_keep: int | torch.Tensor = 0,
+ **kwargs: Unpack[TransformersKwargs],
+ ) -> CausalLMOutputWithPast:
+ r"""
+ Example:
+
+ ```python
+ >>> from transformers import AutoTokenizer, DeepseekV32ForCausalLM
+
+ >>> model = DeepseekV32ForCausalLM.from_pretrained("meta-deepseek_v32/DeepseekV32-2-7b-hf")
+ >>> tokenizer = AutoTokenizer.from_pretrained("meta-deepseek_v32/DeepseekV32-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,
+ **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__ = ["DeepseekV32PreTrainedModel", "DeepseekV32Model", "DeepseekV32ForCausalLM"]
diff --git a/src/transformers/models/deepseek_v32/modular_deepseek_v32.py b/src/transformers/models/deepseek_v32/modular_deepseek_v32.py
new file mode 100644
index 000000000000..77be4aa9c943
--- /dev/null
+++ b/src/transformers/models/deepseek_v32/modular_deepseek_v32.py
@@ -0,0 +1,380 @@
+# 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.
+# 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.
+"""DeepSeek-V3.2-Exp: DeepSeek-V3 plus DeepSeek Sparse Attention (DSA).
+
+This is DeepSeek-V3 with a lightning indexer added to each attention layer: the indexer scores every
+query against the cached keys and keeps the top-`index_topk` tokens, which become an additive sparse
+mask folded into the MLA attention mask. Everything else (MoE, MLA projections, RoPE, the decoder /
+model / causal-LM scaffolding) is inherited unchanged from DeepSeek-V3.
+
+The cross-layer top-k *sharing* variant is a GLM-MoE-DSA innovation and lives in that model, which
+inherits from this one (see `models/glm_moe_dsa/modular_glm_moe_dsa.py`).
+"""
+
+from collections.abc import Callable
+
+import torch
+import torch.nn as nn
+import torch.nn.functional as F
+from huggingface_hub.dataclasses import strict
+
+from ...cache_utils import Cache
+from ...modeling_flash_attention_utils import FlashAttentionKwargs
+from ...modeling_rope_utils import RotaryEmbeddingConfigMixin
+from ...modeling_utils import ALL_ATTENTION_FUNCTIONS
+from ...processing_utils import Unpack
+from ...utils import auto_docstring, logging
+from ..deepseek_v3.modeling_deepseek_v3 import (
+ DeepseekV3Attention,
+ DeepseekV3ForCausalLM,
+ DeepseekV3Model,
+ DeepseekV3PreTrainedModel,
+ DeepseekV3RMSNorm,
+ DeepseekV3RotaryEmbedding,
+ apply_rotary_pos_emb,
+ apply_rotary_pos_emb_interleave,
+ 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__)
+
+
+@auto_docstring(checkpoint="deepseek-ai/DeepSeek-V3.2-Exp")
+@strict
+class DeepseekV32Config(Glm4MoeLiteConfig, RotaryEmbeddingConfigMixin):
+ r"""
+ n_group (`int`, *optional*, defaults to 1):
+ Number of groups for routed experts.
+ mlp_layer_types (`list`, *optional*):
+ MLP type pattern for each layer (`"dense"` or `"sparse"`). Defaults to 3 dense + rest sparse.
+ 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`, *optional*, defaults to 64):
+ Number of heads for the indexer projections (DSA).
+ first_k_dense_replace (`int`, *optional*, defaults to 3):
+ Number of leading layers that use a dense MLP; the rest use the MoE block.
+
+ ```python
+ >>> from transformers import DeepseekV32Config, DeepseekV32Model
+
+ >>> # Initializing a DeepSeek-V3.2 configuration
+ >>> configuration = DeepseekV32Config()
+
+ >>> # Initializing a model from the configuration
+ >>> model = DeepseekV32Model(configuration)
+
+ >>> # Accessing the model configuration
+ >>> configuration = model.config
+ ```"""
+
+ 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.*.mlp.shared_experts.gate_proj": "colwise",
+ "layers.*.mlp.shared_experts.up_proj": "colwise",
+ "layers.*.mlp.shared_experts.down_proj": "rowwise",
+ "layers.*.mlp.gate_proj": "colwise",
+ "layers.*.mlp.up_proj": "colwise",
+ "layers.*.mlp.down_proj": "rowwise",
+ }
+
+ attribute_map = {"num_local_experts": "num_experts"}
+
+ vocab_size: int = 129280
+ hidden_size: int = 7168
+ intermediate_size: int = 18432
+ moe_intermediate_size: int = 2048
+ num_hidden_layers: int = 61
+ num_attention_heads: int = 128
+ num_key_value_heads: int = 128
+ n_shared_experts: int = 1
+ n_routed_experts: int = 256
+ routed_scaling_factor: float = 2.5
+ kv_lora_rank: int = 512
+ q_lora_rank: int = 1536
+ qk_rope_head_dim: int = 64
+ v_head_dim: int = 128
+ qk_nope_head_dim: int = 128
+ n_group: int = 8
+ topk_group: int = 4
+ num_experts_per_tok: int = 8
+ norm_topk_prob: bool = True
+ hidden_act: str = "silu"
+ max_position_embeddings: int = 163840
+ initializer_range: float = 0.02
+ rms_norm_eps: float = 1e-6
+ use_cache: bool = True
+ pad_token_id: int | None = None
+ bos_token_id: int | None = 0
+ eos_token_id: int | list[int] | None = 1
+ tie_word_embeddings: bool = False
+ rope_parameters: dict | None = None
+ 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 = 64
+ mlp_bias: bool = False
+ num_experts: int = 256
+ head_dim: int = 64
+ first_k_dense_replace: int = 3
+ pretraining_tp = AttributeError()
+ rope_interleave = AttributeError()
+ layer_types: list[str] | None = None
+
+ def __post_init__(self, **kwargs):
+ self.qk_head_dim = self.qk_nope_head_dim + self.qk_rope_head_dim
+ # RoPE applies only to the rope slice, so point `head_dim` at it: the inherited (Llama) rotary
+ # embedding reads `config.head_dim` and then computes the right frequencies with no override needed.
+ self.head_dim = self.qk_rope_head_dim
+ # MLP layer types: the first `first_k_dense_replace` layers are dense, the rest are MoE.
+ if self.mlp_layer_types is None:
+ n_dense = min(self.first_k_dense_replace, self.num_hidden_layers)
+ self.mlp_layer_types = ["dense"] * n_dense + ["sparse"] * (self.num_hidden_layers - n_dense)
+ # Every layer is DSA — drives cache-class dispatch.
+ if self.layer_types is None:
+ self.layer_types = ["deepseek_sparse_attention"] * self.num_hidden_layers
+ super().__post_init__(**kwargs)
+
+
+class DeepseekV32RMSNorm(DeepseekV3RMSNorm):
+ pass
+
+
+class DeepseekV32RotaryEmbedding(DeepseekV3RotaryEmbedding):
+ pass
+
+
+class DeepseekV32Indexer(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,
+ and returns the additive top-k sparse mask directly (`0` at the selected tokens, `-inf` elsewhere);
+ the raw top-k indices are only ever scattered into that mask, so they are not surfaced.
+
+ **Cache strategy**: the indexer key cache lives on the per-layer `DynamicIndexedLayer` (or the
+ `StaticIndexedLayer` for static caches) inside the shared cache, accessed via
+ `past_key_values.update_indexer()`.
+ """
+
+ def __init__(self, config: "DeepseekV32Config", layer_idx: int):
+ super().__init__()
+ self.config = config
+ self.layer_idx = layer_idx
+
+ self.hidden_size: int = config.hidden_size
+ self.n_heads: int = config.index_n_heads
+ 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.wq_b = nn.Linear(self.q_lora_rank, self.n_heads * self.head_dim, bias=False)
+ self.wk = nn.Linear(self.hidden_size, self.head_dim, bias=False)
+ self.k_norm = nn.LayerNorm(self.head_dim, eps=1e-6)
+ self.weights_proj = nn.Linear(self.hidden_size, self.n_heads, bias=False)
+ self.softmax_scale = self.head_dim**-0.5
+
+ @torch.no_grad()
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ q_resid: torch.Tensor,
+ position_embeddings: tuple[torch.Tensor, torch.Tensor],
+ attention_mask: torch.Tensor | None,
+ position_ids: torch.Tensor,
+ past_key_values: Cache | None = None,
+ ) -> torch.Tensor:
+ """
+ Selects the top-k tokens per query for DeepSeek Sparse Attention (DSA).
+
+ This is the bf16 equivalent of the reference Indexer which uses `rotate_activation` (Hadamard transform)
+ and `fp8_index` (FP8 quantized scoring kernel). Since the Hadamard transform is orthogonal (dot products
+ are preserved: Hq·Hk = q·k), and FP8 quantization is a precision optimization, we skip both and compute
+ scores directly in bf16/fp32.
+
+ The scoring logic computes:
+ index_score[b,s,t] = Σ_h (weight[b,s,h] · softmax_scale · q[b,s,h,:] · k[b,t,:])
+
+ Args:
+ hidden_states: Input hidden states `[B, S, hidden_size]`.
+ q_resid: Query residual from `q_a_layernorm(q_a_proj(x))`, shape `[B, S, q_lora_rank]`.
+ position_embeddings: `(cos, sin)` from RotaryEmbedding.
+ attention_mask: Causal mask, broadcastable to `[B, S, T]`.
+ past_key_values: Cache object containing the indexer key cache for this layer.
+
+ Returns:
+ `torch.Tensor`: the `int32` top-k token indices of shape `[B, S, topk]`. The eager / SDPA paths
+ turn these into an additive sparse mask; the `flash-mla` kernel consumes them directly.
+ """
+ batch_size, seq_len, _ = hidden_states.shape
+ cos, sin = position_embeddings
+ 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_rot, q_pass = torch.split(q, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1)
+
+ k = self.k_norm(self.wk(hidden_states)).unsqueeze(2) # [B, S, 1, D]
+ k_rot, k_pass = torch.split(k, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1)
+
+ # The indexer uses NON-interleaved (half-split) RoPE — unlike the main MLA attention
+ q_rot, k_rot = apply_rotary_pos_emb(q_rot, k_rot, cos, sin, unsqueeze_dim=2)
+ q = torch.cat([q_rot, q_pass], dim=-1) # [B, S, H, D]
+ k = torch.cat([k_rot, k_pass], dim=-1).squeeze(2) # [B, S, D]
+
+ if past_key_values is not None:
+ k = past_key_values.update_indexer(k, self.layer_idx)
+
+ scores = torch.matmul(q.float(), k.transpose(-1, -2).float().unsqueeze(1)) * self.softmax_scale
+ scores = F.relu(scores)
+
+ # Weight per head and sum across heads: [B, S, 1, H] @ [B, S, H, T] → [B, S, T]
+ weights = self.weights_proj(hidden_states.to(self.weights_proj.weight.dtype)).float() * (self.n_heads**-0.5)
+ index_scores = torch.matmul(weights.unsqueeze(-2), scores).squeeze(-2)
+
+ # Causality needs to be taken into account when computing scores so padding tokens don't affect computation
+ if attention_mask is not None:
+ index_scores = index_scores + attention_mask
+ else:
+ key_positions = torch.arange(index_scores.shape[-1], device=index_scores.device)
+ causal = key_positions[None, None, :] > position_ids[:, :, None] # [B, S, T]
+ index_scores = index_scores.masked_fill(causal, float("-inf"))
+
+ topk = min(self.index_topk, index_scores.shape[-1])
+ return index_scores.topk(topk, dim=-1).indices.to(torch.int32) # [B, S, topk]
+
+
+class DeepseekV32Attention(DeepseekV3Attention):
+ """
+ DeepSeek-V3 MLA, with a DSA indexer whose top-k sparse mask is folded into the attention mask.
+ Qlora rank formulation is dropped as it is never used in released models.
+ """
+
+ def __init__(self, config: DeepseekV32Config, layer_idx: int):
+ super().__init__(config, layer_idx)
+ self.indexer = DeepseekV32Indexer(config, layer_idx)
+
+ 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,
+ position_ids: torch.Tensor | None = None,
+ **kwargs: Unpack[FlashAttentionKwargs],
+ ) -> tuple[torch.Tensor, torch.Tensor | None]:
+ 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)
+
+ q_resid = self.q_a_layernorm(self.q_a_proj(hidden_states))
+ q_states = self.q_b_proj(q_resid).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
+ q_rot, k_rot = apply_rotary_pos_emb_interleave(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:
+ key_states, value_states = past_key_values.update(key_states, value_states, self.layer_idx)
+
+ # The indexer scores against a 3D `[B, S, T]` mask; the attention mask is 4D `[B, 1, S, T]`.
+ indexer_mask = attention_mask[:, 0, :, :] if attention_mask is not None else None
+ topk_indices = self.indexer(
+ hidden_states, q_resid, position_embeddings, indexer_mask, position_ids, past_key_values=past_key_values
+ ) # [B, S, topk]
+
+ sparse_indices = None
+ if self.config._attn_implementation in ("eager", "sdpa"):
+ # Boolean mask: `True` at keys *not* selected by the indexer (to be masked out).
+ index_mask = (
+ topk_indices.new_ones((batch_size, seq_length, key_states.shape[2]), dtype=torch.bool)
+ .scatter(-1, topk_indices.long(), False)
+ .unsqueeze(1)
+ ) # [B, 1, S, T]; True = masked
+ if attention_mask is None:
+ key_positions = torch.arange(key_states.shape[2], device=hidden_states.device)
+ index_mask = index_mask | (key_positions[None, None, None, :] > position_ids[:, None, :, None])
+ attention_mask = hidden_states.new_zeros((batch_size, 1, seq_length, key_states.shape[2]))
+ attention_mask = attention_mask.masked_fill(index_mask, torch.finfo(hidden_states.dtype).min)
+ else:
+ sparse_indices = topk_indices
+
+ attention_interface: Callable = ALL_ATTENTION_FUNCTIONS.get_interface(
+ self.config._attn_implementation, eager_attention_forward
+ )
+ 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,
+ indices=sparse_indices,
+ **kwargs,
+ )
+
+ attn_output = attn_output.reshape(batch_size, seq_length, -1).contiguous()
+ attn_output = self.o_proj(attn_output)
+ return attn_output, attn_weights
+
+
+class DeepseekV32DecoderLayer(Glm4MoeLiteDecoderLayer):
+ pass
+
+
+class DeepseekV32PreTrainedModel(DeepseekV3PreTrainedModel):
+ _keep_in_fp32_modules = ["indexer.weights_proj"]
+ _keep_in_fp32_modules_strict = ["e_score_correction_bias"]
+ _keys_to_ignore_on_load_unexpected = [r"model\.layers\.61.*"]
+ _supports_flash_attn = False # flash-mla kernels need a bit more work in the way we enable them!
+ _supports_sdpa = True
+ _supports_flex_attn = False
+
+
+class DeepseekV32Model(DeepseekV3Model):
+ pass
+
+
+class DeepseekV32ForCausalLM(DeepseekV3ForCausalLM):
+ pass
+
+
+__all__ = [
+ "DeepseekV32Config",
+ "DeepseekV32PreTrainedModel",
+ "DeepseekV32Model",
+ "DeepseekV32ForCausalLM",
+]
diff --git a/src/transformers/models/glm4_moe_lite/modeling_glm4_moe_lite.py b/src/transformers/models/glm4_moe_lite/modeling_glm4_moe_lite.py
index 0b8ccc865775..e0d7b8fddb14 100644
--- a/src/transformers/models/glm4_moe_lite/modeling_glm4_moe_lite.py
+++ b/src/transformers/models/glm4_moe_lite/modeling_glm4_moe_lite.py
@@ -184,9 +184,12 @@ def eager_attention_forward(
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.
+ Applies interleaved Rotary Position Embedding to the query and key tensors.
+
+ DeepSeek lays the rotary dimensions out in interleaved pairs `(x0, x1), (x2, x3), ...`, each rotated by a
+ single frequency. We compute that rotation directly on the even/odd slices instead of de-interleaving with a
+ `view`/`transpose`/`reshape`; the output is bit-identical to the de-interleaved `rotate_half` formulation while
+ avoiding the extra contiguous copy.
Args:
q (`torch.Tensor`): The query tensor.
@@ -206,17 +209,15 @@ def apply_rotary_pos_emb_interleave(q, k, cos, sin, position_ids=None, unsqueeze
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)
+ # `cos`/`sin` are `cat(freqs, freqs)`; the first half holds the per-pair angle.
+ cos = cos[..., : cos.shape[-1] // 2].unsqueeze(unsqueeze_dim)
+ sin = sin[..., : sin.shape[-1] // 2].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)
+ q1, q2 = q[..., 0::2], q[..., 1::2]
+ k1, k2 = k[..., 0::2], k[..., 1::2]
- 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)
+ q_embed = torch.cat([q1 * cos - q2 * sin, q2 * cos + q1 * sin], dim=-1)
+ k_embed = torch.cat([k1 * cos - k2 * sin, k2 * cos + k1 * sin], dim=-1)
return q_embed, k_embed
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 c96a274e1a4f..5ecf328bd170 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
@@ -21,26 +21,30 @@
from huggingface_hub.dataclasses import strict
from ...configuration_utils import PreTrainedConfig
-from ...modeling_rope_utils import RopeParameters
+from ...modeling_rope_utils import RotaryEmbeddingConfigMixin
from ...utils import auto_docstring
@auto_docstring(checkpoint="zai-org/GLM-5")
@strict
-class GlmMoeDsaConfig(PreTrainedConfig):
+class GlmMoeDsaConfig(PreTrainedConfig, RotaryEmbeddingConfigMixin):
r"""
n_group (`int`, *optional*, defaults to 1):
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.
+ MLP type pattern for each layer (`"dense"` or `"sparse"`). Defaults to 3 dense + rest sparse.
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):
+ index_n_heads (`int`, *optional*, defaults to 32):
Number of heads for the indexer projections (DSA).
+ first_k_dense_replace (`int`, *optional*, defaults to 3):
+ Number of leading layers that use a dense MLP; the rest use the MoE block.
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`.
+ Per-layer indexer mode (`"full"` runs the indexer, `"shared"` reuses the previous full
+ layer's top-k). Defaults to the pattern derived from `index_topk_freq` /
+ `index_skip_topk_offset` (or `index_topk_pattern`).
```python
>>> from transformers import GlmMoeDsaConfig, GlmMoeDsaModel
@@ -84,7 +88,6 @@ class GlmMoeDsaConfig(PreTrainedConfig):
}
vocab_size: int = 154880
-
hidden_size: int = 6144
intermediate_size: int = 12288
moe_intermediate_size: int = 2048
@@ -112,23 +115,23 @@ class GlmMoeDsaConfig(PreTrainedConfig):
bos_token_id: int | None = 0
eos_token_id: int | list[int] | None = 1
tie_word_embeddings: bool = False
- rope_parameters: RopeParameters | dict | None = None
+ rope_parameters: dict | None = None
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
+ mlp_bias: bool = False
+ num_experts: int = 256
+ head_dim: int = 64
+ first_k_dense_replace: int = 3
+ layer_types: list[str] | None = None
+ # `"full"` runs the indexer, `"shared"` reuses the previous full layer's index mask.
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 % moe_layer_freq == 0 else "dense" for i in range(self.num_hidden_layers)
- ]
-
+ # Per-layer indexer mode: a pattern (e.g. `"FSSF..."`) overrides the freq/offset schedule.
if self.indexer_types is None:
pattern = kwargs.get("index_topk_pattern")
if pattern is not None:
@@ -141,6 +144,21 @@ def __post_init__(self, **kwargs):
self.indexer_types = [
"full" if (max(i - offset + 1, 0) % freq) == 0 else "shared" for i in range(self.num_hidden_layers)
]
+ self.qk_head_dim = self.qk_nope_head_dim + self.qk_rope_head_dim
+ # RoPE applies only to the rope slice, so point `head_dim` at it: the inherited (Llama) rotary
+ # embedding reads `config.head_dim` and then computes the right frequencies with no override needed.
+ self.head_dim = self.qk_rope_head_dim
+ # MLP layer types: the first `first_k_dense_replace` layers are dense, the rest are MoE.
+ if self.mlp_layer_types is None:
+ n_dense = min(self.first_k_dense_replace, self.num_hidden_layers)
+ self.mlp_layer_types = ["dense"] * n_dense + ["sparse"] * (self.num_hidden_layers - n_dense)
+ # Every layer is DSA — drives cache-class dispatch.
+ if self.layer_types is None:
+ self.layer_types = ["deepseek_sparse_attention"] * self.num_hidden_layers
+ # Default to MoE from the second layer and on
+ if self.mlp_layer_types is None:
+ self.mlp_layer_types = ["dense"] + ["sparse"] * (self.num_hidden_layers - 1)
+ self.qk_head_dim = self.qk_nope_head_dim + self.qk_rope_head_dim
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 d6a7d5a16859..34d7bf139c36 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,17 +18,19 @@
# See the License for the specific language governing permissions and
# limitations under the License.
+import math
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
from ...cache_utils import Cache, DynamicCache
from ...generation import GenerationMixin
-from ...integrations import use_experts_implementation, use_kernel_forward_from_hub
+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
@@ -37,7 +39,7 @@
from ...modeling_utils import ALL_ATTENTION_FUNCTIONS, PreTrainedModel
from ...processing_utils import Unpack
from ...utils import TransformersKwargs, auto_docstring, can_return_tuple
-from ...utils.generic import is_flash_attention_requested, maybe_autocast, merge_with_config_defaults
+from ...utils.generic import maybe_autocast, merge_with_config_defaults
from ...utils.output_capturing import capture_outputs
from .configuration_glm_moe_dsa import GlmMoeDsaConfig
@@ -63,6 +65,71 @@ def extra_repr(self):
return f"{tuple(self.weight.shape)}, eps={self.variance_epsilon}"
+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"]
+ dim = getattr(config, "head_dim", None) or config.hidden_size // config.num_attention_heads
+
+ 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)
+
+
def rotate_half(x):
"""Rotates half the hidden dims of the input."""
x1 = x[..., : x.shape[-1] // 2]
@@ -70,49 +137,43 @@ def rotate_half(x):
return torch.cat((-x2, x1), dim=-1)
-def apply_rotary_pos_emb(
- x: torch.Tensor,
- cos: torch.Tensor,
- sin: torch.Tensor,
- unsqueeze_dim: int = 1,
-) -> torch.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=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`.
+@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:
- x (`torch.Tensor`): Input tensor of shape `[..., head_dim]`.
- cos (`torch.Tensor`): Cosine part from RotaryEmbedding, shape `[batch, seq_len, head_dim]`.
- sin (`torch.Tensor`): Sine part from RotaryEmbedding, shape `[batch, seq_len, head_dim]`.
- unsqueeze_dim (`int`): Dimension along which to unsqueeze cos/sin for broadcasting.
- Use `1` when x is `[B, H, S, D]` (BHSD) and `2` when x is `[B, S, H, D]` (BSHD).
-
+ 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:
- `torch.Tensor`: Tensor with rotary embeddings applied, same shape as input.
+ `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)
- *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)
+ q_embed = (q * cos) + (rotate_half(q) * sin)
+ k_embed = (k * cos) + (rotate_half(k) * sin)
+ return q_embed, k_embed
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 uses the same interleaved pair layout as main MLA
- attention.
+ The Indexer has its own lightweight projections (wq_b, wk) separate from the main MLA attention,
+ and returns the additive top-k sparse mask directly (`0` at the selected tokens, `-inf` elsewhere);
+ the raw top-k indices are only ever scattered into that mask, so they are not surfaced.
- **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
- `num_hidden_layers` attention layers. Keys are concatenated along the sequence dimension
- during autoregressive decode.
+ **Cache strategy**: the indexer key cache lives on the per-layer `DynamicIndexedLayer` (or the
+ `StaticIndexedLayer` for static caches) inside the shared cache, accessed via
+ `past_key_values.update_indexer()`.
"""
def __init__(self, config: "GlmMoeDsaConfig", layer_idx: int):
@@ -127,30 +188,24 @@ def __init__(self, config: "GlmMoeDsaConfig", layer_idx: int):
self.index_topk: int = config.index_topk
self.q_lora_rank: int = config.q_lora_rank
- # Named to match checkpoint: wq_b, wk, k_norm
self.wq_b = nn.Linear(self.q_lora_rank, self.n_heads * self.head_dim, bias=False)
self.wk = nn.Linear(self.hidden_size, self.head_dim, bias=False)
self.k_norm = nn.LayerNorm(self.head_dim, eps=1e-6)
- # Named to match checkpoint: weights_proj
- # In the reference, this is fp32; the HF FP8 checkpoint stores a bf16 tensor.
- # Keeping it as a plain Linear prevents FP8 conversion (see `_keep_in_fp32_modules`).
self.weights_proj = nn.Linear(self.hidden_size, self.n_heads, bias=False)
self.softmax_scale = self.head_dim**-0.5
- # Indexer maintains its own key cache (not in DynamicCache, which is sized for attention layers only)
- self.register_buffer("_cached_keys", None, persistent=False)
-
@torch.no_grad()
def forward(
self,
- hidden_states: torch.Tensor, # [B, S, hidden]
- q_resid: torch.Tensor, # [B, S, q_lora_rank]
+ hidden_states: torch.Tensor,
+ q_resid: torch.Tensor,
position_embeddings: tuple[torch.Tensor, torch.Tensor],
attention_mask: torch.Tensor | None,
- use_cache: bool = False,
- ) -> torch.LongTensor:
+ position_ids: torch.Tensor,
+ past_key_values: Cache | None = None,
+ ) -> torch.Tensor:
"""
- Computes top-k token indices for sparse attention (DSA).
+ Selects the top-k tokens per query for DeepSeek Sparse Attention (DSA).
This is the bf16 equivalent of the reference Indexer which uses `rotate_activation` (Hadamard transform)
and `fp8_index` (FP8 quantized scoring kernel). Since the Hadamard transform is orthogonal (dot products
@@ -165,67 +220,46 @@ def forward(
q_resid: Query residual from `q_a_layernorm(q_a_proj(x))`, shape `[B, S, q_lora_rank]`.
position_embeddings: `(cos, sin)` from RotaryEmbedding.
attention_mask: Causal mask, broadcastable to `[B, S, T]`.
- use_cache: Whether to store/update the indexer's own key cache for autoregressive decode.
+ past_key_values: Cache object containing the indexer key cache for this layer.
Returns:
- `torch.LongTensor`: Top-k token indices of shape `[B, S, topk]`.
+ `torch.Tensor`: the `int32` top-k token indices of shape `[B, S, topk]`. The eager / SDPA paths
+ turn these into an additive sparse mask; the `flash-mla` kernel consumes them directly.
"""
batch_size, seq_len, _ = hidden_states.shape
cos, sin = position_embeddings
-
- # === Queries ===
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 = 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 = torch.cat([k_pe, k_nope], dim=-1) # [B, S, D]
-
- # === Key cache (managed by the indexer, not DynamicCache) ===
- # Reset cache on prefill (new prompt) to avoid stale keys / batch-size mismatch
- if seq_len > 1:
- self._cached_keys = None
-
- if use_cache:
- if self._cached_keys is not None:
- k_cached = torch.cat([self._cached_keys, k], dim=1) # [B, T, D]
- else:
- k_cached = k
- self._cached_keys = k_cached
- else:
- k_cached = k
-
- # === Scoring ===
- # Reference: weights = weights_proj(x.float()) * n_heads^(-0.5)
- # Reference: weights = weights.unsqueeze(-1) * q_scale * softmax_scale
- # Reference: index_score = fp8_index(q_fp8, weights, k_cache, k_scale_cache)
- #
- # In bf16 mode (no FP8), q_scale = 1. The fp8_index kernel computes:
- # score[b,s,t] = sum_h(weights[b,s,h] * dot(q[b,s,h,:], k[b,t,:]))
- # where weights already absorbs n_heads^(-0.5) and softmax_scale.
-
- # Don't force fp32 inputs here: the checkpoint stores `weights_proj.weight` in bf16.
- # Use native dtype for matmul, then upcast the result for scoring stability.
- weights = self.weights_proj(hidden_states).float() * (self.n_heads**-0.5) # [B, S, H]
-
- # q·k^T per head: [B, S, H, D] @ [B, T, D]^T → [B, S, H, T]
- scores = torch.einsum("bshd,btd->bsht", q.float(), k_cached.float()) * self.softmax_scale
+ q_rot, q_pass = torch.split(q, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1)
+
+ k = self.k_norm(self.wk(hidden_states)).unsqueeze(2) # [B, S, 1, D]
+ k_rot, k_pass = torch.split(k, [self.qk_rope_head_dim, self.head_dim - self.qk_rope_head_dim], dim=-1)
+
+ # The indexer uses NON-interleaved (half-split) RoPE — unlike the main MLA attention
+ q_rot, k_rot = apply_rotary_pos_emb(q_rot, k_rot, cos, sin, unsqueeze_dim=2)
+ q = torch.cat([q_rot, q_pass], dim=-1) # [B, S, H, D]
+ k = torch.cat([k_rot, k_pass], dim=-1).squeeze(2) # [B, S, D]
+
+ if past_key_values is not None:
+ k = past_key_values.update_indexer(k, self.layer_idx)
+
+ scores = torch.matmul(q.float(), k.transpose(-1, -2).float().unsqueeze(1)) * self.softmax_scale
scores = F.relu(scores)
- # Weight per head and sum across heads → [B, S, T]
- index_scores = torch.einsum("bsht,bsh->bst", scores, weights)
+ # Weight per head and sum across heads: [B, S, 1, H] @ [B, S, H, T] → [B, S, T]
+ weights = self.weights_proj(hidden_states.to(self.weights_proj.weight.dtype)).float() * (self.n_heads**-0.5)
+ index_scores = torch.matmul(weights.unsqueeze(-2), scores).squeeze(-2)
+
+ # Causality needs to be taken into account when computing scores so padding tokens don't affect computation
if attention_mask is not None:
index_scores = index_scores + attention_mask
+ else:
+ key_positions = torch.arange(index_scores.shape[-1], device=index_scores.device)
+ causal = key_positions[None, None, :] > position_ids[:, :, None] # [B, S, T]
+ index_scores = index_scores.masked_fill(causal, float("-inf"))
- total_len = index_scores.shape[-1]
- topk = min(self.index_topk, total_len)
- topk_indices = index_scores.topk(topk, dim=-1).indices # [B, S, topk]
- return topk_indices
+ topk = min(self.index_topk, index_scores.shape[-1])
+ return index_scores.topk(topk, dim=-1).indices.to(torch.int32) # [B, S, topk]
def repeat_kv(hidden_states: torch.Tensor, n_rep: int) -> torch.Tensor:
@@ -265,27 +299,59 @@ def eager_attention_forward(
return attn_output, attn_weights
-class GlmMoeDsaAttention(nn.Module):
+def apply_rotary_pos_emb_interleave(q, k, cos, sin, position_ids=None, unsqueeze_dim=1):
+ r"""
+ Applies interleaved Rotary Position Embedding to the query and key tensors.
+
+ DeepSeek lays the rotary dimensions out in interleaved pairs `(x0, x1), (x2, x3), ...`, each rotated by a
+ single frequency. We compute that rotation directly on the even/odd slices instead of de-interleaving with a
+ `view`/`transpose`/`reshape`; the output is bit-identical to the de-interleaved `rotate_half` formulation while
+ avoiding the extra contiguous copy.
+
+ 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.
"""
- Multi-head Latent Attention (MLA) with DeepSeek Sparse Attention (DSA) indexer.
+ # `cos`/`sin` are `cat(freqs, freqs)`; the first half holds the per-pair angle.
+ cos = cos[..., : cos.shape[-1] // 2].unsqueeze(unsqueeze_dim)
+ sin = sin[..., : sin.shape[-1] // 2].unsqueeze(unsqueeze_dim)
+
+ q1, q2 = q[..., 0::2], q[..., 1::2]
+ k1, k2 = k[..., 0::2], k[..., 1::2]
+
+ q_embed = torch.cat([q1 * cos - q2 * sin, q2 * cos + q1 * sin], dim=-1)
+ k_embed = torch.cat([k1 * cos - k2 * sin, k2 * cos + k1 * sin], dim=-1)
+ return q_embed, k_embed
- This follows the same architecture as DeepSeek V3.2's MLA:
- - Query: x → q_a_proj → RMSNorm → q_b_proj → split(q_nope, q_pe) → RoPE(q_pe)
- - KV: x → kv_a_proj → split(kv_compressed, k_pe) → RMSNorm(kv_compressed) → kv_b_proj
- → RoPE(k_pe)
- - Cache: fully expanded key_states [B, H, T, qk_head_dim] and value_states [B, H, T, v_head_dim]
- - Indexer: selects top-k tokens via DSA, applied as an additive -inf mask on attention scores
- **Caching strategy**: follows the DeepSeek V3 transformers convention of fully expanding K/V
- before caching. This ensures compatibility with DynamicCache, StaticCache, flash attention,
- 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.
+def yarn_get_mscale(scale=1, mscale=1):
+ if scale <= 1:
+ return 1.0
+ return 0.1 * mscale * math.log(scale) + 1.0
- **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.
+class GlmMoeDsaAttention(nn.Module):
+ """
+ DeepSeek-V3 MLA + a DSA indexer, extended with **cross-layer top-k sharing**.
+
+ `config.indexer_types[layer_idx]` decides whether this layer runs its own indexer (`"full"`) or
+ reuses the previous full layer's top-k selection (`"shared"`).
+ `next_skip_topk` signals that the *next* layer will reuse this
+ layer's top-k, so it is propagated upward via `prev_topk_indices`.
"""
def __init__(self, config: GlmMoeDsaConfig, layer_idx: int):
@@ -304,8 +370,6 @@ def __init__(self, config: GlmMoeDsaConfig, layer_idx: int):
self.qk_head_dim = config.qk_head_dim
self.is_causal = True
-
- # Query projection (with optional LoRA)
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:
@@ -313,7 +377,6 @@ def __init__(self, config: GlmMoeDsaConfig, layer_idx: int):
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 (MLA compressed path)
self.kv_a_proj_with_mqa = nn.Linear(
config.hidden_size,
self.kv_lora_rank + self.qk_rope_head_dim,
@@ -326,7 +389,6 @@ def __init__(self, config: GlmMoeDsaConfig, layer_idx: int):
bias=False,
)
- # Output projection
self.o_proj = nn.Linear(
self.num_heads * self.v_head_dim,
config.hidden_size,
@@ -334,15 +396,14 @@ def __init__(self, config: GlmMoeDsaConfig, layer_idx: int):
)
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
# Refer: https://arxiv.org/abs/2603.12201 for more details.
- # 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(
@@ -351,117 +412,83 @@ def forward(
position_embeddings: tuple[torch.Tensor, torch.Tensor],
attention_mask: torch.Tensor | None,
past_key_values: Cache | None = None,
+ position_ids: torch.Tensor | None = None,
prev_topk_indices: torch.Tensor | None = None,
**kwargs: Unpack[FlashAttentionKwargs],
- ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]:
+ ) -> tuple[torch.Tensor, torch.Tensor | None, torch.Tensor | None]:
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)
+
+ q_resid = self.q_a_layernorm(self.q_a_proj(hidden_states))
+ q_states = self.q_b_proj(q_resid).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
+ q_rot, k_rot = apply_rotary_pos_emb_interleave(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)
- # ===== Query path =====
- if self.q_lora_rank is None:
- query_states = self.q_proj(hidden_states)
- q_resid = None
- else:
- q_resid = self.q_a_layernorm(self.q_a_proj(hidden_states)) # [B, S, q_lora_rank]
- query_states = self.q_b_proj(q_resid)
- 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)
-
- # ===== KV path =====
- compressed_kv = self.kv_a_proj_with_mqa(hidden_states) # [B, S, kv_rank + rope_D]
- k_compressed, k_pe = torch.split(compressed_kv, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1)
- k_compressed = self.kv_a_layernorm(k_compressed) # [B, S, kv_rank]
-
- # Expand KV through kv_b_proj
- kv_expanded = self.kv_b_proj(k_compressed) # [B, S, H * (nope_D + v_D)]
- kv_expanded = kv_expanded.view(batch_size, seq_length, -1, self.qk_nope_head_dim + self.v_head_dim)
- k_nope, value_states = torch.split(kv_expanded, [self.qk_nope_head_dim, self.v_head_dim], dim=-1)
- 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), using interleaved pair layout.
- k_pe = k_pe.view(batch_size, 1, seq_length, self.qk_rope_head_dim) # [B, 1, S, rope_D]
- 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
- query_states = torch.cat([q_nope, q_pe], dim=-1) # [B, H, S, qk_head_dim]
- key_states = torch.cat([k_nope, k_pe], dim=-1) # [B, H, S, qk_head_dim]
-
- # Cache update
if past_key_values is not None:
key_states, value_states = past_key_values.update(key_states, value_states, self.layer_idx)
- # ===== 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
- else attention_mask.unsqueeze(1)
- if attention_mask is not None
- else None
- )
+ # DSA: select this layer's top-k tokens, or reuse the previous full layer's on `"shared"` layers.
+ if self.indexer is not None:
+ indexer_mask = attention_mask[:, 0, :, :] if attention_mask is not None else None
topk_indices = self.indexer(
hidden_states,
q_resid,
position_embeddings,
indexer_mask,
- use_cache=past_key_values is not None,
+ position_ids,
+ past_key_values=past_key_values,
) # [B, S, topk]
else:
- topk_indices = prev_topk_indices # [B, S, topk]
-
- # Build combined DSA + causal mask: -inf everywhere except selected top-k positions
- total_len = key_states.shape[2]
- index_mask = torch.full(
- (batch_size, seq_length, total_len),
- float("-inf"),
- device=hidden_states.device,
- dtype=query_states.dtype,
- )
- index_mask.scatter_(-1, topk_indices, 0.0) # [B, S, T]
- index_mask = index_mask.unsqueeze(1) # [B, 1, S, T]
- if attention_mask is not None and attention_mask.dim() == 4:
- causal_mask = attention_mask[..., :total_len]
- combined_mask = index_mask + causal_mask
- else:
- combined_mask = (
- attention_mask.masked_fill(index_mask == float("-inf"), float("-inf"))
- if attention_mask is not None
- else index_mask
+ if prev_topk_indices is None:
+ raise ValueError("Shared DSA layers require top-k indices from a previous full indexer layer.")
+ topk_indices = prev_topk_indices
+
+ sparse_indices = None
+ if self.config._attn_implementation in ("eager", "sdpa"):
+ index_mask = (
+ topk_indices.new_ones((batch_size, seq_length, key_states.shape[2]), dtype=torch.bool)
+ .scatter(-1, topk_indices.long(), False)
+ .unsqueeze(1)
)
-
- # Flash attention head_dim padding (qk_head_dim != v_head_dim)
- if is_flash_attention_requested(self.config) and self.qk_head_dim != self.v_head_dim:
- value_states = F.pad(value_states, [0, self.qk_head_dim - self.v_head_dim])
+ if attention_mask is None:
+ key_positions = torch.arange(key_states.shape[2], device=hidden_states.device)
+ index_mask = index_mask | (key_positions[None, None, None, :] > position_ids[:, None, :, None])
+ attention_mask = hidden_states.new_zeros((batch_size, 1, seq_length, key_states.shape[2]))
+ attention_mask = attention_mask.masked_fill(index_mask, torch.finfo(hidden_states.dtype).min)
+ else:
+ sparse_indices = topk_indices
attention_interface: Callable = ALL_ATTENTION_FUNCTIONS.get_interface(
self.config._attn_implementation, eager_attention_forward
)
-
attn_output, attn_weights = attention_interface(
self,
query_states,
key_states,
value_states,
- combined_mask,
+ attention_mask,
dropout=0.0 if not self.training else self.attention_dropout,
scaling=self.scaling,
- indices=topk_indices, # flash_mla_with_kvcache
+ indices=sparse_indices, # consumed by flash_mla_with_kvcache; ignored by eager / SDPA
**kwargs,
)
- if is_flash_attention_requested(self.config) 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, topk_indices if self.next_skip_topk else None
+ return attn_output, attn_weights, topk_indices
class GlmMoeDsaMLP(nn.Module):
@@ -618,7 +645,7 @@ def forward(
past_key_values: Cache | None = None,
use_cache: bool | None = False,
position_embeddings: tuple[torch.Tensor, torch.Tensor] | None = None,
- prev_topk_indices: torch.Tensor | None = None,
+ prev_topk_indices: torch.Tensor | None = None, # MAIN DIFF with DSV3.2
**kwargs: Unpack[TransformersKwargs],
) -> tuple[torch.Tensor, torch.Tensor | None]:
residual = hidden_states
@@ -631,12 +658,11 @@ def forward(
past_key_values=past_key_values,
use_cache=use_cache,
position_embeddings=position_embeddings,
- prev_topk_indices=prev_topk_indices,
+ prev_topk_indices=prev_topk_indices, # MAIN DIFF with DSV3.2
**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)
@@ -663,10 +689,7 @@ class GlmMoeDsaPreTrainedModel(PreTrainedModel):
}
_keep_in_fp32_modules_strict = ["e_score_correction_bias"]
_keys_to_ignore_on_load_unexpected = [r"model\.layers\.78.*"]
- # NOTE: FP8 quantization uses `_keep_in_fp32_modules` (not `_strict`) to decide which modules to NOT convert.
- # We must keep `indexer.weights_proj` as a plain Linear to match the checkpoint (no `weight_scale_inv`).
_keep_in_fp32_modules = ["indexer.weights_proj"]
- _compatible_flash_implementations = ["kernels-community/flash-mla"]
@torch.no_grad()
def _init_weights(self, module):
@@ -679,72 +702,6 @@ def _init_weights(self, module):
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=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"]
- 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
-
- @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):
def __init__(self, config: GlmMoeDsaConfig):
@@ -801,7 +758,7 @@ def forward(
hidden_states = inputs_embeds
position_embeddings = self.rotary_emb(hidden_states, position_ids=position_ids)
- topk_indices = None
+ topk_indices = None # MAIN DIFF with DSV3.2
for decoder_layer in self.layers[: self.config.num_hidden_layers]:
hidden_states, topk_indices = decoder_layer(
hidden_states,
@@ -810,7 +767,7 @@ def forward(
position_ids=position_ids,
past_key_values=past_key_values,
use_cache=use_cache,
- prev_topk_indices=topk_indices,
+ prev_topk_indices=topk_indices, # MAIN DIFF with DSV3.2
**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 28d8dd6e17f6..a4e91beb4633 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
@@ -15,84 +15,55 @@
from collections.abc import Callable
import torch
-import torch.nn as nn
-import torch.nn.functional as F
from huggingface_hub.dataclasses import strict
from ...cache_utils import Cache, DynamicCache
-from ...configuration_utils import PreTrainedConfig
from ...masking_utils import create_causal_mask
from ...modeling_flash_attention_utils import FlashAttentionKwargs
from ...modeling_outputs import BaseModelOutputWithPast
from ...modeling_utils import ALL_ATTENTION_FUNCTIONS
-from ...models.llama.modeling_llama import rotate_half
from ...processing_utils import Unpack
from ...utils import TransformersKwargs, auto_docstring, logging
-from ...utils.generic import is_flash_attention_requested
-from ..glm4_moe.modeling_glm4_moe import (
- Glm4MoeForCausalLM,
- Glm4MoeModel,
- Glm4MoePreTrainedModel,
- Glm4MoeRMSNorm,
- Glm4MoeRotaryEmbedding,
-)
-from ..glm4_moe_lite.configuration_glm4_moe_lite import Glm4MoeLiteConfig
-from ..glm4_moe_lite.modeling_glm4_moe_lite import (
- Glm4MoeLiteDecoderLayer,
+from ..deepseek_v3.modeling_deepseek_v3 import (
+ DeepseekV3Attention,
+ DeepseekV3RMSNorm,
+ apply_rotary_pos_emb_interleave,
eager_attention_forward,
)
+from ..deepseek_v32.configuration_deepseek_v32 import DeepseekV32Config
+from ..deepseek_v32.modeling_deepseek_v32 import (
+ DeepseekV32DecoderLayer,
+ DeepseekV32ForCausalLM,
+ DeepseekV32Indexer,
+ DeepseekV32Model,
+ DeepseekV32PreTrainedModel,
+ DeepseekV32RotaryEmbedding,
+)
logger = logging.get_logger(__name__)
-def apply_rotary_pos_emb(
- x: torch.Tensor,
- cos: torch.Tensor,
- sin: torch.Tensor,
- unsqueeze_dim: int = 1,
-) -> torch.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=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]`.
- cos (`torch.Tensor`): Cosine part from RotaryEmbedding, shape `[batch, seq_len, head_dim]`.
- sin (`torch.Tensor`): Sine part from RotaryEmbedding, shape `[batch, seq_len, head_dim]`.
- unsqueeze_dim (`int`): Dimension along which to unsqueeze cos/sin for broadcasting.
- Use `1` when x is `[B, H, S, D]` (BHSD) and `2` when x is `[B, S, H, D]` (BSHD).
-
- Returns:
- `torch.Tensor`: Tensor with rotary embeddings applied, same shape as input.
- """
- 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")
@strict
-class GlmMoeDsaConfig(Glm4MoeLiteConfig):
+class GlmMoeDsaConfig(DeepseekV32Config):
r"""
n_group (`int`, *optional*, defaults to 1):
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.
+ MLP type pattern for each layer (`"dense"` or `"sparse"`). Defaults to 3 dense + rest sparse.
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):
+ index_n_heads (`int`, *optional*, defaults to 32):
Number of heads for the indexer projections (DSA).
+ first_k_dense_replace (`int`, *optional*, defaults to 3):
+ Number of leading layers that use a dense MLP; the rest use the MoE block.
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`.
+ Per-layer indexer mode (`"full"` runs the indexer, `"shared"` reuses the previous full
+ layer's top-k). Defaults to the pattern derived from `index_topk_freq` /
+ `index_skip_topk_offset` (or `index_topk_pattern`).
```python
>>> from transformers import GlmMoeDsaConfig, GlmMoeDsaModel
@@ -107,51 +78,38 @@ class GlmMoeDsaConfig(Glm4MoeLiteConfig):
>>> configuration = model.config
```"""
- 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.*.mlp.shared_experts.gate_proj": "colwise",
- "layers.*.mlp.shared_experts.up_proj": "colwise",
- "layers.*.mlp.shared_experts.down_proj": "rowwise",
- "layers.*.mlp.gate_proj": "colwise",
- "layers.*.mlp.up_proj": "colwise",
- "layers.*.mlp.down_proj": "rowwise",
- }
-
attribute_map = {
"num_local_experts": "n_routed_experts",
}
+ vocab_size: int = 154880
hidden_size: int = 6144
intermediate_size: int = 12288
moe_intermediate_size: int = 2048
num_hidden_layers: int = 78
num_attention_heads: int = 64
num_key_value_heads: int = 64
+ n_shared_experts: int = 1
n_routed_experts: int = 256
routed_scaling_factor: float = 2.5
+ kv_lora_rank: int = 512
q_lora_rank: int = 2048
+ qk_rope_head_dim: int = 64
+ v_head_dim: int = 256
+ qk_nope_head_dim: int = 192
+ n_group: int = 1
+ topk_group: int = 1
num_experts_per_tok: int = 8
+ max_position_embeddings: int = 202752
+ rms_norm_eps: float = 1e-5
index_topk: int = 2048
index_head_dim: int = 128
index_n_heads: int = 32
+ # `"full"` runs the indexer, `"shared"` reuses the previous full layer's index mask.
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
- 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 % moe_layer_freq == 0 else "dense" for i in range(self.num_hidden_layers)
- ]
-
+ # Per-layer indexer mode: a pattern (e.g. `"FSSF..."`) overrides the freq/offset schedule.
if self.indexer_types is None:
pattern = kwargs.get("index_topk_pattern")
if pattern is not None:
@@ -164,218 +122,35 @@ def __post_init__(self, **kwargs):
self.indexer_types = [
"full" if (max(i - offset + 1, 0) % freq) == 0 else "shared" for i in range(self.num_hidden_layers)
]
- PreTrainedConfig.__post_init__(self, **kwargs)
+ super().__post_init__(**kwargs)
-class GlmMoeDsaRMSNorm(Glm4MoeRMSNorm):
+class GlmMoeDsaRMSNorm(DeepseekV3RMSNorm):
pass
-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 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
- `num_hidden_layers` attention layers. Keys are concatenated along the sequence dimension
- during autoregressive decode.
- """
-
- def __init__(self, config: "GlmMoeDsaConfig", layer_idx: int):
- super().__init__()
- self.config = config
- self.layer_idx = layer_idx
-
- self.hidden_size: int = config.hidden_size
- self.n_heads: int = config.index_n_heads
- 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
-
- # Named to match checkpoint: wq_b, wk, k_norm
- self.wq_b = nn.Linear(self.q_lora_rank, self.n_heads * self.head_dim, bias=False)
- self.wk = nn.Linear(self.hidden_size, self.head_dim, bias=False)
- self.k_norm = nn.LayerNorm(self.head_dim, eps=1e-6)
- # Named to match checkpoint: weights_proj
- # In the reference, this is fp32; the HF FP8 checkpoint stores a bf16 tensor.
- # Keeping it as a plain Linear prevents FP8 conversion (see `_keep_in_fp32_modules`).
- self.weights_proj = nn.Linear(self.hidden_size, self.n_heads, bias=False)
- self.softmax_scale = self.head_dim**-0.5
-
- # Indexer maintains its own key cache (not in DynamicCache, which is sized for attention layers only)
- self.register_buffer("_cached_keys", None, persistent=False)
-
- @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,
- use_cache: bool = False,
- ) -> torch.LongTensor:
- """
- Computes top-k token indices for sparse attention (DSA).
-
- This is the bf16 equivalent of the reference Indexer which uses `rotate_activation` (Hadamard transform)
- and `fp8_index` (FP8 quantized scoring kernel). Since the Hadamard transform is orthogonal (dot products
- are preserved: Hq·Hk = q·k), and FP8 quantization is a precision optimization, we skip both and compute
- scores directly in bf16/fp32.
-
- The scoring logic computes:
- index_score[b,s,t] = Σ_h (weight[b,s,h] · softmax_scale · q[b,s,h,:] · k[b,t,:])
-
- Args:
- hidden_states: Input hidden states `[B, S, hidden_size]`.
- q_resid: Query residual from `q_a_layernorm(q_a_proj(x))`, shape `[B, S, q_lora_rank]`.
- position_embeddings: `(cos, sin)` from RotaryEmbedding.
- attention_mask: Causal mask, broadcastable to `[B, S, T]`.
- use_cache: Whether to store/update the indexer's own key cache for autoregressive decode.
-
- Returns:
- `torch.LongTensor`: Top-k token indices of shape `[B, S, topk]`.
- """
- batch_size, seq_len, _ = hidden_states.shape
- cos, sin = position_embeddings
-
- # === Queries ===
- 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 = 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 = torch.cat([k_pe, k_nope], dim=-1) # [B, S, D]
-
- # === Key cache (managed by the indexer, not DynamicCache) ===
- # Reset cache on prefill (new prompt) to avoid stale keys / batch-size mismatch
- if seq_len > 1:
- self._cached_keys = None
-
- if use_cache:
- if self._cached_keys is not None:
- k_cached = torch.cat([self._cached_keys, k], dim=1) # [B, T, D]
- else:
- k_cached = k
- self._cached_keys = k_cached
- else:
- k_cached = k
-
- # === Scoring ===
- # Reference: weights = weights_proj(x.float()) * n_heads^(-0.5)
- # Reference: weights = weights.unsqueeze(-1) * q_scale * softmax_scale
- # Reference: index_score = fp8_index(q_fp8, weights, k_cache, k_scale_cache)
- #
- # In bf16 mode (no FP8), q_scale = 1. The fp8_index kernel computes:
- # score[b,s,t] = sum_h(weights[b,s,h] * dot(q[b,s,h,:], k[b,t,:]))
- # where weights already absorbs n_heads^(-0.5) and softmax_scale.
-
- # Don't force fp32 inputs here: the checkpoint stores `weights_proj.weight` in bf16.
- # Use native dtype for matmul, then upcast the result for scoring stability.
- weights = self.weights_proj(hidden_states).float() * (self.n_heads**-0.5) # [B, S, H]
-
- # q·k^T per head: [B, S, H, D] @ [B, T, D]^T → [B, S, H, T]
- scores = torch.einsum("bshd,btd->bsht", q.float(), k_cached.float()) * self.softmax_scale
- scores = F.relu(scores)
- # Weight per head and sum across heads → [B, S, T]
- index_scores = torch.einsum("bsht,bsh->bst", scores, weights)
+class GlmMoeDsaRotaryEmbedding(DeepseekV32RotaryEmbedding):
+ pass
- if attention_mask is not None:
- index_scores = index_scores + attention_mask
- total_len = index_scores.shape[-1]
- topk = min(self.index_topk, total_len)
- topk_indices = index_scores.topk(topk, dim=-1).indices # [B, S, topk]
- return topk_indices
+class GlmMoeDsaIndexer(DeepseekV32Indexer):
+ pass
-class GlmMoeDsaAttention(nn.Module):
+class GlmMoeDsaAttention(DeepseekV3Attention):
"""
- Multi-head Latent Attention (MLA) with DeepSeek Sparse Attention (DSA) indexer.
-
- This follows the same architecture as DeepSeek V3.2's MLA:
- - Query: x → q_a_proj → RMSNorm → q_b_proj → split(q_nope, q_pe) → RoPE(q_pe)
- - KV: x → kv_a_proj → split(kv_compressed, k_pe) → RMSNorm(kv_compressed) → kv_b_proj
- → RoPE(k_pe)
- - Cache: fully expanded key_states [B, H, T, qk_head_dim] and value_states [B, H, T, v_head_dim]
- - Indexer: selects top-k tokens via DSA, applied as an additive -inf mask on attention scores
+ DeepSeek-V3 MLA + a DSA indexer, extended with **cross-layer top-k sharing**.
- **Caching strategy**: follows the DeepSeek V3 transformers convention of fully expanding K/V
- before caching. This ensures compatibility with DynamicCache, StaticCache, flash attention,
- 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.
+ `config.indexer_types[layer_idx]` decides whether this layer runs its own indexer (`"full"`) or
+ reuses the previous full layer's top-k selection (`"shared"`).
+ `next_skip_topk` signals that the *next* layer will reuse this
+ layer's top-k, so it is propagated upward via `prev_topk_indices`.
"""
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.is_causal = True
-
- # Query projection (with optional LoRA)
- 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 (MLA compressed path)
- 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,
- )
-
- self.scaling = self.qk_head_dim ** (-0.5)
-
+ super().__init__(config, layer_idx)
# Refer: https://arxiv.org/abs/2603.12201 for more details.
- # 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(
@@ -384,120 +159,86 @@ def forward(
position_embeddings: tuple[torch.Tensor, torch.Tensor],
attention_mask: torch.Tensor | None,
past_key_values: Cache | None = None,
+ position_ids: torch.Tensor | None = None,
prev_topk_indices: torch.Tensor | None = None,
**kwargs: Unpack[FlashAttentionKwargs],
- ) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor | None]:
+ ) -> tuple[torch.Tensor, torch.Tensor | None, torch.Tensor | None]:
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)
+
+ q_resid = self.q_a_layernorm(self.q_a_proj(hidden_states))
+ q_states = self.q_b_proj(q_resid).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
+ q_rot, k_rot = apply_rotary_pos_emb_interleave(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)
- # ===== Query path =====
- if self.q_lora_rank is None:
- query_states = self.q_proj(hidden_states)
- q_resid = None
- else:
- q_resid = self.q_a_layernorm(self.q_a_proj(hidden_states)) # [B, S, q_lora_rank]
- query_states = self.q_b_proj(q_resid)
- 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)
-
- # ===== KV path =====
- compressed_kv = self.kv_a_proj_with_mqa(hidden_states) # [B, S, kv_rank + rope_D]
- k_compressed, k_pe = torch.split(compressed_kv, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1)
- k_compressed = self.kv_a_layernorm(k_compressed) # [B, S, kv_rank]
-
- # Expand KV through kv_b_proj
- kv_expanded = self.kv_b_proj(k_compressed) # [B, S, H * (nope_D + v_D)]
- kv_expanded = kv_expanded.view(batch_size, seq_length, -1, self.qk_nope_head_dim + self.v_head_dim)
- k_nope, value_states = torch.split(kv_expanded, [self.qk_nope_head_dim, self.v_head_dim], dim=-1)
- 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), using interleaved pair layout.
- k_pe = k_pe.view(batch_size, 1, seq_length, self.qk_rope_head_dim) # [B, 1, S, rope_D]
- 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
- query_states = torch.cat([q_nope, q_pe], dim=-1) # [B, H, S, qk_head_dim]
- key_states = torch.cat([k_nope, k_pe], dim=-1) # [B, H, S, qk_head_dim]
-
- # Cache update
if past_key_values is not None:
key_states, value_states = past_key_values.update(key_states, value_states, self.layer_idx)
- # ===== 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
- else attention_mask.unsqueeze(1)
- if attention_mask is not None
- else None
- )
+ # DSA: select this layer's top-k tokens, or reuse the previous full layer's on `"shared"` layers.
+ if self.indexer is not None:
+ indexer_mask = attention_mask[:, 0, :, :] if attention_mask is not None else None
topk_indices = self.indexer(
hidden_states,
q_resid,
position_embeddings,
indexer_mask,
- use_cache=past_key_values is not None,
+ position_ids,
+ past_key_values=past_key_values,
) # [B, S, topk]
else:
- topk_indices = prev_topk_indices # [B, S, topk]
-
- # Build combined DSA + causal mask: -inf everywhere except selected top-k positions
- total_len = key_states.shape[2]
- index_mask = torch.full(
- (batch_size, seq_length, total_len),
- float("-inf"),
- device=hidden_states.device,
- dtype=query_states.dtype,
- )
- index_mask.scatter_(-1, topk_indices, 0.0) # [B, S, T]
- index_mask = index_mask.unsqueeze(1) # [B, 1, S, T]
- if attention_mask is not None and attention_mask.dim() == 4:
- causal_mask = attention_mask[..., :total_len]
- combined_mask = index_mask + causal_mask
- else:
- combined_mask = (
- attention_mask.masked_fill(index_mask == float("-inf"), float("-inf"))
- if attention_mask is not None
- else index_mask
+ if prev_topk_indices is None:
+ raise ValueError("Shared DSA layers require top-k indices from a previous full indexer layer.")
+ topk_indices = prev_topk_indices
+
+ sparse_indices = None
+ if self.config._attn_implementation in ("eager", "sdpa"):
+ index_mask = (
+ topk_indices.new_ones((batch_size, seq_length, key_states.shape[2]), dtype=torch.bool)
+ .scatter(-1, topk_indices.long(), False)
+ .unsqueeze(1)
)
-
- # Flash attention head_dim padding (qk_head_dim != v_head_dim)
- if is_flash_attention_requested(self.config) and self.qk_head_dim != self.v_head_dim:
- value_states = F.pad(value_states, [0, self.qk_head_dim - self.v_head_dim])
+ if attention_mask is None:
+ key_positions = torch.arange(key_states.shape[2], device=hidden_states.device)
+ index_mask = index_mask | (key_positions[None, None, None, :] > position_ids[:, None, :, None])
+ attention_mask = hidden_states.new_zeros((batch_size, 1, seq_length, key_states.shape[2]))
+ attention_mask = attention_mask.masked_fill(index_mask, torch.finfo(hidden_states.dtype).min)
+ else:
+ sparse_indices = topk_indices
attention_interface: Callable = ALL_ATTENTION_FUNCTIONS.get_interface(
self.config._attn_implementation, eager_attention_forward
)
-
attn_output, attn_weights = attention_interface(
self,
query_states,
key_states,
value_states,
- combined_mask,
+ attention_mask,
dropout=0.0 if not self.training else self.attention_dropout,
scaling=self.scaling,
- indices=topk_indices, # flash_mla_with_kvcache
+ indices=sparse_indices, # consumed by flash_mla_with_kvcache; ignored by eager / SDPA
**kwargs,
)
- if is_flash_attention_requested(self.config) 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, topk_indices if self.next_skip_topk else None
+ return attn_output, attn_weights, topk_indices
-class GlmMoeDsaDecoderLayer(Glm4MoeLiteDecoderLayer):
+class GlmMoeDsaDecoderLayer(DeepseekV32DecoderLayer):
def forward(
self,
hidden_states: torch.Tensor,
@@ -506,7 +247,7 @@ def forward(
past_key_values: Cache | None = None,
use_cache: bool | None = False,
position_embeddings: tuple[torch.Tensor, torch.Tensor] | None = None,
- prev_topk_indices: torch.Tensor | None = None,
+ prev_topk_indices: torch.Tensor | None = None, # MAIN DIFF with DSV3.2
**kwargs: Unpack[TransformersKwargs],
) -> tuple[torch.Tensor, torch.Tensor | None]:
residual = hidden_states
@@ -519,12 +260,11 @@ def forward(
past_key_values=past_key_values,
use_cache=use_cache,
position_embeddings=position_embeddings,
- prev_topk_indices=prev_topk_indices,
+ prev_topk_indices=prev_topk_indices, # MAIN DIFF with DSV3.2
**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)
@@ -532,39 +272,11 @@ def forward(
return hidden_states, topk_indices
-class GlmMoeDsaPreTrainedModel(Glm4MoePreTrainedModel):
- # NOTE: FP8 quantization uses `_keep_in_fp32_modules` (not `_strict`) to decide which modules to NOT convert.
- # We must keep `indexer.weights_proj` as a plain Linear to match the checkpoint (no `weight_scale_inv`).
- _keep_in_fp32_modules = ["indexer.weights_proj"]
- _keep_in_fp32_modules_strict = ["e_score_correction_bias"]
+class GlmMoeDsaPreTrainedModel(DeepseekV32PreTrainedModel):
_keys_to_ignore_on_load_unexpected = [r"model\.layers\.78.*"]
- _supports_flash_attn = False # flash-mla kernels need a bit more work in the way we enable them!
- _supports_sdpa = True
- _supports_flex_attn = False
- _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):
+class GlmMoeDsaModel(DeepseekV32Model):
def forward(
self,
input_ids: torch.LongTensor | None = None,
@@ -600,7 +312,7 @@ def forward(
hidden_states = inputs_embeds
position_embeddings = self.rotary_emb(hidden_states, position_ids=position_ids)
- topk_indices = None
+ topk_indices = None # MAIN DIFF with DSV3.2
for decoder_layer in self.layers[: self.config.num_hidden_layers]:
hidden_states, topk_indices = decoder_layer(
hidden_states,
@@ -609,7 +321,7 @@ def forward(
position_ids=position_ids,
past_key_values=past_key_values,
use_cache=use_cache,
- prev_topk_indices=topk_indices,
+ prev_topk_indices=topk_indices, # MAIN DIFF with DSV3.2
**kwargs,
)
@@ -620,7 +332,7 @@ def forward(
)
-class GlmMoeDsaForCausalLM(Glm4MoeForCausalLM):
+class GlmMoeDsaForCausalLM(DeepseekV32ForCausalLM):
pass
diff --git a/src/transformers/models/longcat_flash/modeling_longcat_flash.py b/src/transformers/models/longcat_flash/modeling_longcat_flash.py
index d5ac6e237742..4a287fde6e02 100644
--- a/src/transformers/models/longcat_flash/modeling_longcat_flash.py
+++ b/src/transformers/models/longcat_flash/modeling_longcat_flash.py
@@ -245,13 +245,6 @@ def forward(self, hidden_states):
return hidden_states
-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)
-
-
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,
@@ -291,9 +284,12 @@ def eager_attention_forward(
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.
+ Applies interleaved Rotary Position Embedding to the query and key tensors.
+
+ DeepSeek lays the rotary dimensions out in interleaved pairs `(x0, x1), (x2, x3), ...`, each rotated by a
+ single frequency. We compute that rotation directly on the even/odd slices instead of de-interleaving with a
+ `view`/`transpose`/`reshape`; the output is bit-identical to the de-interleaved `rotate_half` formulation while
+ avoiding the extra contiguous copy.
Args:
q (`torch.Tensor`): The query tensor.
@@ -313,17 +309,15 @@ def apply_rotary_pos_emb_interleave(q, k, cos, sin, position_ids=None, unsqueeze
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)
+ # `cos`/`sin` are `cat(freqs, freqs)`; the first half holds the per-pair angle.
+ cos = cos[..., : cos.shape[-1] // 2].unsqueeze(unsqueeze_dim)
+ sin = sin[..., : sin.shape[-1] // 2].unsqueeze(unsqueeze_dim)
- b, h, s, d = k.shape
- k = k.view(b, h, s, d // 2, 2).transpose(4, 3).reshape(b, h, s, d)
+ q1, q2 = q[..., 0::2], q[..., 1::2]
+ k1, k2 = k[..., 0::2], k[..., 1::2]
- q_embed = (q * cos) + (rotate_half(q) * sin)
- k_embed = (k * cos) + (rotate_half(k) * sin)
+ q_embed = torch.cat([q1 * cos - q2 * sin, q2 * cos + q1 * sin], dim=-1)
+ k_embed = torch.cat([k1 * cos - k2 * sin, k2 * cos + k1 * sin], dim=-1)
return q_embed, k_embed
diff --git a/src/transformers/models/mistral4/modeling_mistral4.py b/src/transformers/models/mistral4/modeling_mistral4.py
index 006ddad187bf..23bfe1091b69 100644
--- a/src/transformers/models/mistral4/modeling_mistral4.py
+++ b/src/transformers/models/mistral4/modeling_mistral4.py
@@ -327,9 +327,12 @@ def eager_attention_forward(
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.
+ Applies interleaved Rotary Position Embedding to the query and key tensors.
+
+ DeepSeek lays the rotary dimensions out in interleaved pairs `(x0, x1), (x2, x3), ...`, each rotated by a
+ single frequency. We compute that rotation directly on the even/odd slices instead of de-interleaving with a
+ `view`/`transpose`/`reshape`; the output is bit-identical to the de-interleaved `rotate_half` formulation while
+ avoiding the extra contiguous copy.
Args:
q (`torch.Tensor`): The query tensor.
@@ -349,17 +352,15 @@ def apply_rotary_pos_emb_interleave(q, k, cos, sin, position_ids=None, unsqueeze
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)
+ # `cos`/`sin` are `cat(freqs, freqs)`; the first half holds the per-pair angle.
+ cos = cos[..., : cos.shape[-1] // 2].unsqueeze(unsqueeze_dim)
+ sin = sin[..., : sin.shape[-1] // 2].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)
+ q1, q2 = q[..., 0::2], q[..., 1::2]
+ k1, k2 = k[..., 0::2], k[..., 1::2]
- 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)
+ q_embed = torch.cat([q1 * cos - q2 * sin, q2 * cos + q1 * sin], dim=-1)
+ k_embed = torch.cat([k1 * cos - k2 * sin, k2 * cos + k1 * sin], dim=-1)
return q_embed, k_embed
diff --git a/src/transformers/models/youtu/modeling_youtu.py b/src/transformers/models/youtu/modeling_youtu.py
index d40bef358da6..f293235f5cbb 100644
--- a/src/transformers/models/youtu/modeling_youtu.py
+++ b/src/transformers/models/youtu/modeling_youtu.py
@@ -224,9 +224,12 @@ def eager_attention_forward(
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.
+ Applies interleaved Rotary Position Embedding to the query and key tensors.
+
+ DeepSeek lays the rotary dimensions out in interleaved pairs `(x0, x1), (x2, x3), ...`, each rotated by a
+ single frequency. We compute that rotation directly on the even/odd slices instead of de-interleaving with a
+ `view`/`transpose`/`reshape`; the output is bit-identical to the de-interleaved `rotate_half` formulation while
+ avoiding the extra contiguous copy.
Args:
q (`torch.Tensor`): The query tensor.
@@ -246,17 +249,15 @@ def apply_rotary_pos_emb_interleave(q, k, cos, sin, position_ids=None, unsqueeze
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)
+ # `cos`/`sin` are `cat(freqs, freqs)`; the first half holds the per-pair angle.
+ cos = cos[..., : cos.shape[-1] // 2].unsqueeze(unsqueeze_dim)
+ sin = sin[..., : sin.shape[-1] // 2].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)
+ q1, q2 = q[..., 0::2], q[..., 1::2]
+ k1, k2 = k[..., 0::2], k[..., 1::2]
- 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)
+ q_embed = torch.cat([q1 * cos - q2 * sin, q2 * cos + q1 * sin], dim=-1)
+ k_embed = torch.cat([k1 * cos - k2 * sin, k2 * cos + k1 * sin], dim=-1)
return q_embed, k_embed
diff --git a/tests/models/deepseek_v32/__init__.py b/tests/models/deepseek_v32/__init__.py
new file mode 100644
index 000000000000..e69de29bb2d1
diff --git a/tests/models/deepseek_v32/test_modeling_deepseek_v32.py b/tests/models/deepseek_v32/test_modeling_deepseek_v32.py
new file mode 100644
index 000000000000..c8ceb53d45a9
--- /dev/null
+++ b/tests/models/deepseek_v32/test_modeling_deepseek_v32.py
@@ -0,0 +1,351 @@
+# 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.
+# 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 DeepSeekV3.2 model."""
+
+import unittest
+
+import pytest
+from parameterized import parameterized
+
+from transformers import Cache, is_torch_available
+from transformers.testing_utils import require_torch, require_torch_accelerator, slow
+
+from ...causal_lm_tester import CausalLMModelTest, CausalLMModelTester
+from ...test_modeling_common import (
+ TEST_EAGER_MATCHES_BATCHED_AND_GROUPED_INFERENCE_PARAMETERIZATION,
+ TEST_EAGER_MATCHES_SDPA_INFERENCE_PARAMETERIZATION,
+)
+
+
+if is_torch_available():
+ import torch
+
+ from transformers import (
+ AutoTokenizer,
+ DeepseekV32ForCausalLM,
+ DeepseekV32Model,
+ )
+
+
+# Opening of "Alice's Adventures in Wonderland" by Lewis Carroll (public domain, Project Gutenberg).
+# Used as a real, coherent long-context prompt (~2.8K tokens > index_topk) for the DSA indexer test.
+LONG_CONTEXT_PROMPT = """Alice was beginning to get very tired of sitting by her sister on the bank, and of having nothing to do: once or twice she had peeped into the book her sister was reading, but it had no pictures or conversations in it, “and what is the use of a book,” thought Alice “without pictures or conversations?”
+
+So she was considering in her own mind (as well as she could, for the hot day made her feel very sleepy and stupid), whether the pleasure of making a daisy-chain would be worth the trouble of getting up and picking the daisies, when suddenly a White Rabbit with pink eyes ran close by her.
+
+There was nothing so _very_ remarkable in that; nor did Alice think it so _very_ much out of the way to hear the Rabbit say to itself, “Oh dear! Oh dear! I shall be late!” (when she thought it over afterwards, it occurred to her that she ought to have wondered at this, but at the time it all seemed quite natural); but when the Rabbit actually _took a watch out of its waistcoat-pocket_, and looked at it, and then hurried on, Alice started to her feet, for it flashed across her mind that she had never before seen a rabbit with either a waistcoat-pocket, or a watch to take out of it, and burning with curiosity, she ran across the field after it, and fortunately was just in time to see it pop down a large rabbit-hole under the hedge.
+
+In another moment down went Alice after it, never once considering how in the world she was to get out again.
+
+The rabbit-hole went straight on like a tunnel for some way, and then dipped suddenly down, so suddenly that Alice had not a moment to think about stopping herself before she found herself falling down a very deep well.
+
+Either the well was very deep, or she fell very slowly, for she had plenty of time as she went down to look about her and to wonder what was going to happen next. First, she tried to look down and make out what she was coming to, but it was too dark to see anything; then she looked at the sides of the well, and noticed that they were filled with cupboards and book-shelves; here and there she saw maps and pictures hung upon pegs. She took down a jar from one of the shelves as she passed; it was labelled “ORANGE MARMALADE”, but to her great disappointment it was empty: she did not like to drop the jar for fear of killing somebody underneath, so managed to put it into one of the cupboards as she fell past it.
+
+“Well!” thought Alice to herself, “after such a fall as this, I shall think nothing of tumbling down stairs! How brave they’ll all think me at home! Why, I wouldn’t say anything about it, even if I fell off the top of the house!” (Which was very likely true.)
+
+Down, down, down. Would the fall _never_ come to an end? “I wonder how many miles I’ve fallen by this time?” she said aloud. “I must be getting somewhere near the centre of the earth. Let me see: that would be four thousand miles down, I think—” (for, you see, Alice had learnt several things of this sort in her lessons in the schoolroom, and though this was not a _very_ good opportunity for showing off her knowledge, as there was no one to listen to her, still it was good practice to say it over) “—yes, that’s about the right distance—but then I wonder what Latitude or Longitude I’ve got to?” (Alice had no idea what Latitude was, or Longitude either, but thought they were nice grand words to say.)
+
+Presently she began again. “I wonder if I shall fall right _through_ the earth! How funny it’ll seem to come out among the people that walk with their heads downward! The Antipathies, I think—” (she was rather glad there _was_ no one listening, this time, as it didn’t sound at all the right word) “—but I shall have to ask them what the name of the country is, you know. Please, Ma’am, is this New Zealand or Australia?” (and she tried to curtsey as she spoke—fancy _curtseying_ as you’re falling through the air! Do you think you could manage it?) “And what an ignorant little girl she’ll think me for asking! No, it’ll never do to ask: perhaps I shall see it written up somewhere.”
+
+Down, down, down. There was nothing else to do, so Alice soon began talking again. “Dinah’ll miss me very much to-night, I should think!” (Dinah was the cat.) “I hope they’ll remember her saucer of milk at tea-time. Dinah my dear! I wish you were down here with me! There are no mice in the air, I’m afraid, but you might catch a bat, and that’s very like a mouse, you know. But do cats eat bats, I wonder?” And here Alice began to get rather sleepy, and went on saying to herself, in a dreamy sort of way, “Do cats eat bats? Do cats eat bats?” and sometimes, “Do bats eat cats?” for, you see, as she couldn’t answer either question, it didn’t much matter which way she put it. She felt that she was dozing off, and had just begun to dream that she was walking hand in hand with Dinah, and saying to her very earnestly, “Now, Dinah, tell me the truth: did you ever eat a bat?” when suddenly, thump! thump! down she came upon a heap of sticks and dry leaves, and the fall was over.
+
+Alice was not a bit hurt, and she jumped up on to her feet in a moment: she looked up, but it was all dark overhead; before her was another long passage, and the White Rabbit was still in sight, hurrying down it. There was not a moment to be lost: away went Alice like the wind, and was just in time to hear it say, as it turned a corner, “Oh my ears and whiskers, how late it’s getting!” She was close behind it when she turned the corner, but the Rabbit was no longer to be seen: she found herself in a long, low hall, which was lit up by a row of lamps hanging from the roof.
+
+There were doors all round the hall, but they were all locked; and when Alice had been all the way down one side and up the other, trying every door, she walked sadly down the middle, wondering how she was ever to get out again.
+
+Suddenly she came upon a little three-legged table, all made of solid glass; there was nothing on it except a tiny golden key, and Alice’s first thought was that it might belong to one of the doors of the hall; but, alas! either the locks were too large, or the key was too small, but at any rate it would not open any of them. However, on the second time round, she came upon a low curtain she had not noticed before, and behind it was a little door about fifteen inches high: she tried the little golden key in the lock, and to her great delight it fitted!
+
+Alice opened the door and found that it led into a small passage, not much larger than a rat-hole: she knelt down and looked along the passage into the loveliest garden you ever saw. How she longed to get out of that dark hall, and wander about among those beds of bright flowers and those cool fountains, but she could not even get her head through the doorway; “and even if my head would go through,” thought poor Alice, “it would be of very little use without my shoulders. Oh, how I wish I could shut up like a telescope! I think I could, if I only knew how to begin.” For, you see, so many out-of-the-way things had happened lately, that Alice had begun to think that very few things indeed were really impossible.
+
+There seemed to be no use in waiting by the little door, so she went back to the table, half hoping she might find another key on it, or at any rate a book of rules for shutting people up like telescopes: this time she found a little bottle on it, (“which certainly was not here before,” said Alice,) and round the neck of the bottle was a paper label, with the words “DRINK ME,” beautifully printed on it in large letters.
+
+It was all very well to say “Drink me,” but the wise little Alice was not going to do _that_ in a hurry. “No, I’ll look first,” she said, “and see whether it’s marked ‘_poison_’ or not”; for she had read several nice little histories about children who had got burnt, and eaten up by wild beasts and other unpleasant things, all because they _would_ not remember the simple rules their friends had taught them: such as, that a red-hot poker will burn you if you hold it too long; and that if you cut your finger _very_ deeply with a knife, it usually bleeds; and she had never forgotten that, if you drink much from a bottle marked “poison,” it is almost certain to disagree with you, sooner or later.
+
+However, this bottle was _not_ marked “poison,” so Alice ventured to taste it, and finding it very nice, (it had, in fact, a sort of mixed flavour of cherry-tart, custard, pine-apple, roast turkey, toffee, and hot buttered toast,) she very soon finished it off.
+
+* * * * * * *
+
+* * * * * *
+
+* * * * * * *
+
+“What a curious feeling!” said Alice; “I must be shutting up like a telescope.”
+
+And so it was indeed: she was now only ten inches high, and her face brightened up at the thought that she was now the right size for going through the little door into that lovely garden. First, however, she waited for a few minutes to see if she was going to shrink any further: she felt a little nervous about this; “for it might end, you know,” said Alice to herself, “in my going out altogether, like a candle. I wonder what I should be like then?” And she tried to fancy what the flame of a candle is like after the candle is blown out, for she could not remember ever having seen such a thing.
+
+After a while, finding that nothing more happened, she decided on going into the garden at once; but, alas for poor Alice! when she got to the door, she found she had forgotten the little golden key, and when she went back to the table for it, she found she could not possibly reach it: she could see it quite plainly through the glass, and she tried her best to climb up one of the legs of the table, but it was too slippery; and when she had tired herself out with trying, the poor little thing sat down and cried.
+
+“Come, there’s no use in crying like that!” said Alice to herself, rather sharply; “I advise you to leave off this minute!” She generally gave herself very good advice, (though she very seldom followed it), and sometimes she scolded herself so severely as to bring tears into her eyes; and once she remembered trying to box her own ears for having cheated herself in a game of croquet she was playing against herself, for this curious child was very fond of pretending to be two people. “But it’s no use now,” thought poor Alice, “to pretend to be two people! Why, there’s hardly enough of me left to make _one_ respectable person!”
+
+Soon her eye fell on a little glass box that was lying under the table: she opened it, and found in it a very small cake, on which the words “EAT ME” were beautifully marked in currants. “Well, I’ll eat it,” said Alice, “and if it makes me grow larger, I can reach the key; and if it makes me grow smaller, I can creep under the door; so either way I’ll get into the garden, and I don’t care which happens!”
+
+She ate a little bit, and said anxiously to herself, “Which way? Which way?”, holding her hand on the top of her head to feel which way it was growing, and she was quite surprised to find that she remained the same size: to be sure, this generally happens when one eats cake, but Alice had got so much into the way of expecting nothing but out-of-the-way things to happen, that it seemed quite dull and stupid for life to go on in the common way.
+
+So she set to work, and very soon finished off the cake.
+
+* * * * * * *
+
+* * * * * *
+
+* * * * * * *
+
+
+
+CHAPTER II. The Pool of Tears
+
+“Curiouser and curiouser!” cried Alice (she was so much surprised, that for the moment she quite forgot how to speak good English); “now I’m opening out like the largest telescope that ever was! Good-bye, feet!” (for when she looked down at her feet, they seemed to be almost out of sight, they were getting so far off)."""
+
+
+class DeepseekV32ModelTester(CausalLMModelTester):
+ if is_torch_available():
+ base_model_class = DeepseekV32Model
+
+ def __init__(
+ self,
+ parent,
+ n_routed_experts=8,
+ kv_lora_rank=32,
+ q_lora_rank=16,
+ qk_nope_head_dim=64,
+ qk_rope_head_dim=64,
+ first_k_dense_replace=1,
+ n_group=1,
+ topk_group=1,
+ ):
+ super().__init__(parent=parent)
+ 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.first_k_dense_replace = first_k_dense_replace
+ self.n_group = n_group
+ self.topk_group = topk_group
+
+
+@require_torch
+class DeepseekV32ModelTest(CausalLMModelTest, unittest.TestCase):
+ pipeline_model_mapping = (
+ {
+ "feature-extraction": DeepseekV32Model,
+ "text-generation": DeepseekV32ForCausalLM,
+ }
+ if is_torch_available()
+ else {}
+ )
+ fx_compatible = False
+ test_torchscript = False
+ test_all_params_have_gradient = False
+ model_tester_class = DeepseekV32ModelTester
+ model_split_percents = [0.5, 0.7, 0.8]
+
+ # used in `test_torch_compile_for_training`
+ _torch_compile_train_cls = DeepseekV32ForCausalLM if is_torch_available() else None
+
+ @unittest.skip("DeepseekV32 applies RoPE to 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("DeepseekV32 applies RoPE to qk_rope_head_dim; generic rope scaling tests assume config.head_dim")
+ def test_model_rope_scaling_from_config(self, scaling_type):
+ pass
+
+ def _check_past_key_values_for_generate(self, batch_size, past_key_values, seq_length, config):
+ """Needs to be overridden as deepseek 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)
+
+ @parameterized.expand([("random",), ("same",)])
+ @unittest.skip("DeepseekV32 uses MLA so it is not compatible with assisted decoding")
+ def test_assisted_decoding_matches_greedy_search(self, assistant_type):
+ pass
+
+ @unittest.skip("DeepseekV32 uses MLA so it is not compatible with assisted decoding")
+ def test_prompt_lookup_decoding_matches_greedy_search(self):
+ pass
+
+ @unittest.skip("DeepseekV32 uses MLA so it is not compatible with assisted decoding")
+ def test_assisted_decoding_sample(self):
+ pass
+
+ @unittest.skip("DeepseekV32 uses MLA so it is not compatible with the standard cache format")
+ def test_beam_search_generate_dict_outputs_use_cache(self):
+ pass
+
+ @unittest.skip("DeepseekV32 uses MLA so it is not compatible with the standard cache format")
+ def test_greedy_generate_dict_outputs_use_cache(self):
+ pass
+
+ @unittest.skip(reason="SDPA can't dispatch on flash due to unsupported head dims")
+ def test_sdpa_can_dispatch_on_flash(self):
+ pass
+
+ @unittest.skip("Dynamic control flow in MoE")
+ @pytest.mark.torch_compile_test
+ def test_torch_compile_for_training(self):
+ pass
+
+ # DeepSeek Sparse Attention selects tokens with a hard top-k, which is discontinuous: a tiny numerical
+ # difference in the indexer scores (attention backend, padding, batching, sequence packing) can flip
+ # which tokens are selected and thus change the output. These exact cross-backend / padding-equivalence
+ # tests therefore do not hold for DSA (dense models like DeepSeek-V3 pass them).
+ @parameterized.expand(TEST_EAGER_MATCHES_SDPA_INFERENCE_PARAMETERIZATION)
+ @unittest.skip("DSA hard top-k selection is sensitive to tiny numerical differences across backends.")
+ def test_eager_matches_sdpa_inference(self, *args, **kwargs):
+ pass
+
+ @parameterized.expand(TEST_EAGER_MATCHES_BATCHED_AND_GROUPED_INFERENCE_PARAMETERIZATION)
+ @unittest.skip("DSA hard top-k selection is sensitive to tiny numerical differences across batching.")
+ def test_eager_matches_batched_and_grouped_inference(self, *args, **kwargs):
+ pass
+
+ @unittest.skip("DSA hard top-k selection is sensitive to padding shifts (selection can flip).")
+ def test_left_padding_compatibility(self):
+ pass
+
+ @unittest.skip("DSA hard top-k selection is sensitive to sequence packing (selection can flip).")
+ def test_eager_padding_matches_padding_free_with_position_ids(self):
+ pass
+
+ @unittest.skip("DSA hard top-k selection is sensitive to sequence packing (selection can flip).")
+ def test_sdpa_padding_matches_padding_free_with_position_ids(self):
+ pass
+
+ @unittest.skip("MoE routing on a tiny randomly-initialized model makes the overfit target unstable.")
+ def test_training_overfit(self):
+ pass
+
+
+@slow
+@require_torch_accelerator
+class DeepseekV32IntegrationTest(unittest.TestCase):
+ def test_deepseek_v32(self):
+ EXPECTED_TEXT = ['An attention function can be described as mapping a query and a set of key-value pairs to an output, where the query, keys, values, and output are all vectors. The output is computed as a weighted sum of the values, where the weight assigned to each value is computed by a compatibility function of the query with the corresponding key.\n\nWe call our particular attention "Scaled Dot-Product Attention" (Figure (left'] # fmt: skip
+
+ tokenizer = AutoTokenizer.from_pretrained("deepseek-ai/DeepSeek-V3.2-Exp")
+ model = DeepseekV32ForCausalLM.from_pretrained(
+ "deepseek-ai/DeepSeek-V3.2-Exp",
+ device_map="auto",
+ dtype=torch.bfloat16,
+ )
+
+ input_text = [
+ "An attention function can be described as mapping a query and a set of key-value pairs to an output, where the query, keys, values, and output are all vectors." # fmt: skip
+ ]
+ model_inputs = tokenizer(input_text, return_tensors="pt").to(model.device)
+
+ generated_ids = model.generate(**model_inputs, max_new_tokens=50, do_sample=False)
+ generated_text = tokenizer.decode(generated_ids, skip_special_tokens=True)
+ self.assertEqual(generated_text, EXPECTED_TEXT)
+
+ def test_logits_eager(self):
+ input_ids = [1, 306, 4658, 278, 6593, 310, 2834, 338]
+
+ model = DeepseekV32ForCausalLM.from_pretrained(
+ "deepseek-ai/DeepSeek-V3.2-Exp",
+ device_map="auto",
+ dtype=torch.bfloat16,
+ attn_implementation="eager",
+ )
+
+ with torch.no_grad():
+ out = model(torch.tensor([input_ids]).to(model.device))
+
+ EXPECTED_MEAN = torch.tensor([[5.0182, 6.1787, 5.3601, 5.7569, 5.7146, 5.1751, 4.2580, 3.3002]], device=out.logits.device) # fmt: skip
+ torch.testing.assert_close(out.logits.float().mean(-1), EXPECTED_MEAN, atol=1e-3, rtol=1e-3)
+
+ EXPECTED_SLICE = torch.tensor([17.3750, 12.4375, 2.1406, 15.0625, 13.5625, 14.8750, 13.7500, 13.6250, 13.5625, 14.0000, 13.0000, 15.1875, 13.6250, 13.3750, 15.3750], device=out.logits.device) # fmt: skip
+ torch.testing.assert_close(out.logits[0, 0, :15].float(), EXPECTED_SLICE, atol=1e-3, rtol=1e-3)
+
+ def test_logits_long_context(self):
+ prompt = LONG_CONTEXT_PROMPT
+
+ tokenizer = AutoTokenizer.from_pretrained("deepseek-ai/DeepSeek-V3.2-Exp")
+ model = DeepseekV32ForCausalLM.from_pretrained(
+ "deepseek-ai/DeepSeek-V3.2-Exp",
+ device_map="auto",
+ dtype=torch.bfloat16,
+ attn_implementation="eager",
+ )
+
+ inputs = tokenizer(prompt, return_tensors="pt").to(model.device)
+ self.assertGreater(inputs["input_ids"].shape[1], model.config.index_topk)
+
+ with torch.no_grad():
+ out = model(**inputs)
+
+ EXPECTED_MEAN = torch.tensor([1.6388, 0.5092, 1.8607, 1.3185, 0.3267, 1.0360, 2.3001, 3.5533], device=out.logits.device) # fmt: skip
+ torch.testing.assert_close(out.logits[0, -8:].float().mean(-1), EXPECTED_MEAN, atol=1e-3, rtol=1e-3)
+
+ EXPECTED_SLICE = torch.tensor([-2.8594, 23.7500, -0.8086, 18.1250, 19.8750, 15.8125, 12.8750, 14.6875, 21.3750, 19.0000, 20.3750, 20.0000, 18.7500, 16.3750, 22.0000], device=out.logits.device) # fmt: skip
+ torch.testing.assert_close(out.logits[0, -1, :15].float(), EXPECTED_SLICE, atol=1e-3, rtol=1e-3)
+
+ gen = model.generate(**inputs, max_new_tokens=40, do_sample=False)
+ continuation = tokenizer.decode(gen[0, inputs["input_ids"].shape[1] :], skip_special_tokens=True)
+ EXPECTED_GENERATION = " “Oh, my poor little feet, I wonder who will put on your shoes and stockings for you now, dears? I’m sure _I_ shan’t be able!" # fmt: skip
+ self.assertEqual(EXPECTED_GENERATION, continuation)
+
+ def test_batched_generation_padding(self):
+ # Batch two prompts of different lengths so the batch must be padded. Left padding is the correct
+ # mode for decoder-only generation: every row stays coherent. Right padding is known to corrupt the
+ # rows that actually receive padding (transformers emits a warning for it), but the longest row gets
+ # no padding, so its generation must be identical regardless of padding side — which verifies the
+ # attention mask is applied correctly end to end.
+ prompts = [
+ "The capital of France is",
+ "An attention function can be described as mapping a query and a set of key-value pairs to an output, where the query, keys, values, and output are all vectors.", # fmt: skip
+ ]
+
+ EXPECTED_LEFT = [
+ "The capital of France is Paris. 法国的首都是巴黎。\nThe capital of the United States is Washington, D.C. 美国的首都是华盛顿。\nThe capital of the United Kingdom is London. 英国的首都是伦敦", # fmt: skip
+ 'An attention function can be described as mapping a query and a set of key-value pairs to an output, where the query, keys, values, and output are all vectors. The output is computed as a weighted sum of the values, where the weight assigned to each value is computed by a compatibility function of the query with the corresponding key.\n\nWe call our particular attention "Scal', # fmt: skip
+ ]
+
+ model = DeepseekV32ForCausalLM.from_pretrained(
+ "deepseek-ai/DeepSeek-V3.2-Exp",
+ device_map="auto",
+ dtype=torch.bfloat16,
+ attn_implementation="eager",
+ )
+
+ # Left padding: the whole batch is correct.
+ tok_left = AutoTokenizer.from_pretrained("deepseek-ai/DeepSeek-V3.2-Exp", padding_side="left")
+ if tok_left.pad_token is None:
+ tok_left.pad_token = tok_left.eos_token
+ inputs_left = tok_left(prompts, return_tensors="pt", padding=True).to(model.device)
+ gen_left = model.generate(**inputs_left, max_new_tokens=40, do_sample=False)
+ text_left = tok_left.batch_decode(gen_left, skip_special_tokens=True)
+ self.assertEqual(EXPECTED_LEFT, text_left)
+
+ # Right padding: the longest prompt is not padded, so its continuation must match left padding.
+ tok_right = AutoTokenizer.from_pretrained("deepseek-ai/DeepSeek-V3.2-Exp", padding_side="right")
+ if tok_right.pad_token is None:
+ tok_right.pad_token = tok_right.eos_token
+ inputs_right = tok_right(prompts, return_tensors="pt", padding=True).to(model.device)
+ gen_right = model.generate(**inputs_right, max_new_tokens=40, do_sample=False)
+ text_right = tok_right.batch_decode(gen_right, skip_special_tokens=True)
+ self.assertEqual(EXPECTED_LEFT[1], text_right[1])
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 e52c9aec4f13..a69f2e3bd029 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
@@ -35,6 +35,7 @@
from ...causal_lm_tester import CausalLMModelTest, CausalLMModelTester
from ...test_modeling_common import (
+ TEST_EAGER_MATCHES_BATCHED_AND_GROUPED_INFERENCE_PARAMETERIZATION,
TEST_EAGER_MATCHES_SDPA_INFERENCE_PARAMETERIZATION,
)
@@ -110,21 +111,29 @@ def test_indexer_types_respect_skip_topk_offset(self):
["full", "full", "full", "shared", "shared", "shared", "full", "shared"],
)
+ # DSA selects tokens with a hard top-k, which is discontinuous: a tiny numerical difference in the
+ # indexer scores (attention backend, padding, batching, sequence packing) can flip which tokens are
+ # selected and thus change the output, so these exact-equivalence tests do not hold for DSA.
@parameterized.expand(TEST_EAGER_MATCHES_SDPA_INFERENCE_PARAMETERIZATION)
- @unittest.skip("Won't fix: Blip2 + T5 backbone needs custom input preparation for this test")
+ @unittest.skip("DSA hard top-k selection is sensitive to tiny numerical differences across backends.")
def test_eager_matches_sdpa_inference(self, *args):
pass
- @unittest.skip("Not sure MoE can pass this + indexer outputs are not deterministic wrt padding")
- def test_left_padding_compatibility(
- self,
- ):
+ @parameterized.expand(TEST_EAGER_MATCHES_BATCHED_AND_GROUPED_INFERENCE_PARAMETERIZATION)
+ @unittest.skip("DSA hard top-k selection is sensitive to tiny numerical differences across batching.")
+ def test_eager_matches_batched_and_grouped_inference(self, *args):
pass
- @unittest.skip("Not sure MoE can pass this + indexer outputs are not deterministic wrt padding")
- def test_sdpa_padding_matches_padding_free_with_position_ids(
- self,
- ):
+ @unittest.skip("DSA hard top-k selection is sensitive to padding shifts (selection can flip).")
+ def test_left_padding_compatibility(self):
+ pass
+
+ @unittest.skip("DSA hard top-k selection is sensitive to sequence packing (selection can flip).")
+ def test_eager_padding_matches_padding_free_with_position_ids(self):
+ pass
+
+ @unittest.skip("DSA hard top-k selection is sensitive to sequence packing (selection can flip).")
+ def test_sdpa_padding_matches_padding_free_with_position_ids(self):
pass
@unittest.skip("Not sure MoE can pass this + indexer outputs are not deterministic wrt padding")
@@ -208,8 +217,8 @@ def test_glm_moe_dsa_fp8_inference(self):
max_new_tokens=16,
)
- output = tokenizer.decode(outputs, skip_special_tokens=False)
- self.assertqual(
+ output = tokenizer.batch_decode(outputs, skip_special_tokens=False)
+ self.assertEqual(
output,
[
"<|endoftext|><|endoftext|><|endoftext|>Hi, introduce yourself!\nI'm a 18 years old boy from Italy and I'm a student",
diff --git a/utils/check_config_attributes.py b/utils/check_config_attributes.py
index ca3a013d9d05..25a4f7e00edc 100644
--- a/utils/check_config_attributes.py
+++ b/utils/check_config_attributes.py
@@ -132,6 +132,8 @@
"num_nextn_predict_layers",
"router_jitter_noise",
],
+ "DeepseekV32Config": ["head_dim", "layer_types", "mlp_bias", "first_k_dense_replace"],
+ "GlmMoeDsaConfig": ["head_dim", "layer_types", "mlp_bias", "first_k_dense_replace"],
"EsmFoldConfig": ["esm_ablate_pairwise", "esm_ablate_sequence", "esm_input_dropout", "esm_type"],
"TrunkConfig": ["cpu_grad_checkpoint", "layer_drop"],
"SeamlessM4TConfig": True,