From 45ed767a7af61a64dc41e29be1b6c0cc14c21d9c Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 12:37:21 +0200 Subject: [PATCH 01/12] [PyTorch] Add DeepSeekV3Layer skeleton (MLA + MoE) Co-Authored-By: Claude Fable 5 Signed-off-by: Pawel Gadzinski --- .../pytorch/deepseek/__init__.py | 11 ++++++++ transformer_engine/pytorch/deepseek/moe.py | 27 +++++++++++++++++++ .../deepseek/multi_latent_attention.py | 24 +++++++++++++++++ .../pytorch/deepseek/transformer_layer.py | 25 +++++++++++++++++ 4 files changed, 87 insertions(+) create mode 100644 transformer_engine/pytorch/deepseek/__init__.py create mode 100644 transformer_engine/pytorch/deepseek/moe.py create mode 100644 transformer_engine/pytorch/deepseek/multi_latent_attention.py create mode 100644 transformer_engine/pytorch/deepseek/transformer_layer.py diff --git a/transformer_engine/pytorch/deepseek/__init__.py b/transformer_engine/pytorch/deepseek/__init__.py new file mode 100644 index 0000000000..5dafdf1fff --- /dev/null +++ b/transformer_engine/pytorch/deepseek/__init__.py @@ -0,0 +1,11 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""DeepSeekV3 transformer layer built from Transformer Engine MoE building blocks.""" + +from transformer_engine.pytorch.deepseek.multi_latent_attention import MultiLatentAttention +from transformer_engine.pytorch.deepseek.moe import DeepSeekV3MoE +from transformer_engine.pytorch.deepseek.transformer_layer import DeepSeekV3Layer + +__all__ = ["DeepSeekV3Layer", "DeepSeekV3MoE", "MultiLatentAttention"] diff --git a/transformer_engine/pytorch/deepseek/moe.py b/transformer_engine/pytorch/deepseek/moe.py new file mode 100644 index 0000000000..f4f787743b --- /dev/null +++ b/transformer_engine/pytorch/deepseek/moe.py @@ -0,0 +1,27 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""DeepSeekV3 MoE block: sigmoid router with aux-loss-free bias, shared + +routed experts.""" + +import torch + +__all__ = ["DeepSeekV3MoE"] + + +class DeepSeekV3MoE(torch.nn.Module): + """ + DeepSeekV3-style Mixture of Experts block composed from TE MoE + primitives: ``fused_topk_with_score_function`` (sigmoid score function, + expert bias, grouped top-k), ``moe_permute_with_probs``/``moe_unpermute``, + :class:`GroupedLinear` routed experts, a shared expert + (:class:`LayerNormMLP`), ``Fp8Padding``/``Fp8Unpadding`` and optional + expert parallelism via ``ep_dispatch``/``ep_combine``. + + .. warning:: Work in progress, not functional yet. + """ + + def __init__(self, *args, **kwargs): + super().__init__() + raise NotImplementedError("DeepSeekV3MoE is under development") diff --git a/transformer_engine/pytorch/deepseek/multi_latent_attention.py b/transformer_engine/pytorch/deepseek/multi_latent_attention.py new file mode 100644 index 0000000000..6c2bb7420b --- /dev/null +++ b/transformer_engine/pytorch/deepseek/multi_latent_attention.py @@ -0,0 +1,24 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Multi-Latent Attention (MLA) block as used in DeepSeekV3.""" + +import torch + +__all__ = ["MultiLatentAttention"] + + +class MultiLatentAttention(torch.nn.Module): + """ + Multi-Latent Attention with low-rank Q/KV down-projections and a + decoupled RoPE/NoPE head split, composed from :class:`Linear`, + :class:`LayerNormLinear` and :class:`DotProductAttention` + (``kv_channels=(head_dim_qk, head_dim_v)``). + + .. warning:: Work in progress, not functional yet. + """ + + def __init__(self, *args, **kwargs): + super().__init__() + raise NotImplementedError("MultiLatentAttention is under development") diff --git a/transformer_engine/pytorch/deepseek/transformer_layer.py b/transformer_engine/pytorch/deepseek/transformer_layer.py new file mode 100644 index 0000000000..2a28a6ceb3 --- /dev/null +++ b/transformer_engine/pytorch/deepseek/transformer_layer.py @@ -0,0 +1,25 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""DeepSeekV3 transformer layer.""" + +import torch + +__all__ = ["DeepSeekV3Layer"] + + +class DeepSeekV3Layer(torch.nn.Module): + """ + A full DeepSeekV3 transformer layer, analogous to + :class:`TransformerLayer`: :class:`MultiLatentAttention` followed by + either a dense :class:`LayerNormMLP` (first layers) or + :class:`DeepSeekV3MoE`, with the same residual and fused + bias-dropout-add plumbing as :class:`TransformerLayer`. + + .. warning:: Work in progress, not functional yet. + """ + + def __init__(self, *args, **kwargs): + super().__init__() + raise NotImplementedError("DeepSeekV3Layer is under development") From c306c6f840bfe187cbdc3ecad39ad5aa517b669e Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 13:56:21 +0200 Subject: [PATCH 02/12] Move DeepSeekV3 skeleton to models/deepseek_v3 subpackage Co-Authored-By: Claude Fable 5 Signed-off-by: Pawel Gadzinski --- transformer_engine/pytorch/models/__init__.py | 13 +++++++++++++ .../{deepseek => models/deepseek_v3}/__init__.py | 8 +++++--- .../pytorch/{deepseek => models/deepseek_v3}/moe.py | 0 .../deepseek_v3}/multi_latent_attention.py | 0 .../deepseek_v3}/transformer_layer.py | 0 5 files changed, 18 insertions(+), 3 deletions(-) create mode 100644 transformer_engine/pytorch/models/__init__.py rename transformer_engine/pytorch/{deepseek => models/deepseek_v3}/__init__.py (50%) rename transformer_engine/pytorch/{deepseek => models/deepseek_v3}/moe.py (100%) rename transformer_engine/pytorch/{deepseek => models/deepseek_v3}/multi_latent_attention.py (100%) rename transformer_engine/pytorch/{deepseek => models/deepseek_v3}/transformer_layer.py (100%) diff --git a/transformer_engine/pytorch/models/__init__.py b/transformer_engine/pytorch/models/__init__.py new file mode 100644 index 0000000000..bee5474c81 --- /dev/null +++ b/transformer_engine/pytorch/models/__init__.py @@ -0,0 +1,13 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Model-specific transformer layers composed from Transformer Engine modules.""" + +from transformer_engine.pytorch.models.deepseek_v3 import ( + DeepSeekV3Layer, + DeepSeekV3MoE, + MultiLatentAttention, +) + +__all__ = ["DeepSeekV3Layer", "DeepSeekV3MoE", "MultiLatentAttention"] diff --git a/transformer_engine/pytorch/deepseek/__init__.py b/transformer_engine/pytorch/models/deepseek_v3/__init__.py similarity index 50% rename from transformer_engine/pytorch/deepseek/__init__.py rename to transformer_engine/pytorch/models/deepseek_v3/__init__.py index 5dafdf1fff..a7cbb50ae2 100644 --- a/transformer_engine/pytorch/deepseek/__init__.py +++ b/transformer_engine/pytorch/models/deepseek_v3/__init__.py @@ -4,8 +4,10 @@ """DeepSeekV3 transformer layer built from Transformer Engine MoE building blocks.""" -from transformer_engine.pytorch.deepseek.multi_latent_attention import MultiLatentAttention -from transformer_engine.pytorch.deepseek.moe import DeepSeekV3MoE -from transformer_engine.pytorch.deepseek.transformer_layer import DeepSeekV3Layer +from transformer_engine.pytorch.models.deepseek_v3.multi_latent_attention import ( + MultiLatentAttention, +) +from transformer_engine.pytorch.models.deepseek_v3.moe import DeepSeekV3MoE +from transformer_engine.pytorch.models.deepseek_v3.transformer_layer import DeepSeekV3Layer __all__ = ["DeepSeekV3Layer", "DeepSeekV3MoE", "MultiLatentAttention"] diff --git a/transformer_engine/pytorch/deepseek/moe.py b/transformer_engine/pytorch/models/deepseek_v3/moe.py similarity index 100% rename from transformer_engine/pytorch/deepseek/moe.py rename to transformer_engine/pytorch/models/deepseek_v3/moe.py diff --git a/transformer_engine/pytorch/deepseek/multi_latent_attention.py b/transformer_engine/pytorch/models/deepseek_v3/multi_latent_attention.py similarity index 100% rename from transformer_engine/pytorch/deepseek/multi_latent_attention.py rename to transformer_engine/pytorch/models/deepseek_v3/multi_latent_attention.py diff --git a/transformer_engine/pytorch/deepseek/transformer_layer.py b/transformer_engine/pytorch/models/deepseek_v3/transformer_layer.py similarity index 100% rename from transformer_engine/pytorch/deepseek/transformer_layer.py rename to transformer_engine/pytorch/models/deepseek_v3/transformer_layer.py From f73d04edaf0e16740f428b537c78d70d0c9c1ec1 Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 14:00:07 +0200 Subject: [PATCH 03/12] Add DeepSeekV3 layer entries to PyTorch API docs Co-Authored-By: Claude Fable 5 Signed-off-by: Pawel Gadzinski --- docs/api/pytorch.rst | 9 +++++++++ 1 file changed, 9 insertions(+) diff --git a/docs/api/pytorch.rst b/docs/api/pytorch.rst index 5fac0a89a6..bd3099b590 100644 --- a/docs/api/pytorch.rst +++ b/docs/api/pytorch.rst @@ -59,6 +59,15 @@ PyTorch .. autoapifunction:: transformer_engine.pytorch.deinterleave_glu_tensor +Model-specific layers +--------------------- + +.. autoapiclass:: transformer_engine.pytorch.models.DeepSeekV3Layer(**kwargs) + +.. autoapiclass:: transformer_engine.pytorch.models.DeepSeekV3MoE(**kwargs) + +.. autoapiclass:: transformer_engine.pytorch.models.MultiLatentAttention(**kwargs) + Data types ---------- From 09f28a9a3903b85ea28acff8ef63149738f38ea8 Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 14:13:44 +0200 Subject: [PATCH 04/12] [PyTorch] Implement DeepSeekV3Layer: MLA + MoE from TE building blocks MultiLatentAttention: low-rank q/kv latents (RMSNorm fused into LayerNormLinear up-projections), decoupled RoPE/NoPE head split with a shared key rope head, DotProductAttention with kv_channels=(qk, v) for the cuDNN fused backend. DeepSeekV3MoE: fused sigmoid router with aux-loss-free expert bias and grouped top-k, routed experts as te.ops GroupedLinear+ScaledSwiGLU+ GroupedLinear (CuTe fused grouped MLP on supported HW), probs applied per-token in the activation, local permute/unpermute or NCCL expert parallelism via ep_dispatch/ep_combine, optional shared expert. DeepSeekV3Layer: pre-RMSNorm + MLA and dense LayerNormMLP (RMSNorm, swiglu) or MoE with residual connections. Co-Authored-By: Claude Fable 5 Signed-off-by: Pawel Gadzinski --- tests/pytorch/test_deepseek.py | 124 +++++++++ transformer_engine/pytorch/__init__.py | 1 + .../pytorch/models/deepseek_v3/moe.py | 245 +++++++++++++++++- .../deepseek_v3/multi_latent_attention.py | 193 +++++++++++++- .../models/deepseek_v3/transformer_layer.py | 163 +++++++++++- 5 files changed, 702 insertions(+), 24 deletions(-) create mode 100644 tests/pytorch/test_deepseek.py diff --git a/tests/pytorch/test_deepseek.py b/tests/pytorch/test_deepseek.py new file mode 100644 index 0000000000..7778d0448c --- /dev/null +++ b/tests/pytorch/test_deepseek.py @@ -0,0 +1,124 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +import pytest +import torch + +from transformer_engine.pytorch.utils import deinterleave_glu_tensor +from transformer_engine.pytorch.models import ( + DeepSeekV3Layer, + DeepSeekV3MoE, + MultiLatentAttention, +) + +SEQ_LEN = 128 +BATCH = 2 +HIDDEN = 256 +HEADS = 4 +DTYPE = torch.bfloat16 + +MLA_KWARGS = dict( + q_lora_rank=96, + kv_lora_rank=64, + qk_nope_head_dim=64, + qk_rope_head_dim=32, + v_head_dim=64, +) + + +def _input(requires_grad=True): + torch.manual_seed(1234) + return torch.randn( + SEQ_LEN, BATCH, HIDDEN, dtype=DTYPE, device="cuda", requires_grad=requires_grad + ) + + +def test_mla_forward_backward(): + torch.manual_seed(0) + mla = MultiLatentAttention(HIDDEN, HEADS, params_dtype=DTYPE, **MLA_KWARGS) + x = _input() + out = mla(x) + assert out.shape == x.shape + out.sum().backward() + assert x.grad is not None and torch.isfinite(x.grad).all() + + +@pytest.mark.parametrize("shared", [False, True], ids=["no_shared", "shared"]) +@pytest.mark.parametrize("grouped", [False, True], ids=["ungrouped", "grouped"]) +def test_moe_forward_backward(shared, grouped): + torch.manual_seed(0) + moe = DeepSeekV3MoE( + HIDDEN, + moe_ffn_hidden_size=128, + num_experts=8, + topk=2, + num_groups=4 if grouped else None, + group_topk=2 if grouped else None, + shared_expert_ffn_hidden_size=128 if shared else None, + params_dtype=DTYPE, + ) + x = _input() + out = moe(x) + assert out.shape == x.shape + out.sum().backward() + assert x.grad is not None and torch.isfinite(x.grad).all() + + counts = moe._last_tokens_per_expert + assert counts.sum().item() == SEQ_LEN * BATCH * 2 + bias_before = moe.expert_bias.clone() + moe.update_expert_bias() + assert not torch.equal(bias_before, moe.expert_bias) + + +def test_moe_matches_dense_reference(): + """topk == num_experts with uniform probs must reduce to a sum of expert MLPs.""" + torch.manual_seed(0) + num_experts = 4 + moe = DeepSeekV3MoE( + HIDDEN, + moe_ffn_hidden_size=128, + num_experts=num_experts, + topk=num_experts, + routed_scaling_factor=1.0, + params_dtype=DTYPE, + ) + x = _input(requires_grad=False) + out = moe(x) + + tokens = x.reshape(-1, HIDDEN) + probs, _ = moe._route(moe.gate(tokens).float()) + fc1, _, fc2 = moe.experts + ref = torch.zeros_like(tokens) + for e in range(num_experts): + w1 = deinterleave_glu_tensor(getattr(fc1, f"weight{e}"), 32) + w2 = getattr(fc2, f"weight{e}") + gate_part, lin_part = (tokens @ w1.t()).chunk(2, dim=-1) + act = torch.nn.functional.silu(gate_part.float()) * lin_part.float() + ref += (act.to(DTYPE) * probs[:, e : e + 1].to(DTYPE)) @ w2.t() + torch.testing.assert_close(out.reshape(-1, HIDDEN), ref, rtol=0.05, atol=0.05) + + +@pytest.mark.parametrize("num_experts", [None, 8], ids=["dense", "moe"]) +def test_layer_forward_backward(num_experts): + torch.manual_seed(0) + layer = ( + DeepSeekV3Layer( + HIDDEN, + HEADS, + ffn_hidden_size=512, + num_experts=num_experts, + moe_ffn_hidden_size=128 if num_experts else None, + topk=2 if num_experts else None, + shared_expert_ffn_hidden_size=128 if num_experts else None, + params_dtype=DTYPE, + **MLA_KWARGS, + ) + if num_experts + else DeepSeekV3Layer(HIDDEN, HEADS, ffn_hidden_size=512, params_dtype=DTYPE, **MLA_KWARGS) + ) + x = _input() + out = layer(x) + assert out.shape == x.shape + out.sum().backward() + assert x.grad is not None and torch.isfinite(x.grad).all() diff --git a/transformer_engine/pytorch/__init__.py b/transformer_engine/pytorch/__init__.py index 2b1803bfb2..fae4d973e5 100644 --- a/transformer_engine/pytorch/__init__.py +++ b/transformer_engine/pytorch/__init__.py @@ -34,6 +34,7 @@ from transformer_engine.pytorch.attention import InferenceParams from transformer_engine.pytorch.attention import RotaryPositionEmbedding from transformer_engine.pytorch.transformer import TransformerLayer +from transformer_engine.pytorch import models from transformer_engine.pytorch.permutation import ( moe_permute, moe_permute_with_probs, diff --git a/transformer_engine/pytorch/models/deepseek_v3/moe.py b/transformer_engine/pytorch/models/deepseek_v3/moe.py index f4f787743b..5a1c8d650c 100644 --- a/transformer_engine/pytorch/models/deepseek_v3/moe.py +++ b/transformer_engine/pytorch/models/deepseek_v3/moe.py @@ -5,23 +5,248 @@ """DeepSeekV3 MoE block: sigmoid router with aux-loss-free bias, shared + routed experts.""" +from typing import Optional, Union + import torch +import transformer_engine.pytorch.ops as te_ops +from transformer_engine.pytorch.router import fused_topk_with_score_function +from transformer_engine.pytorch.permutation import moe_permute_with_probs, moe_unpermute + __all__ = ["DeepSeekV3MoE"] +def _make_expert_mlp(num_experts, hidden_size, ffn_hidden_size, dtype, device): + # GroupedLinear + ScaledSwiGLU + GroupedLinear fuses into a single CuTe + # grouped MLP on supported hardware; elsewhere it runs as three ops with + # the same API and checkpoint layout. + return te_ops.Sequential( + te_ops.GroupedLinear( + num_experts, hidden_size, 2 * ffn_hidden_size, bias=False, dtype=dtype, device=device + ), + te_ops.ScaledSwiGLU(glu_interleave_size=32), + te_ops.GroupedLinear( + num_experts, ffn_hidden_size, hidden_size, bias=False, dtype=dtype, device=device + ), + ) + + class DeepSeekV3MoE(torch.nn.Module): """ - DeepSeekV3-style Mixture of Experts block composed from TE MoE - primitives: ``fused_topk_with_score_function`` (sigmoid score function, - expert bias, grouped top-k), ``moe_permute_with_probs``/``moe_unpermute``, - :class:`GroupedLinear` routed experts, a shared expert - (:class:`LayerNormMLP`), ``Fp8Padding``/``Fp8Unpadding`` and optional - expert parallelism via ``ep_dispatch``/``ep_combine``. - - .. warning:: Work in progress, not functional yet. + DeepSeekV3-style Mixture of Experts block. + + Routing uses the fused sigmoid router with aux-loss-free expert bias and + node-limited (grouped) top-k (``fused_topk_with_score_function``). Routed + experts run as a grouped SwiGLU MLP built from ``te.ops`` (fusable into a + single CuTe grouped-GEMM kernel); routing probabilities are applied + per-token inside the expert MLP, so unpermute/combine is a plain + accumulation. Token routing is either local + (``moe_permute_with_probs``/``moe_unpermute``) or, when ``ep_group`` is + given, expert-parallel over NCCL (``ep_dispatch``/``ep_combine``). + + When expert parallelism is used, ``transformer_engine.pytorch.ep.ep_bootstrap`` + must be called once per process before the first forward, and inputs must + be bfloat16. + + Parameters + ---------- + hidden_size : int + size of each input sample. + moe_ffn_hidden_size : int + ffn size of each routed expert. + num_experts : int + total number of routed experts. + topk : int, default = 8 + number of experts per token. + num_groups : int, optional + number of expert groups for node-limited routing. + group_topk : int, optional + number of groups each token is limited to. + routed_scaling_factor : float, default = 2.5 + scaling applied to the routing probabilities. + shared_expert_ffn_hidden_size : int, optional + ffn size of the shared expert; ``None`` + disables the shared expert. + expert_bias_update_rate : float, default = 1e-3 + step size of the aux-loss-free bias update + (see :meth:`update_expert_bias`). + params_dtype : torch.dtype, optional + dtype of module parameters. + ep_group : ProcessGroup, optional + expert-parallel process group; enables the NCCL EP path. + ep_max_tokens_per_rank : int, optional + max local tokens per forward (required with EP). + ep_recv_capacity_per_rank : int, optional + receive-buffer capacity; defaults to + ``ep_size * ep_max_tokens_per_rank * topk``. + ep_alignment : int, default = 128 + per-expert row alignment of the EP receive buffer. """ - def __init__(self, *args, **kwargs): + def __init__( + self, + hidden_size: int, + moe_ffn_hidden_size: int, + num_experts: int, + topk: int = 8, + num_groups: Optional[int] = None, + group_topk: Optional[int] = None, + routed_scaling_factor: float = 2.5, + shared_expert_ffn_hidden_size: Optional[int] = None, + expert_bias_update_rate: float = 1e-3, + params_dtype: Optional[torch.dtype] = None, + device: Union[torch.device, str] = "cuda", + ep_group: Optional[torch.distributed.ProcessGroup] = None, + ep_max_tokens_per_rank: Optional[int] = None, + ep_recv_capacity_per_rank: Optional[int] = None, + ep_alignment: int = 128, + ) -> None: super().__init__() - raise NotImplementedError("DeepSeekV3MoE is under development") + + dtype = params_dtype if params_dtype is not None else torch.get_default_dtype() + self.hidden_size = hidden_size + self.num_experts = num_experts + self.topk = topk + self.num_groups = num_groups + self.group_topk = group_topk + self.routed_scaling_factor = routed_scaling_factor + self.expert_bias_update_rate = expert_bias_update_rate + + self.gate = torch.nn.Linear( + hidden_size, num_experts, bias=False, dtype=dtype, device=device + ) + self.register_buffer( + "expert_bias", torch.zeros(num_experts, dtype=torch.float32, device=device) + ) + self._last_tokens_per_expert: Optional[torch.Tensor] = None + + self.ep_group = ep_group + self.ep_size = 1 if ep_group is None else torch.distributed.get_world_size(ep_group) + assert num_experts % self.ep_size == 0 + num_local_experts = num_experts // self.ep_size + + self.experts = _make_expert_mlp( + num_local_experts, hidden_size, moe_ffn_hidden_size, dtype, device + ) + + self.shared_expert = None + if shared_expert_ffn_hidden_size is not None: + self.shared_expert = te_ops.Sequential( + te_ops.Linear( + hidden_size, + 2 * shared_expert_ffn_hidden_size, + bias=False, + dtype=dtype, + device=device, + ), + te_ops.SwiGLU(), + te_ops.Linear( + shared_expert_ffn_hidden_size, + hidden_size, + bias=False, + dtype=dtype, + device=device, + ), + ) + + self.ep_buffer = None + if ep_group is not None: + from transformer_engine.pytorch.ep import EpBuffer + + assert ep_max_tokens_per_rank is not None, "EP requires ep_max_tokens_per_rank." + if ep_recv_capacity_per_rank is None: + ep_recv_capacity_per_rank = self.ep_size * ep_max_tokens_per_rank * topk + self.ep_buffer = EpBuffer( + top_k=topk, + max_tokens_per_rank=ep_max_tokens_per_rank, + hidden_dim=hidden_size, + num_local_experts=num_local_experts, + recv_capacity_per_rank=ep_recv_capacity_per_rank, + alignment=ep_alignment, + device=device, + ) + + def _route(self, logits: torch.Tensor, topk_indices: Optional[torch.Tensor] = None): + return fused_topk_with_score_function( + logits=logits, + topk=self.topk, + use_pre_softmax=False, + num_groups=self.num_groups, + group_topk=self.group_topk, + scaling_factor=self.routed_scaling_factor, + score_function="sigmoid", + expert_bias=self.expert_bias, + topk_indices=topk_indices, + ) + + def _forward_local(self, tokens: torch.Tensor) -> torch.Tensor: + probs, routing_map = self._route(self.gate(tokens).float()) + tokens_per_expert = routing_map.sum(dim=0) + self._last_tokens_per_expert = tokens_per_expert.detach() + + num_out = tokens.shape[0] * self.topk + permuted, permuted_probs, row_id_map = moe_permute_with_probs( + tokens, probs, routing_map, num_out_tokens=num_out + ) + + # The fused grouped MLP requires the total row count to be a multiple + # of 128; rows beyond sum(tokens_per_expert) fall outside every group. + pad = (-num_out) % 128 + if pad: + permuted = torch.nn.functional.pad(permuted, (0, 0, 0, pad)) + permuted_probs = torch.nn.functional.pad(permuted_probs, (0, pad)) + + out = self.experts( + permuted, tokens_per_expert, permuted_probs.to(tokens.dtype), tokens_per_expert + ) + return moe_unpermute(out[:num_out], row_id_map, restore_shape=tokens.shape) + + def _forward_ep(self, tokens: torch.Tensor) -> torch.Tensor: + from transformer_engine.pytorch.ep import ep_dispatch, ep_combine + + assert tokens.dtype == torch.bfloat16, "The EP path requires bfloat16 inputs." + topk_idx = torch.empty( + (tokens.shape[0], self.topk), dtype=torch.int64, device=tokens.device + ) + probs, topk_idx = self._route(self.gate(tokens).float(), topk_indices=topk_idx) + self._last_tokens_per_expert = torch.bincount( + topk_idx.flatten(), minlength=self.num_experts + ) + topk_weights = probs.gather(1, topk_idx).float() + + recv_tokens, recv_weights, tokens_per_expert = ep_dispatch( + self.ep_buffer, tokens, topk_idx, topk_weights + ) + expert_out = self.experts( + recv_tokens, tokens_per_expert, recv_weights.to(tokens.dtype), tokens_per_expert + ) + return ep_combine(self.ep_buffer, expert_out, num_local_tokens=tokens.shape[0]) + + def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: + """ + Parameters + ---------- + hidden_states : torch.Tensor + input of shape ``[..., hidden_size]``. + """ + tokens = hidden_states.reshape(-1, self.hidden_size) + if self.ep_group is not None: + out = self._forward_ep(tokens) + else: + out = self._forward_local(tokens) + if self.shared_expert is not None: + out = out + self.shared_expert(tokens) + return out.view_as(hidden_states) + + @torch.no_grad() + def update_expert_bias(self) -> None: + """Aux-loss-free bias update from the last forward's routing counts. + + With data/expert parallelism, all-reduce ``_last_tokens_per_expert`` + across ranks before calling (or call on identically-routed ranks). + """ + counts = self._last_tokens_per_expert + if counts is None: + return + err = counts.float().mean() - counts.float() + self.expert_bias += self.expert_bias_update_rate * torch.sign(err) diff --git a/transformer_engine/pytorch/models/deepseek_v3/multi_latent_attention.py b/transformer_engine/pytorch/models/deepseek_v3/multi_latent_attention.py index 6c2bb7420b..a36075f2c5 100644 --- a/transformer_engine/pytorch/models/deepseek_v3/multi_latent_attention.py +++ b/transformer_engine/pytorch/models/deepseek_v3/multi_latent_attention.py @@ -4,21 +4,200 @@ """Multi-Latent Attention (MLA) block as used in DeepSeekV3.""" +from typing import Optional, Union + import torch +from transformer_engine.pytorch.module import Linear, LayerNormLinear +from transformer_engine.pytorch.attention import DotProductAttention, RotaryPositionEmbedding +from transformer_engine.pytorch.attention.rope import apply_rotary_pos_emb + __all__ = ["MultiLatentAttention"] class MultiLatentAttention(torch.nn.Module): """ - Multi-Latent Attention with low-rank Q/KV down-projections and a - decoupled RoPE/NoPE head split, composed from :class:`Linear`, - :class:`LayerNormLinear` and :class:`DotProductAttention` - (``kv_channels=(head_dim_qk, head_dim_v)``). + Multi-Latent Attention as used in DeepSeekV3. - .. warning:: Work in progress, not functional yet. + Queries and key-values are projected through low-rank latents + (``q_lora_rank``, ``kv_lora_rank``); RMSNorm on each latent is fused into + the up-projection (:class:`LayerNormLinear` with RMSNorm). Each query/key + head is split into a ``qk_nope_head_dim`` part and a ``qk_rope_head_dim`` + part; RoPE is applied only to the rope part, and the key rope part comes + from a single shared head broadcast to all heads. Attention runs through + :class:`DotProductAttention` with asymmetric head dims + ``kv_channels=(qk_nope_head_dim + qk_rope_head_dim, v_head_dim)``, which + supports the cuDNN fused attention backend. + + Parameters + ---------- + hidden_size : int + size of each input sample. + num_attention_heads : int + number of attention heads. + q_lora_rank : int, default = 1536 + rank of the query latent. + kv_lora_rank : int, default = 512 + rank of the key-value latent. + qk_nope_head_dim : int, default = 128 + per-head dim of the non-rotary query/key part. + qk_rope_head_dim : int, default = 64 + per-head dim of the rotary query/key part. + v_head_dim : int, default = 128 + per-head dim of the values. + attention_dropout : float, default = 0.0 + dropout probability on attention scores. + attn_mask_type : str, default = "causal" + attention mask type passed to :class:`DotProductAttention`. + rotary_base : float, default = 10000.0 + RoPE base. + softmax_scale : float, optional + softmax scale; defaults to ``1/sqrt(qk head dim)`` inside + :class:`DotProductAttention`. + qkv_format : str, default = "sbhd" + layout of the input/output tensors. + params_dtype : torch.dtype, optional + dtype of module parameters. + tp_group : ProcessGroup, optional + tensor-parallel process group for the up/output projections. + tp_size : int, default = 1 + tensor-parallel world size. """ - def __init__(self, *args, **kwargs): + def __init__( + self, + hidden_size: int, + num_attention_heads: int, + q_lora_rank: int = 1536, + kv_lora_rank: int = 512, + qk_nope_head_dim: int = 128, + qk_rope_head_dim: int = 64, + v_head_dim: int = 128, + attention_dropout: float = 0.0, + attn_mask_type: str = "causal", + rotary_base: float = 10000.0, + softmax_scale: Optional[float] = None, + qkv_format: str = "sbhd", + params_dtype: Optional[torch.dtype] = None, + tp_group: Optional[torch.distributed.ProcessGroup] = None, + tp_size: int = 1, + device: Union[torch.device, str] = "cuda", + ) -> None: super().__init__() - raise NotImplementedError("MultiLatentAttention is under development") + + assert qkv_format in ("sbhd", "bshd"), "MultiLatentAttention supports sbhd/bshd formats." + assert num_attention_heads % tp_size == 0 + + self.qkv_format = qkv_format + self.num_attention_heads = num_attention_heads + self.num_attention_heads_per_partition = num_attention_heads // tp_size + self.qk_nope_head_dim = qk_nope_head_dim + self.qk_rope_head_dim = qk_rope_head_dim + self.qk_head_dim = qk_nope_head_dim + qk_rope_head_dim + self.v_head_dim = v_head_dim + self.kv_lora_rank = kv_lora_rank + + common = {"bias": False, "params_dtype": params_dtype, "device": device} + tp = {"tp_group": tp_group, "tp_size": tp_size} + + self.q_down_proj = Linear(hidden_size, q_lora_rank, **common) + self.q_up_proj = LayerNormLinear( + q_lora_rank, + num_attention_heads * self.qk_head_dim, + normalization="RMSNorm", + parallel_mode="column" if tp_size > 1 else None, + **tp, + **common, + ) + self.kv_down_proj = Linear(hidden_size, kv_lora_rank + qk_rope_head_dim, **common) + self.kv_up_proj = LayerNormLinear( + kv_lora_rank, + num_attention_heads * (qk_nope_head_dim + v_head_dim), + normalization="RMSNorm", + parallel_mode="column" if tp_size > 1 else None, + **tp, + **common, + ) + self.out_proj = Linear( + num_attention_heads * v_head_dim, + hidden_size, + parallel_mode="row" if tp_size > 1 else None, + **tp, + **common, + ) + + self.rope = RotaryPositionEmbedding(qk_rope_head_dim, rotary_base=rotary_base) + self._rope_freqs: Optional[torch.Tensor] = None + + self.core_attention = DotProductAttention( + num_attention_heads, + kv_channels=(self.qk_head_dim, v_head_dim), + attention_dropout=attention_dropout, + qkv_format=qkv_format, + attn_mask_type=attn_mask_type, + softmax_scale=softmax_scale, + tp_group=tp_group, + tp_size=tp_size, + ) + + def _rope_freqs_for(self, seq_len: int, device: torch.device) -> torch.Tensor: + if self._rope_freqs is None or self._rope_freqs.shape[0] < seq_len: + self._rope_freqs = self.rope(seq_len).to(device) + return self._rope_freqs[:seq_len] + + def forward( + self, + hidden_states: torch.Tensor, + attention_mask: Optional[torch.Tensor] = None, + attn_mask_type: Optional[str] = None, + checkpoint_core_attention: bool = False, + ) -> torch.Tensor: + """ + Parameters + ---------- + hidden_states : torch.Tensor + input of shape ``[sq, b, h]`` (sbhd) or ``[b, sq, h]`` (bshd). + attention_mask : torch.Tensor, optional + boolean mask passed to :class:`DotProductAttention`. + attn_mask_type : str, optional + override of the constructor's mask type. + checkpoint_core_attention : bool, default = False + checkpoint the core attention computation. + """ + seq_dim = 0 if self.qkv_format == "sbhd" else 1 + seq_len = hidden_states.shape[seq_dim] + heads = self.num_attention_heads_per_partition + + q = self.q_up_proj(self.q_down_proj(hidden_states)) + q = q.view(*q.shape[:-1], heads, self.qk_head_dim) + + kv_down = self.kv_down_proj(hidden_states) + kv_latent, k_pos = torch.split(kv_down, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1) + kv = self.kv_up_proj(kv_latent) + kv = kv.view(*kv.shape[:-1], heads, self.qk_nope_head_dim + self.v_head_dim) + k_nope, v = torch.split(kv, [self.qk_nope_head_dim, self.v_head_dim], dim=-1) + + freqs = self._rope_freqs_for(seq_len, hidden_states.device) + q_rope = apply_rotary_pos_emb( + q[..., self.qk_nope_head_dim :].contiguous(), + freqs, + tensor_format=self.qkv_format, + fused=True, + ) + k_rope = apply_rotary_pos_emb( + k_pos.unsqueeze(-2), freqs, tensor_format=self.qkv_format, fused=True + ) + + q = torch.cat([q[..., : self.qk_nope_head_dim], q_rope], dim=-1) + k = torch.cat([k_nope, k_rope.expand(*k_nope.shape[:-1], -1)], dim=-1) + + context = self.core_attention( + q, + k, + v.contiguous(), + attention_mask=attention_mask, + qkv_format=self.qkv_format, + attn_mask_type=attn_mask_type, + checkpoint_core_attention=checkpoint_core_attention, + ) + return self.out_proj(context) diff --git a/transformer_engine/pytorch/models/deepseek_v3/transformer_layer.py b/transformer_engine/pytorch/models/deepseek_v3/transformer_layer.py index 2a28a6ceb3..af1eeb1a95 100644 --- a/transformer_engine/pytorch/models/deepseek_v3/transformer_layer.py +++ b/transformer_engine/pytorch/models/deepseek_v3/transformer_layer.py @@ -4,22 +4,171 @@ """DeepSeekV3 transformer layer.""" +from typing import Optional, Union + import torch +from transformer_engine.pytorch.module import LayerNormMLP, RMSNorm +from transformer_engine.pytorch.models.deepseek_v3.multi_latent_attention import ( + MultiLatentAttention, +) +from transformer_engine.pytorch.models.deepseek_v3.moe import DeepSeekV3MoE + __all__ = ["DeepSeekV3Layer"] class DeepSeekV3Layer(torch.nn.Module): """ A full DeepSeekV3 transformer layer, analogous to - :class:`TransformerLayer`: :class:`MultiLatentAttention` followed by - either a dense :class:`LayerNormMLP` (first layers) or - :class:`DeepSeekV3MoE`, with the same residual and fused - bias-dropout-add plumbing as :class:`TransformerLayer`. + :class:`TransformerLayer`: pre-RMSNorm + :class:`MultiLatentAttention`, + then either a dense SwiGLU MLP (:class:`LayerNormMLP` with RMSNorm, used + for the first dense layers of DeepSeekV3) or :class:`DeepSeekV3MoE`, each + with a residual connection. - .. warning:: Work in progress, not functional yet. + Parameters + ---------- + hidden_size : int + size of each input sample. + num_attention_heads : int + number of attention heads. + ffn_hidden_size : int + ffn size of the dense MLP (used when ``num_experts`` is + ``None``). + num_experts : int, optional + number of routed experts; ``None`` makes this a dense layer. + moe_ffn_hidden_size : int, optional + ffn size of each routed expert (required with MoE). + hidden_dropout : float, default = 0.0 + dropout probability on the residual branches. + kwargs common to the submodules (``q_lora_rank``, ``kv_lora_rank``, + ``qk_nope_head_dim``, ``qk_rope_head_dim``, ``v_head_dim``, + ``attention_dropout``, ``attn_mask_type``, ``qkv_format``, ``topk``, + ``num_groups``, ``group_topk``, ``routed_scaling_factor``, + ``shared_expert_ffn_hidden_size``, EP options, ...) are forwarded to + :class:`MultiLatentAttention` and :class:`DeepSeekV3MoE`. """ - def __init__(self, *args, **kwargs): + _MLA_KWARGS = frozenset( + { + "q_lora_rank", + "kv_lora_rank", + "qk_nope_head_dim", + "qk_rope_head_dim", + "v_head_dim", + "attention_dropout", + "attn_mask_type", + "rotary_base", + "softmax_scale", + "qkv_format", + "tp_group", + "tp_size", + } + ) + _MOE_KWARGS = frozenset( + { + "topk", + "num_groups", + "group_topk", + "routed_scaling_factor", + "shared_expert_ffn_hidden_size", + "expert_bias_update_rate", + "ep_group", + "ep_max_tokens_per_rank", + "ep_recv_capacity_per_rank", + "ep_alignment", + } + ) + + def __init__( + self, + hidden_size: int, + num_attention_heads: int, + ffn_hidden_size: Optional[int] = None, + num_experts: Optional[int] = None, + moe_ffn_hidden_size: Optional[int] = None, + hidden_dropout: float = 0.0, + layernorm_epsilon: float = 1e-5, + params_dtype: Optional[torch.dtype] = None, + device: Union[torch.device, str] = "cuda", + **kwargs, + ) -> None: super().__init__() - raise NotImplementedError("DeepSeekV3Layer is under development") + + unknown = set(kwargs) - self._MLA_KWARGS - self._MOE_KWARGS + if unknown: + raise TypeError(f"Unexpected keyword arguments: {sorted(unknown)}") + mla_kwargs = {k: v for k, v in kwargs.items() if k in self._MLA_KWARGS} + moe_kwargs = {k: v for k, v in kwargs.items() if k in self._MOE_KWARGS} + + self.hidden_dropout = hidden_dropout + + self.input_layernorm = RMSNorm( + hidden_size, eps=layernorm_epsilon, device=device, dtype=params_dtype + ) + self.self_attention = MultiLatentAttention( + hidden_size, + num_attention_heads, + params_dtype=params_dtype, + device=device, + **mla_kwargs, + ) + + if num_experts is None: + assert ffn_hidden_size is not None, "Dense layers require ffn_hidden_size." + self.pre_mlp_layernorm = None + self.mlp = LayerNormMLP( + hidden_size, + ffn_hidden_size, + eps=layernorm_epsilon, + normalization="RMSNorm", + activation="swiglu", + bias=False, + params_dtype=params_dtype, + device=device, + ) + else: + assert moe_ffn_hidden_size is not None, "MoE layers require moe_ffn_hidden_size." + self.pre_mlp_layernorm = RMSNorm( + hidden_size, eps=layernorm_epsilon, device=device, dtype=params_dtype + ) + self.mlp = DeepSeekV3MoE( + hidden_size, + moe_ffn_hidden_size, + num_experts, + params_dtype=params_dtype, + device=device, + **moe_kwargs, + ) + + def _residual_add(self, out: torch.Tensor, residual: torch.Tensor) -> torch.Tensor: + out = torch.nn.functional.dropout(out, p=self.hidden_dropout, training=self.training) + return residual + out + + def forward( + self, + hidden_states: torch.Tensor, + attention_mask: Optional[torch.Tensor] = None, + checkpoint_core_attention: bool = False, + ) -> torch.Tensor: + """ + Parameters + ---------- + hidden_states : torch.Tensor + input of shape ``[sq, b, h]`` (sbhd) or ``[b, sq, h]`` (bshd). + attention_mask : torch.Tensor, optional + boolean attention mask. + checkpoint_core_attention : bool, default = False + checkpoint the core attention computation. + """ + attention_out = self.self_attention( + self.input_layernorm(hidden_states), + attention_mask=attention_mask, + checkpoint_core_attention=checkpoint_core_attention, + ) + hidden_states = self._residual_add(attention_out, hidden_states) + + if self.pre_mlp_layernorm is not None: + mlp_out = self.mlp(self.pre_mlp_layernorm(hidden_states)) + else: + mlp_out = self.mlp(hidden_states) + return self._residual_add(mlp_out, hidden_states) From e23100b73cef8e7f7c955af8f626f83177aef06e Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 14:28:46 +0200 Subject: [PATCH 05/12] Add distributed EP test for DeepSeekV3 MoE/layer run_deepseek_ep.py checks the EP path against the all-experts-local path numerically (forward, input/gate grads, all-reduced expert wgrads) and smoke-tests the full layer with EP. Also size the default EP recv capacity for per-expert alignment padding and the fused grouped MLP's row-count requirement. Co-Authored-By: Claude Fable 5 Signed-off-by: Pawel Gadzinski --- tests/pytorch/distributed/run_deepseek_ep.py | 185 ++++++++++++++++++ .../distributed/run_test_deepseek_ep.sh | 52 +++++ tests/pytorch/distributed/test_deepseek_ep.py | 26 +++ .../pytorch/models/deepseek_v3/moe.py | 6 +- 4 files changed, 268 insertions(+), 1 deletion(-) create mode 100644 tests/pytorch/distributed/run_deepseek_ep.py create mode 100644 tests/pytorch/distributed/run_test_deepseek_ep.sh create mode 100644 tests/pytorch/distributed/test_deepseek_ep.py diff --git a/tests/pytorch/distributed/run_deepseek_ep.py b/tests/pytorch/distributed/run_deepseek_ep.py new file mode 100644 index 0000000000..0edae05961 --- /dev/null +++ b/tests/pytorch/distributed/run_deepseek_ep.py @@ -0,0 +1,185 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. +"""Multi-process DeepSeekV3 MoE/layer EP tests, launched via torchrun.""" + +import os +import sys +import unittest + +import torch +import torch.distributed as dist + +from transformer_engine.pytorch.ep import ep_bootstrap, ep_finalize, release_symm_mem_pool +from transformer_engine.pytorch.models import DeepSeekV3Layer, DeepSeekV3MoE + +HIDDEN = 256 +MOE_FFN = 128 +SHARED_FFN = 128 +NUM_LOCAL_EXPERTS = 2 +TOP_K = 2 +TOKENS_PER_RANK = 64 +HEADS = 4 +DTYPE = torch.bfloat16 + +MLA_KWARGS = dict( + q_lora_rank=96, + kv_lora_rank=64, + qk_nope_head_dim=64, + qk_rope_head_dim=32, + v_head_dim=64, +) + + +def _device_sm() -> int: + major, minor = torch.cuda.get_device_capability() + return major * 10 + minor + + +def _recv_capacity(ep_size: int) -> int: + cap = ep_size * TOKENS_PER_RANK * TOP_K + NUM_LOCAL_EXPERTS * 128 + return -(-cap // 128) * 128 + + +def _broadcast_params(module: torch.nn.Module) -> None: + for t in list(module.parameters()) + list(module.buffers()): + dist.broadcast(t.detach(), src=0) + + +class TestDeepSeekEP(unittest.TestCase): + @classmethod + def setUpClass(cls): + if _device_sm() < 90: + raise unittest.SkipTest(f"NCCL EP requires SM>=90 (got SM{_device_sm()})") + cls.rank = dist.get_rank() + cls.ep_size = dist.get_world_size() + cls.num_experts = NUM_LOCAL_EXPERTS * cls.ep_size + world_pg = dist.distributed_c10d._get_default_group() + cls.ep_group = dist.new_group(ranks=list(range(world_pg.size())), backend="nccl") + ep_bootstrap( + cls.ep_group, + num_experts=cls.num_experts, + max_tokens_per_rank=TOKENS_PER_RANK, + hidden_dim=HIDDEN, + num_topk=TOP_K, + recv_capacity_per_rank=_recv_capacity(cls.ep_size), + ) + + def _make_moe(self, ep: bool, shared: bool = True) -> DeepSeekV3MoE: + return DeepSeekV3MoE( + HIDDEN, + moe_ffn_hidden_size=MOE_FFN, + num_experts=self.num_experts, + topk=TOP_K, + shared_expert_ffn_hidden_size=SHARED_FFN if shared else None, + params_dtype=DTYPE, + ep_group=self.ep_group if ep else None, + ep_max_tokens_per_rank=TOKENS_PER_RANK if ep else None, + ep_recv_capacity_per_rank=_recv_capacity(self.ep_size) if ep else None, + ) + + def _copy_local_expert_weights(self, ep_moe: DeepSeekV3MoE, ref: DeepSeekV3MoE) -> None: + with torch.no_grad(): + ep_moe.gate.weight.copy_(ref.gate.weight) + if ref.shared_expert is not None: + for dst, src in zip( + ep_moe.shared_expert.parameters(), ref.shared_expert.parameters() + ): + dst.copy_(src) + ep_fc1, _, ep_fc2 = ep_moe.experts + ref_fc1, _, ref_fc2 = ref.experts + for local_e in range(NUM_LOCAL_EXPERTS): + global_e = self.rank * NUM_LOCAL_EXPERTS + local_e + getattr(ep_fc1, f"weight{local_e}").copy_(getattr(ref_fc1, f"weight{global_e}")) + getattr(ep_fc2, f"weight{local_e}").copy_(getattr(ref_fc2, f"weight{global_e}")) + + def test_moe_ep_matches_local(self): + """EP MoE must match the single-GPU (all-experts-local) path numerically.""" + torch.manual_seed(0) + ref = self._make_moe(ep=False) + _broadcast_params(ref) + ep_moe = self._make_moe(ep=True) + self._copy_local_expert_weights(ep_moe, ref) + + torch.manual_seed(1234 + self.rank) + x = torch.randn(TOKENS_PER_RANK, HIDDEN, dtype=DTYPE, device="cuda") + x_ep = x.clone().requires_grad_(True) + x_ref = x.clone().requires_grad_(True) + + out_ep = ep_moe(x_ep) + out_ref = ref(x_ref) + torch.testing.assert_close(out_ep, out_ref, rtol=0.05, atol=0.05) + + grad_out = torch.randn_like(out_ep) + out_ep.backward(grad_out) + out_ref.backward(grad_out) + torch.testing.assert_close(x_ep.grad, x_ref.grad, rtol=0.05, atol=0.05) + torch.testing.assert_close( + ep_moe.gate.weight.grad, ref.gate.weight.grad, rtol=0.1, atol=0.1 + ) + + # A local expert's wgrad on its owner rank equals the sum of the + # reference wgrads over all ranks. + ep_fc1, _, ep_fc2 = ep_moe.experts + ref_fc1, _, ref_fc2 = ref.experts + for ep_fc, ref_fc in ((ep_fc1, ref_fc1), (ep_fc2, ref_fc2)): + for local_e in range(NUM_LOCAL_EXPERTS): + global_e = self.rank * NUM_LOCAL_EXPERTS + local_e + ref_grad = getattr(ref_fc, f"weight{global_e}").grad.float() + dist.all_reduce(ref_grad) + ep_grad = getattr(ep_fc, f"weight{local_e}").grad.float() + torch.testing.assert_close(ep_grad, ref_grad, rtol=0.1, atol=0.1) + + counts = ep_moe._last_tokens_per_expert.clone() + dist.all_reduce(counts) + self.assertEqual(counts.sum().item(), self.ep_size * TOKENS_PER_RANK * TOP_K) + + def test_layer_ep_forward_backward(self): + """Full DeepSeekV3Layer smoke test with an EP MoE block.""" + torch.manual_seed(10 + self.rank) + layer = DeepSeekV3Layer( + HIDDEN, + HEADS, + num_experts=self.num_experts, + moe_ffn_hidden_size=MOE_FFN, + topk=TOP_K, + shared_expert_ffn_hidden_size=SHARED_FFN, + params_dtype=DTYPE, + ep_group=self.ep_group, + ep_max_tokens_per_rank=TOKENS_PER_RANK, + ep_recv_capacity_per_rank=_recv_capacity(self.ep_size), + **MLA_KWARGS, + ) + x = torch.randn( + TOKENS_PER_RANK // 2, 2, HIDDEN, dtype=DTYPE, device="cuda", requires_grad=True + ) + out = layer(x) + self.assertEqual(out.shape, x.shape) + out.sum().backward() + self.assertIsNotNone(x.grad) + self.assertTrue(torch.isfinite(x.grad).all()) + + layer.mlp.update_expert_bias() + self.assertTrue(torch.isfinite(layer.mlp.expert_bias).all()) + + +def _init_distributed(): + dist.init_process_group(backend="nccl") + torch.cuda.set_device(int(os.environ["LOCAL_RANK"])) + try: + from torch.distributed import _symmetric_memory as _symm_mem + + _symm_mem.set_backend("NCCL") + except (ImportError, RuntimeError): + pass + + +if __name__ == "__main__": + _init_distributed() + suite = unittest.TestLoader().loadTestsFromTestCase(TestDeepSeekEP) + result = unittest.TextTestRunner(stream=sys.stdout, verbosity=2).run(suite) + dist.barrier() + ep_finalize() + release_symm_mem_pool() + dist.destroy_process_group() + sys.exit(0 if result.wasSuccessful() else 1) diff --git a/tests/pytorch/distributed/run_test_deepseek_ep.sh b/tests/pytorch/distributed/run_test_deepseek_ep.sh new file mode 100644 index 0000000000..8c0bbbc5b9 --- /dev/null +++ b/tests/pytorch/distributed/run_test_deepseek_ep.sh @@ -0,0 +1,52 @@ +#!/bin/bash +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. +# +# Launcher for tests/pytorch/distributed/run_deepseek_ep.py. Auto-detects GPU count. + +set -uo pipefail + +SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)" + +DETECTED_GPUS=$(nvidia-smi -L 2>/dev/null | wc -l) +if [ "${DETECTED_GPUS}" -lt 2 ]; then + echo "DeepSeek EP test requires >= 2 GPUs (found ${DETECTED_GPUS}); SKIPPING." + exit 0 +fi + +# NCCL EP requires active NVLink P2P among ranks on the node. +if ! nvidia-smi nvlink --status 2>/dev/null | grep -qE 'Link [0-9]+:.*GB/s'; then + echo "No NVLink between GPUs (PCIe-only fabric); NCCL EP is unsupported here. SKIPPING." + exit 0 +fi + +NUM_RANKS="${NVTE_TEST_EP_NUM_RANKS:-${DETECTED_GPUS}}" +if [ "${NUM_RANKS}" -gt 8 ]; then NUM_RANKS=8; fi + +TEST_TIMEOUT_S="${TEST_TIMEOUT_S:-180}" + +: ${NCCL_EP_JIT_CACHE_DIR:="${TMPDIR:-/tmp}/nccl_ep_jit_cache_$(id -u)"} +export NCCL_EP_JIT_CACHE_DIR +mkdir -p "$NCCL_EP_JIT_CACHE_DIR" + +SCRIPT="${SCRIPT_DIR}/run_deepseek_ep.py" +LOG="stdout_deepseek_ep.txt" + +echo "=== Running ${SCRIPT} on ${NUM_RANKS} GPUs (timeout=${TEST_TIMEOUT_S}s) ===" +setsid timeout --foreground --kill-after=10 --signal=TERM "${TEST_TIMEOUT_S}" \ + torchrun --standalone --nnodes=1 --nproc-per-node="${NUM_RANKS}" \ + "${SCRIPT}" 2>&1 | tee "${LOG}" +RC=${PIPESTATUS[0]} +pkill -9 -f "tests/pytorch/distributed/run_deepseek_ep.py" 2>/dev/null || true + +RET=0 +if [ "${RC}" -ne 0 ]; then echo "torchrun exited with ${RC}"; RET=1; fi +if grep -qE "(^|]:)FAILED|(^|]:)Traceback" "${LOG}"; then RET=1; fi +if ! grep -qE "Ran [0-9]+ test|^OK$" "${LOG}"; then + echo "ERROR: no test summary — likely hang or early crash" + RET=1 +fi +if [ -z "${KEEP_EP_LOGS:-}" ]; then rm -f "${LOG}"; fi + +exit $RET diff --git a/tests/pytorch/distributed/test_deepseek_ep.py b/tests/pytorch/distributed/test_deepseek_ep.py new file mode 100644 index 0000000000..4a4d9a8dea --- /dev/null +++ b/tests/pytorch/distributed/test_deepseek_ep.py @@ -0,0 +1,26 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. +"""Pytest driver — spawns run_deepseek_ep.py under torchrun and asserts it passed.""" + +import os +import subprocess +from pathlib import Path + +import pytest +import torch + +TEST_ROOT = Path(__file__).parent.resolve() +LAUNCHER = TEST_ROOT / "run_test_deepseek_ep.sh" + + +@pytest.mark.skipif(torch.cuda.device_count() < 2, reason="DeepSeek EP requires >= 2 GPUs") +def test_multi_process_deepseek_ep(): + timeout_s = int(os.environ.get("NVTE_TEST_EP_TIMEOUT_S", "180")) + proc = subprocess.run( + ["bash", str(LAUNCHER)], + env={**os.environ, "KEEP_EP_LOGS": "1", "TEST_TIMEOUT_S": str(timeout_s)}, + timeout=timeout_s + 30, + check=False, + ) + assert proc.returncode == 0, f"DeepSeek EP test suite failed (rc={proc.returncode})" diff --git a/transformer_engine/pytorch/models/deepseek_v3/moe.py b/transformer_engine/pytorch/models/deepseek_v3/moe.py index 5a1c8d650c..f413221bd1 100644 --- a/transformer_engine/pytorch/models/deepseek_v3/moe.py +++ b/transformer_engine/pytorch/models/deepseek_v3/moe.py @@ -155,7 +155,11 @@ def __init__( assert ep_max_tokens_per_rank is not None, "EP requires ep_max_tokens_per_rank." if ep_recv_capacity_per_rank is None: - ep_recv_capacity_per_rank = self.ep_size * ep_max_tokens_per_rank * topk + # Worst case plus per-expert alignment padding, rounded up to + # the multiple of 128 required by the fused grouped MLP. + cap = self.ep_size * ep_max_tokens_per_rank * topk + cap += num_local_experts * max(ep_alignment, 1) + ep_recv_capacity_per_rank = -(-cap // 128) * 128 self.ep_buffer = EpBuffer( top_k=topk, max_tokens_per_rank=ep_max_tokens_per_rank, From 4c6e1e8aff62def1cd0bd72ce9bfb960a4aab035 Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 15:49:08 +0200 Subject: [PATCH 06/12] Fix EP wgrad test collective + zero EP recv/grad buffers The per-expert wgrad check called all_reduce on different tensors per rank (rank-local experts), corrupting the reference grads; reduce every expert's grad on every rank instead. Also pass zero-filled recv/grad buffers to ep_dispatch/ep_combine so alignment-padding rows inside the grouped-GEMM m_splits can never poison expert wgrads. Verified on lyris (4x GB300, arm64): run_test_deepseek_ep.sh passes on all ranks (EP forward/dgrad/gate-grad/expert-wgrad match the all-local reference; full-layer EP smoke passes). Co-Authored-By: Claude Fable 5 Signed-off-by: Pawel Gadzinski --- tests/pytorch/distributed/run_deepseek_ep.py | 12 +++++++---- .../pytorch/models/deepseek_v3/moe.py | 20 +++++++++++++++++-- 2 files changed, 26 insertions(+), 6 deletions(-) diff --git a/tests/pytorch/distributed/run_deepseek_ep.py b/tests/pytorch/distributed/run_deepseek_ep.py index 0edae05961..bf756b69ad 100644 --- a/tests/pytorch/distributed/run_deepseek_ep.py +++ b/tests/pytorch/distributed/run_deepseek_ep.py @@ -119,16 +119,20 @@ def test_moe_ep_matches_local(self): ) # A local expert's wgrad on its owner rank equals the sum of the - # reference wgrads over all ranks. + # reference wgrads over all ranks. all_reduce is collective, so every + # rank must reduce every expert's grad (in the same order). ep_fc1, _, ep_fc2 = ep_moe.experts ref_fc1, _, ref_fc2 = ref.experts for ep_fc, ref_fc in ((ep_fc1, ref_fc1), (ep_fc2, ref_fc2)): + ref_grads = [ + getattr(ref_fc, f"weight{e}").grad.float().clone() for e in range(self.num_experts) + ] + for g in ref_grads: + dist.all_reduce(g) for local_e in range(NUM_LOCAL_EXPERTS): global_e = self.rank * NUM_LOCAL_EXPERTS + local_e - ref_grad = getattr(ref_fc, f"weight{global_e}").grad.float() - dist.all_reduce(ref_grad) ep_grad = getattr(ep_fc, f"weight{local_e}").grad.float() - torch.testing.assert_close(ep_grad, ref_grad, rtol=0.1, atol=0.1) + torch.testing.assert_close(ep_grad, ref_grads[global_e], rtol=0.1, atol=0.1) counts = ep_moe._last_tokens_per_expert.clone() dist.all_reduce(counts) diff --git a/transformer_engine/pytorch/models/deepseek_v3/moe.py b/transformer_engine/pytorch/models/deepseek_v3/moe.py index f413221bd1..3c182b4405 100644 --- a/transformer_engine/pytorch/models/deepseek_v3/moe.py +++ b/transformer_engine/pytorch/models/deepseek_v3/moe.py @@ -218,13 +218,29 @@ def _forward_ep(self, tokens: torch.Tensor) -> torch.Tensor: ) topk_weights = probs.gather(1, topk_idx).float() + # Zero-filled recv/grad buffers: per-expert alignment padding lands + # inside the grouped-GEMM m_splits, so uninitialized rows would poison + # the expert wgrads. + cap = self.ep_buffer.recv_capacity_per_rank recv_tokens, recv_weights, tokens_per_expert = ep_dispatch( - self.ep_buffer, tokens, topk_idx, topk_weights + self.ep_buffer, + tokens, + topk_idx, + topk_weights, + recv_tokens=torch.zeros( + (cap, self.hidden_size), dtype=tokens.dtype, device=tokens.device + ), + recv_topk_weights=torch.zeros((cap,), dtype=torch.float32, device=tokens.device), ) expert_out = self.experts( recv_tokens, tokens_per_expert, recv_weights.to(tokens.dtype), tokens_per_expert ) - return ep_combine(self.ep_buffer, expert_out, num_local_tokens=tokens.shape[0]) + return ep_combine( + self.ep_buffer, + expert_out, + num_local_tokens=tokens.shape[0], + grad_out=torch.zeros_like(expert_out), + ) def forward(self, hidden_states: torch.Tensor) -> torch.Tensor: """ From aa17c37fb9a0e8cd74c3b5d67a5d34365f4133d7 Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 17:23:10 +0200 Subject: [PATCH 07/12] Use fused MLA RoPE kernels in MultiLatentAttention Move the Triton MLA RoPE kernels (Megatron-LM fused_mla_yarn_rope_apply port) from tests/pytorch/attention/ mla_rope_utils.py into models/deepseek_v3/mla_rope.py and use them in MultiLatentAttention: the q kernel rotates the rope slice in place and the kv kernel assembles key/value in a single pass, removing the torch.cat/expand/contiguous copies (~10% of layer GPU time). PyTorch fallback (same convention) covers missing Triton and bshd. Fix a latent bug from the test util: the q backward kernel assumed a contiguous incoming gradient, but cuDNN attention backward can hand over a strided one (allocator-state dependent IMA). The old test file stays as a compat shim. Add a Triton-vs-PyTorch parity test. Co-Authored-By: Claude Fable 5 Signed-off-by: Pawel Gadzinski --- tests/pytorch/attention/mla_rope_utils.py | 652 +----------------- tests/pytorch/test_deepseek.py | 60 ++ .../pytorch/models/deepseek_v3/mla_rope.py | 495 +++++++++++++ .../deepseek_v3/multi_latent_attention.py | 55 +- 4 files changed, 601 insertions(+), 661 deletions(-) create mode 100644 transformer_engine/pytorch/models/deepseek_v3/mla_rope.py diff --git a/tests/pytorch/attention/mla_rope_utils.py b/tests/pytorch/attention/mla_rope_utils.py index 90eebfc66a..d022757886 100644 --- a/tests/pytorch/attention/mla_rope_utils.py +++ b/tests/pytorch/attention/mla_rope_utils.py @@ -2,26 +2,17 @@ # # See LICENSE for license information. -"""MLA RoPE for DSv3 671B - Triton forward and backward kernels. - -Source: Megatron-LM megatron/core/fusions/fused_mla_yarn_rope_apply.py -Falls back to pure PyTorch when Triton is unavailable. - -Note: DSv3 uses YaRN-scaled RoPE for long-context extrapolation. This test -intentionally uses plain RoPE (base=10000) because it only validates MXFP8 -attention path wiring, tensor shapes, forward/backward flow, and relative BF16 -vs MXFP8 behavior. Both reference and MXFP8 paths use the same RoPE tables. -""" +"""Compat shim: the MLA RoPE kernels moved to +``transformer_engine.pytorch.models.deepseek_v3.mla_rope``.""" import torch -try: - import triton - import triton.language as tl - - HAVE_TRITON = True -except ImportError: - HAVE_TRITON = False +from transformer_engine.pytorch.models.deepseek_v3.mla_rope import ( # noqa: F401 + HAVE_TRITON, + apply_mla_rope_kv, + apply_mla_rope_q, + build_rope_tables, +) HEAD_DIM_ROPE = 64 HEAD_DIM_NOPE = 128 @@ -29,576 +20,6 @@ ROTARY_BASE = 10000 -def build_rope_tables( - seq_len: int, - emb_dim: int = HEAD_DIM_ROPE, - base: int = ROTARY_BASE, - device: torch.device = None, -) -> tuple[torch.Tensor, torch.Tensor]: - inv_freq = 1.0 / ( - base ** (torch.arange(0, emb_dim, 2, dtype=torch.float32, device=device) / emb_dim) - ) - t = torch.arange(seq_len, device=device, dtype=torch.float32) - freqs = torch.outer(t, inv_freq) - freqs = torch.cat([freqs, freqs], dim=-1) - return torch.cos(freqs).contiguous(), torch.sin(freqs).contiguous() - - -if HAVE_TRITON: - - # Not used for non-packed batches; kept for THD compatibility. - @triton.jit - def _get_thd_token_idx(cu_seqlens, pid_m, seq_num, cp_rank, cp_size): - token_idx = -1 - this_seq_len = 0 - seq_idx = 0 - last_cum_seqlen = tl.load(cu_seqlens) // cp_size - while seq_idx < seq_num: - cur_cum_seqlen = tl.load(cu_seqlens + seq_idx + 1) // cp_size - if token_idx == -1 and cur_cum_seqlen > pid_m: - token_idx = pid_m - last_cum_seqlen - this_seq_len = cur_cum_seqlen - last_cum_seqlen - last_cum_seqlen = cur_cum_seqlen - seq_idx += 1 - if cp_size > 1: - if token_idx < this_seq_len // 2: - token_idx = token_idx + cp_rank * this_seq_len // 2 - else: - token_idx = (token_idx - this_seq_len // 2) + ( - 2 * cp_size - cp_rank - 1 - ) * this_seq_len // 2 - return token_idx - - @triton.autotune( - configs=[ - triton.Config({"BLOCK_H": 1}), - triton.Config({"BLOCK_H": 2}), - triton.Config({"BLOCK_H": 4}), - triton.Config({"BLOCK_H": 8}), - triton.Config({"BLOCK_H": 16}), - triton.Config({"BLOCK_H": 32}), - triton.Config({"BLOCK_H": 64}), - triton.Config({"BLOCK_H": 128}), - ], - key=["emb_dim", "head_num"], - restore_value=["Q"], - ) - @triton.jit - def rotary_fwd_q_kernel( - Q, - COS, - SIN, - qk_head_dim, - emb_dim: tl.constexpr, - head_num: tl.constexpr, - batch_size, - seq_num, - cu_seqlens_q, - stride_x_seq, - stride_x_nheads, - cp_rank, - cp_size, - BLOCK_H: tl.constexpr, - ): - pid_m = tl.program_id(axis=0) - pid_head = tl.program_id(axis=1) - if cu_seqlens_q is None: - token_idx = pid_m // batch_size - else: - token_idx = _get_thd_token_idx(cu_seqlens_q, pid_m, seq_num, cp_rank, cp_size) - cos_left = tl.load(COS + token_idx * emb_dim + tl.arange(0, emb_dim // 2)) - sin_left = tl.load(SIN + token_idx * emb_dim + tl.arange(0, emb_dim // 2)) - cos_right = tl.load(COS + token_idx * emb_dim + emb_dim // 2 + tl.arange(0, emb_dim // 2)) - sin_right = tl.load(SIN + token_idx * emb_dim + emb_dim // 2 + tl.arange(0, emb_dim // 2)) - cos_left = cos_left.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) - sin_left = sin_left.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) - cos_right = cos_right.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) - sin_right = sin_right.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) - head_offsets = pid_head * BLOCK_H + tl.arange(0, BLOCK_H) - Q = Q + pid_m * stride_x_seq - x_off = head_offsets[:, None] * stride_x_nheads + qk_head_dim - mask = head_offsets[:, None] < head_num - x_1_off = x_off + tl.arange(0, emb_dim // 2)[None, :] * 2 - x_2_off = x_1_off + 1 - x_1 = tl.load(Q + x_1_off, mask=mask) - x_2 = tl.load(Q + x_2_off, mask=mask) - x_left = x_1 * cos_left - x_2 * sin_left - x_right = x_2 * cos_right + x_1 * sin_right - x_left_off = x_off + tl.arange(0, emb_dim // 2)[None, :] - x_right_off = x_left_off + emb_dim // 2 - tl.store(Q + x_left_off, x_left, mask=mask) - tl.store(Q + x_right_off, x_right, mask=mask) - - @triton.autotune( - configs=[ - triton.Config({"BLOCK_H": 1}), - triton.Config({"BLOCK_H": 2}), - triton.Config({"BLOCK_H": 4}), - triton.Config({"BLOCK_H": 8}), - triton.Config({"BLOCK_H": 16}), - triton.Config({"BLOCK_H": 32}), - triton.Config({"BLOCK_H": 64}), - triton.Config({"BLOCK_H": 128}), - ], - key=["emb_dim", "head_num"], - restore_value=["DO"], - ) - @triton.jit - def rotary_bwd_q_kernel( - DO, - COS, - SIN, - qk_head_dim, - emb_dim: tl.constexpr, - head_num: tl.constexpr, - batch_size, - seq_num, - cu_seqlens_q, - stride_x_seq, - stride_x_nheads, - cp_rank, - cp_size, - BLOCK_H: tl.constexpr, - ): - pid_m = tl.program_id(axis=0) - pid_head = tl.program_id(axis=1) - if cu_seqlens_q is None: - token_idx = pid_m // batch_size - else: - token_idx = _get_thd_token_idx(cu_seqlens_q, pid_m, seq_num, cp_rank, cp_size) - cos_left = tl.load(COS + token_idx * emb_dim + tl.arange(0, emb_dim // 2)) - sin_left = tl.load(SIN + token_idx * emb_dim + tl.arange(0, emb_dim // 2)) - cos_right = tl.load(COS + token_idx * emb_dim + emb_dim // 2 + tl.arange(0, emb_dim // 2)) - sin_right = tl.load(SIN + token_idx * emb_dim + emb_dim // 2 + tl.arange(0, emb_dim // 2)) - cos_left = cos_left.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) - sin_left = sin_left.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) - cos_right = cos_right.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) - sin_right = sin_right.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) - head_offsets = pid_head * BLOCK_H + tl.arange(0, BLOCK_H) - DO = DO + pid_m * stride_x_seq - x_off = head_offsets[:, None] * stride_x_nheads + qk_head_dim - mask = head_offsets[:, None] < head_num - x_left_off = x_off + tl.arange(0, emb_dim // 2)[None, :] - x_right_off = x_left_off + emb_dim // 2 - x_left = tl.load(DO + x_left_off, mask=mask) - x_right = tl.load(DO + x_right_off, mask=mask) - x_1 = x_left * cos_left + x_right * sin_right - x_2 = -x_left * sin_left + x_right * cos_right - x_1_off = x_off + tl.arange(0, emb_dim // 2)[None, :] * 2 - x_2_off = x_1_off + 1 - tl.store(DO + x_1_off, x_1, mask=mask) - tl.store(DO + x_2_off, x_2, mask=mask) - - @triton.autotune( - configs=[ - triton.Config({"BLOCK_H": 1}), - triton.Config({"BLOCK_H": 2}), - triton.Config({"BLOCK_H": 4}), - triton.Config({"BLOCK_H": 8}), - triton.Config({"BLOCK_H": 16}), - triton.Config({"BLOCK_H": 32}), - triton.Config({"BLOCK_H": 64}), - triton.Config({"BLOCK_H": 128}), - ], - key=["emb_dim", "k_dim", "v_dim", "head_num"], - ) - @triton.jit - def rotary_fwd_kv_kernel( - KV, - K_POS_EMB, - O_KEY, - O_VALUE, - COS, - SIN, - emb_dim: tl.constexpr, - k_dim: tl.constexpr, - v_dim: tl.constexpr, - head_num: tl.constexpr, - batch_size, - seq_num, - cu_seqlens_kv, - stride_kv_seq, - stride_kv_nheads, - stride_emb_seq, - stride_k_seq, - stride_k_nheads, - stride_v_seq, - stride_v_nheads, - cp_rank, - cp_size, - BLOCK_H: tl.constexpr, - ): - pid_m = tl.program_id(axis=0) - pid_head = tl.program_id(axis=1) - if cu_seqlens_kv is None: - token_idx = pid_m // batch_size - else: - token_idx = _get_thd_token_idx(cu_seqlens_kv, pid_m, seq_num, cp_rank, cp_size) - cos_left = tl.load(COS + token_idx * emb_dim + tl.arange(0, emb_dim // 2)) - sin_left = tl.load(SIN + token_idx * emb_dim + tl.arange(0, emb_dim // 2)) - cos_right = tl.load(COS + token_idx * emb_dim + emb_dim // 2 + tl.arange(0, emb_dim // 2)) - sin_right = tl.load(SIN + token_idx * emb_dim + emb_dim // 2 + tl.arange(0, emb_dim // 2)) - head_offsets = pid_head * BLOCK_H + tl.arange(0, BLOCK_H) - KV_ptr = KV + pid_m * stride_kv_seq - kv_off = head_offsets[:, None] * stride_kv_nheads - mask = head_offsets[:, None] < head_num - k_in_off = kv_off + tl.arange(0, k_dim)[None, :] - v_in_off = kv_off + k_dim + tl.arange(0, v_dim)[None, :] - k = tl.load(KV_ptr + k_in_off, mask=mask) - v = tl.load(KV_ptr + v_in_off, mask=mask) - K_ptr = O_KEY + pid_m * stride_k_seq + pid_head * BLOCK_H * stride_k_nheads - V_ptr = O_VALUE + pid_m * stride_v_seq + pid_head * BLOCK_H * stride_v_nheads - k_out_off = tl.arange(0, BLOCK_H)[:, None] * stride_k_nheads + tl.arange(0, k_dim)[None, :] - v_out_off = tl.arange(0, BLOCK_H)[:, None] * stride_v_nheads + tl.arange(0, v_dim)[None, :] - tl.store(K_ptr + k_out_off, k, mask=mask) - tl.store(V_ptr + v_out_off, v, mask=mask) - EMB = K_POS_EMB + pid_m * stride_emb_seq - x_1 = tl.load(EMB + tl.arange(0, emb_dim // 2) * 2) - x_2 = tl.load(EMB + tl.arange(0, emb_dim // 2) * 2 + 1) - x_left = x_1 * cos_left - x_2 * sin_left - x_right = x_2 * cos_right + x_1 * sin_right - x_left = x_left.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) - x_right = x_right.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) - x_left_off = ( - tl.arange(0, BLOCK_H)[:, None] * stride_k_nheads - + k_dim - + tl.arange(0, emb_dim // 2)[None, :] - ) - x_right_off = x_left_off + emb_dim // 2 - tl.store(K_ptr + x_left_off, x_left, mask=mask) - tl.store(K_ptr + x_right_off, x_right, mask=mask) - - @triton.autotune( - configs=[ - triton.Config({"BLOCK_H": 1}), - triton.Config({"BLOCK_H": 2}), - triton.Config({"BLOCK_H": 4}), - triton.Config({"BLOCK_H": 8}), - triton.Config({"BLOCK_H": 16}), - triton.Config({"BLOCK_H": 32}), - triton.Config({"BLOCK_H": 64}), - triton.Config({"BLOCK_H": 128}), - ], - key=["emb_dim", "k_dim", "v_dim", "head_num"], - ) - @triton.jit - def rotary_bwd_kv_kernel( - dK, - dV, - dKV, - dEMB, - COS, - SIN, - emb_dim: tl.constexpr, - k_dim: tl.constexpr, - v_dim: tl.constexpr, - head_num: tl.constexpr, - batch_size, - seq_num, - cu_seqlens_kv, - stride_dk_seq, - stride_dk_nheads, - stride_dv_seq, - stride_dv_nheads, - stride_dkv_seq, - stride_dkv_nheads, - stride_demb_seq, - cp_rank, - cp_size, - BLOCK_H: tl.constexpr, - ): - pid_m = tl.program_id(axis=0) - pid_head = tl.program_id(axis=1) - if cu_seqlens_kv is None: - token_idx = pid_m // batch_size - else: - token_idx = _get_thd_token_idx(cu_seqlens_kv, pid_m, seq_num, cp_rank, cp_size) - head_offsets = pid_head * BLOCK_H + tl.arange(0, BLOCK_H) - dKV_ptr = dKV + pid_m * stride_dkv_seq - dkv_off = head_offsets[:, None] * stride_dkv_nheads - mask = head_offsets[:, None] < head_num - dk_out_off = dkv_off + tl.arange(0, k_dim)[None, :] - dv_out_off = dkv_off + k_dim + tl.arange(0, v_dim)[None, :] - dK_ptr = dK + pid_m * stride_dk_seq + pid_head * BLOCK_H * stride_dk_nheads - dV_ptr = dV + pid_m * stride_dv_seq + pid_head * BLOCK_H * stride_dv_nheads - dk_in_off = tl.arange(0, BLOCK_H)[:, None] * stride_dk_nheads + tl.arange(0, k_dim)[None, :] - dv_in_off = tl.arange(0, BLOCK_H)[:, None] * stride_dv_nheads + tl.arange(0, v_dim)[None, :] - dk = tl.load(dK_ptr + dk_in_off, mask=mask) - dv = tl.load(dV_ptr + dv_in_off, mask=mask) - tl.store(dKV_ptr + dk_out_off, dk, mask=mask) - tl.store(dKV_ptr + dv_out_off, dv, mask=mask) - if pid_head == 0: - x_left_accum = tl.zeros((BLOCK_H, emb_dim // 2), dtype=tl.float32) - x_right_accum = tl.zeros((BLOCK_H, emb_dim // 2), dtype=tl.float32) - for i in tl.static_range(triton.cdiv(head_num, BLOCK_H)): - head_offsets_i = i * BLOCK_H + tl.arange(0, BLOCK_H) - dK_ptr_i = dK + pid_m * stride_dk_seq - x_off = head_offsets_i[:, None] * stride_dk_nheads + k_dim - mask_i = head_offsets_i[:, None] < head_num - x_left_off = x_off + tl.arange(0, emb_dim // 2)[None, :] - x_right_off = x_left_off + emb_dim // 2 - x_left_accum += tl.load(dK_ptr_i + x_left_off, mask=mask_i) - x_right_accum += tl.load(dK_ptr_i + x_right_off, mask=mask_i) - x_left_accum = tl.sum(x_left_accum, axis=0) - x_right_accum = tl.sum(x_right_accum, axis=0) - x_left_accum = x_left_accum.to(dEMB.dtype.element_ty) - x_right_accum = x_right_accum.to(dEMB.dtype.element_ty) - cos_left = tl.load(COS + token_idx * emb_dim + tl.arange(0, emb_dim // 2)) - sin_left = tl.load(SIN + token_idx * emb_dim + tl.arange(0, emb_dim // 2)) - cos_right = tl.load( - COS + token_idx * emb_dim + emb_dim // 2 + tl.arange(0, emb_dim // 2) - ) - sin_right = tl.load( - SIN + token_idx * emb_dim + emb_dim // 2 + tl.arange(0, emb_dim // 2) - ) - x_1 = x_left_accum * cos_left + x_right_accum * sin_right - x_2 = -x_left_accum * sin_left + x_right_accum * cos_right - dEMB_ptr = dEMB + pid_m * stride_demb_seq - tl.store(dEMB_ptr + tl.arange(0, emb_dim // 2) * 2, x_1) - tl.store(dEMB_ptr + tl.arange(0, emb_dim // 2) * 2 + 1, x_2) - - def _flattened_token_stride(tensor: torch.Tensor) -> int: - if tensor.dim() == 4: - return tensor.stride(1) - return tensor.stride(0) - - class _MLARoPEQTriton(torch.autograd.Function): - @staticmethod - def forward(ctx, q, cos, sin, head_dim_nope, head_dim_rope): - s, b, nheads, _ = q.shape - total = s * b - - grid_q = lambda META: (total, triton.cdiv(nheads, META["BLOCK_H"])) - rotary_fwd_q_kernel[grid_q]( - q, - cos, - sin, - head_dim_nope, - head_dim_rope, - nheads, - b, - None, - None, - _flattened_token_stride(q), - q.stride(2), - 0, - 1, - ) - - ctx.save_for_backward(cos, sin) - ctx.head_dim_nope = head_dim_nope - ctx.head_dim_rope = head_dim_rope - ctx.nheads = nheads - ctx.s = s - ctx.b = b - return q - - @staticmethod - def backward(ctx, dq): - cos, sin = ctx.saved_tensors - s, b, nheads = ctx.s, ctx.b, ctx.nheads - total = s * b - - grid_q = lambda META: (total, triton.cdiv(nheads, META["BLOCK_H"])) - rotary_bwd_q_kernel[grid_q]( - dq, - cos, - sin, - ctx.head_dim_nope, - ctx.head_dim_rope, - nheads, - b, - None, - None, - _flattened_token_stride(dq), - dq.stride(2), - 0, - 1, - ) - return dq, None, None, None, None - - class _MLARoPEKVTriton(torch.autograd.Function): - @staticmethod - def forward(ctx, kv, k_pos_emb, cos, sin, head_dim_nope, head_dim_rope, head_dim_v): - s, b, nheads, _ = kv.shape - total = s * b - - o_key = kv.new_empty(s, b, nheads, head_dim_nope + head_dim_rope) - o_value = kv.new_empty(s, b, nheads, head_dim_v) - grid_kv = lambda META: (total, triton.cdiv(nheads, META["BLOCK_H"])) - rotary_fwd_kv_kernel[grid_kv]( - kv, - k_pos_emb, - o_key, - o_value, - cos, - sin, - head_dim_rope, - head_dim_nope, - head_dim_v, - nheads, - b, - None, - None, - _flattened_token_stride(kv), - kv.stride(2), - _flattened_token_stride(k_pos_emb), - _flattened_token_stride(o_key), - o_key.stride(2), - _flattened_token_stride(o_value), - o_value.stride(2), - 0, - 1, - ) - - ctx.save_for_backward(cos, sin) - ctx.head_dim_nope = head_dim_nope - ctx.head_dim_rope = head_dim_rope - ctx.head_dim_v = head_dim_v - ctx.nheads = nheads - ctx.s = s - ctx.b = b - return o_key, o_value - - @staticmethod - def backward(ctx, dk_out, dv_out): - cos, sin = ctx.saved_tensors - s, b, nheads = ctx.s, ctx.b, ctx.nheads - ndp, ndr, ndv = ctx.head_dim_nope, ctx.head_dim_rope, ctx.head_dim_v - total = s * b - - d_kv = dk_out.new_empty(s, b, nheads, ndp + ndv) - d_emb = dk_out.new_empty(s, b, 1, ndr) - grid_kv = lambda META: (total, triton.cdiv(nheads, META["BLOCK_H"])) - rotary_bwd_kv_kernel[grid_kv]( - dk_out, - dv_out, - d_kv, - d_emb, - cos, - sin, - ndr, - ndp, - ndv, - nheads, - b, - None, - None, - _flattened_token_stride(dk_out), - dk_out.stride(2), - _flattened_token_stride(dv_out), - dv_out.stride(2), - _flattened_token_stride(d_kv), - d_kv.stride(2), - _flattened_token_stride(d_emb), - 0, - 1, - ) - return d_kv, d_emb, None, None, None, None, None - - -def _apply_mla_rope_q_with_tables( - q: torch.Tensor, - cos_table: torch.Tensor, - sin_table: torch.Tensor, - head_dim_nope: int = HEAD_DIM_NOPE, - head_dim_rope: int = HEAD_DIM_ROPE, -) -> torch.Tensor: - if HAVE_TRITON: - return _MLARoPEQTriton.apply( - q, - cos_table, - sin_table, - head_dim_nope, - head_dim_rope, - ) - return _apply_pytorch_q(q, cos_table, sin_table, head_dim_nope, head_dim_rope) - - -def _apply_mla_rope_kv_with_tables( - kv: torch.Tensor, - k_pos_emb: torch.Tensor, - cos_table: torch.Tensor, - sin_table: torch.Tensor, - head_dim_nope: int = HEAD_DIM_NOPE, - head_dim_rope: int = HEAD_DIM_ROPE, - head_dim_v: int = HEAD_DIM_V, -) -> tuple[torch.Tensor, torch.Tensor]: - if HAVE_TRITON: - return _MLARoPEKVTriton.apply( - kv, - k_pos_emb, - cos_table, - sin_table, - head_dim_nope, - head_dim_rope, - head_dim_v, - ) - return _apply_pytorch_kv( - kv, - k_pos_emb, - cos_table, - sin_table, - head_dim_nope, - head_dim_rope, - head_dim_v, - ) - - -def apply_mla_rope_q( - q: torch.Tensor, - head_dim_nope: int = HEAD_DIM_NOPE, - head_dim_rope: int = HEAD_DIM_ROPE, - base: int = ROTARY_BASE, - cos_table: torch.Tensor | None = None, - sin_table: torch.Tensor | None = None, -) -> torch.Tensor: - if cos_table is None or sin_table is None: - s = q.shape[0] - cos_table, sin_table = build_rope_tables( - s, - emb_dim=head_dim_rope, - base=base, - device=q.device, - ) - return _apply_mla_rope_q_with_tables( - q, - cos_table, - sin_table, - head_dim_nope, - head_dim_rope, - ) - - -def apply_mla_rope_kv( - kv: torch.Tensor, - k_pos_emb: torch.Tensor, - head_dim_nope: int = HEAD_DIM_NOPE, - head_dim_rope: int = HEAD_DIM_ROPE, - head_dim_v: int = HEAD_DIM_V, - base: int = ROTARY_BASE, - cos_table: torch.Tensor | None = None, - sin_table: torch.Tensor | None = None, -) -> tuple[torch.Tensor, torch.Tensor]: - if cos_table is None or sin_table is None: - s = kv.shape[0] - cos_table, sin_table = build_rope_tables( - s, - emb_dim=head_dim_rope, - base=base, - device=kv.device, - ) - return _apply_mla_rope_kv_with_tables( - kv, - k_pos_emb, - cos_table, - sin_table, - head_dim_nope, - head_dim_rope, - head_dim_v, - ) - - def apply_mla_rope( q: torch.Tensor, kv: torch.Tensor, @@ -609,60 +30,13 @@ def apply_mla_rope( base: int = ROTARY_BASE, cos_table: torch.Tensor | None = None, sin_table: torch.Tensor | None = None, -) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor]: +): if cos_table is None or sin_table is None: - s = q.shape[0] cos_table, sin_table = build_rope_tables( - s, - emb_dim=head_dim_rope, - base=base, - device=q.device, + q.shape[0], head_dim_rope, base=base, device=q.device ) - q = _apply_mla_rope_q_with_tables(q, cos_table, sin_table, head_dim_nope, head_dim_rope) - k, v = _apply_mla_rope_kv_with_tables( - kv, - k_pos_emb, - cos_table, - sin_table, - head_dim_nope, - head_dim_rope, - head_dim_v, + q = apply_mla_rope_q(q, cos_table, sin_table, head_dim_nope, head_dim_rope) + k, v = apply_mla_rope_kv( + kv, k_pos_emb, cos_table, sin_table, head_dim_nope, head_dim_rope, head_dim_v ) return q, k, v - - -def _rotate_interleaved_to_neox( - x: torch.Tensor, cos_table: torch.Tensor, sin_table: torch.Tensor -) -> torch.Tensor: - cos_ = cos_table[:, None, None, :].to(x.dtype) - sin_ = sin_table[:, None, None, :].to(x.dtype) - half_dim = x.shape[-1] // 2 - x_1 = x[..., 0::2] - x_2 = x[..., 1::2] - x_left = x_1 * cos_[..., :half_dim] - x_2 * sin_[..., :half_dim] - x_right = x_2 * cos_[..., half_dim:] + x_1 * sin_[..., half_dim:] - return torch.cat((x_left, x_right), dim=-1) - - -def _apply_pytorch_q(q, cos_table, sin_table, head_dim_nope, head_dim_rope): - q_nope = q[..., :head_dim_nope] - q_rope = q[..., head_dim_nope : head_dim_nope + head_dim_rope] - q_rope = _rotate_interleaved_to_neox(q_rope, cos_table, sin_table) - return torch.cat((q_nope, q_rope), dim=-1) - - -def _apply_pytorch_kv( - kv, - k_pos_emb, - cos_table, - sin_table, - head_dim_nope, - head_dim_rope, - head_dim_v, -): - k_nope = kv[..., :head_dim_nope] - v = kv[..., head_dim_nope : head_dim_nope + head_dim_v] - k_rope = _rotate_interleaved_to_neox(k_pos_emb, cos_table, sin_table).expand( - -1, -1, kv.shape[2], -1 - ) - return torch.cat((k_nope, k_rope), dim=-1), v diff --git a/tests/pytorch/test_deepseek.py b/tests/pytorch/test_deepseek.py index 7778d0448c..4c4aea0a92 100644 --- a/tests/pytorch/test_deepseek.py +++ b/tests/pytorch/test_deepseek.py @@ -34,6 +34,66 @@ def _input(requires_grad=True): ) +def test_mla_rope_triton_matches_pytorch(): + from transformer_engine.pytorch.models.deepseek_v3 import mla_rope + + if not mla_rope.HAVE_TRITON: + pytest.skip("Triton unavailable") + s, b, h = 64, 2, 4 + nope, rope, vdim = 64, 32, 64 + cos, sin = mla_rope.build_rope_tables(s, rope, device="cuda") + + torch.manual_seed(0) + q_leaf = torch.randn(s, b, h, nope + rope, device="cuda", requires_grad=True) + kv_leaf = torch.randn(s, b, h, nope + vdim, device="cuda", requires_grad=True) + pos_leaf = torch.randn(s, b, 1, rope, device="cuda", requires_grad=True) + grad_q = torch.randn(s, b, h, nope + rope, device="cuda") + grad_k = torch.randn(s, b, h, nope + rope, device="cuda") + grad_v = torch.randn(s, b, h, vdim, device="cuda") + + def run(fmt): + # non-leaf copies: the Triton q kernel rotates in place + q, kv, pos = q_leaf * 1.0, kv_leaf * 1.0, pos_leaf * 1.0 + q_out = mla_rope.apply_mla_rope_q(q, cos, sin, nope, rope, fmt) + k_out, v_out = mla_rope.apply_mla_rope_kv(kv, pos, cos, sin, nope, rope, vdim, fmt) + # fresh grad clones: the Triton q backward modifies its input grad in place + torch.autograd.backward( + [q_out, k_out, v_out], [grad_q.clone(), grad_k.clone(), grad_v.clone()] + ) + grads = (q_leaf.grad.clone(), kv_leaf.grad.clone(), pos_leaf.grad.clone()) + q_leaf.grad = kv_leaf.grad = pos_leaf.grad = None + return (q_out.clone(), k_out, v_out), grads + + (q_t, k_t, v_t), grads_t = run("sbhd") + + seq_dim = 0 + q_ref = torch.cat( + ( + (q_leaf * 1.0)[..., :nope], + mla_rope._rotate_interleaved_to_neox((q_leaf * 1.0)[..., nope:], cos, sin, seq_dim), + ), + dim=-1, + ) + k_ref = torch.cat( + ( + (kv_leaf * 1.0)[..., :nope], + mla_rope._rotate_interleaved_to_neox(pos_leaf * 1.0, cos, sin, seq_dim).expand( + s, b, h, rope + ), + ), + dim=-1, + ) + v_ref = (kv_leaf * 1.0)[..., nope:] + torch.autograd.backward([q_ref, k_ref, v_ref], [grad_q.clone(), grad_k.clone(), grad_v.clone()]) + + torch.testing.assert_close(q_t, q_ref, rtol=1e-5, atol=1e-5) + torch.testing.assert_close(k_t, k_ref, rtol=1e-5, atol=1e-5) + torch.testing.assert_close(v_t, v_ref, rtol=1e-5, atol=1e-5) + torch.testing.assert_close(grads_t[0], q_leaf.grad, rtol=1e-5, atol=1e-5) + torch.testing.assert_close(grads_t[1], kv_leaf.grad, rtol=1e-5, atol=1e-5) + torch.testing.assert_close(grads_t[2], pos_leaf.grad, rtol=1e-5, atol=1e-5) + + def test_mla_forward_backward(): torch.manual_seed(0) mla = MultiLatentAttention(HIDDEN, HEADS, params_dtype=DTYPE, **MLA_KWARGS) diff --git a/transformer_engine/pytorch/models/deepseek_v3/mla_rope.py b/transformer_engine/pytorch/models/deepseek_v3/mla_rope.py new file mode 100644 index 0000000000..350bedb69b --- /dev/null +++ b/transformer_engine/pytorch/models/deepseek_v3/mla_rope.py @@ -0,0 +1,495 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Fused MLA RoPE kernels (DeepSeekV3-style decoupled RoPE/NoPE). + +Triton forward/backward kernels adapted from Megatron-LM +``megatron/core/fusions/fused_mla_yarn_rope_apply.py``. The query kernel +rotates the trailing ``head_dim_rope`` slice in place (no concat); the KV +kernel builds the final key (nope | broadcast-rotated shared rope head) and +value tensors in a single pass. Falls back to pure PyTorch when Triton is +unavailable or for the ``bshd`` layout (the Triton path is ``sbhd``-only). + +Rotation convention: the rope slice is read interleaved (as stored in +HF/Megatron DeepSeekV3 checkpoints) and written in NeoX half-split layout, +matching the Megatron fused kernel semantics. +""" + +from typing import Optional, Tuple + +import torch + +try: + import triton + import triton.language as tl + + HAVE_TRITON = True +except ImportError: + HAVE_TRITON = False + +__all__ = ["build_rope_tables", "apply_mla_rope_q", "apply_mla_rope_kv"] + + +def build_rope_tables( + seq_len: int, + emb_dim: int, + base: float = 10000.0, + device: Optional[torch.device] = None, +) -> Tuple[torch.Tensor, torch.Tensor]: + """cos/sin tables of shape ``[seq_len, emb_dim]`` (fp32, NeoX duplicated halves).""" + inv_freq = 1.0 / ( + base ** (torch.arange(0, emb_dim, 2, dtype=torch.float32, device=device) / emb_dim) + ) + t = torch.arange(seq_len, device=device, dtype=torch.float32) + freqs = torch.outer(t, inv_freq) + freqs = torch.cat([freqs, freqs], dim=-1) + return torch.cos(freqs).contiguous(), torch.sin(freqs).contiguous() + + +if HAVE_TRITON: + + # Not used for non-packed batches; kept for THD compatibility. + @triton.jit + def _get_thd_token_idx(cu_seqlens, pid_m, seq_num, cp_rank, cp_size): + token_idx = -1 + this_seq_len = 0 + seq_idx = 0 + last_cum_seqlen = tl.load(cu_seqlens) // cp_size + while seq_idx < seq_num: + cur_cum_seqlen = tl.load(cu_seqlens + seq_idx + 1) // cp_size + if token_idx == -1 and cur_cum_seqlen > pid_m: + token_idx = pid_m - last_cum_seqlen + this_seq_len = cur_cum_seqlen - last_cum_seqlen + last_cum_seqlen = cur_cum_seqlen + seq_idx += 1 + if cp_size > 1: + if token_idx < this_seq_len // 2: + token_idx = token_idx + cp_rank * this_seq_len // 2 + else: + token_idx = (token_idx - this_seq_len // 2) + ( + 2 * cp_size - cp_rank - 1 + ) * this_seq_len // 2 + return token_idx + + _AUTOTUNE_CONFIGS = [triton.Config({"BLOCK_H": h}) for h in (1, 2, 4, 8, 16, 32, 64, 128)] + + @triton.autotune(configs=_AUTOTUNE_CONFIGS, key=["emb_dim", "head_num"], restore_value=["Q"]) + @triton.jit + def rotary_fwd_q_kernel( + Q, + COS, + SIN, + qk_head_dim, + emb_dim: tl.constexpr, + head_num: tl.constexpr, + batch_size, + seq_num, + cu_seqlens_q, + stride_x_seq, + stride_x_nheads, + cp_rank, + cp_size, + BLOCK_H: tl.constexpr, + ): + pid_m = tl.program_id(axis=0) + pid_head = tl.program_id(axis=1) + if cu_seqlens_q is None: + token_idx = pid_m // batch_size + else: + token_idx = _get_thd_token_idx(cu_seqlens_q, pid_m, seq_num, cp_rank, cp_size) + cos_left = tl.load(COS + token_idx * emb_dim + tl.arange(0, emb_dim // 2)) + sin_left = tl.load(SIN + token_idx * emb_dim + tl.arange(0, emb_dim // 2)) + cos_right = tl.load(COS + token_idx * emb_dim + emb_dim // 2 + tl.arange(0, emb_dim // 2)) + sin_right = tl.load(SIN + token_idx * emb_dim + emb_dim // 2 + tl.arange(0, emb_dim // 2)) + cos_left = cos_left.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) + sin_left = sin_left.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) + cos_right = cos_right.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) + sin_right = sin_right.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) + head_offsets = pid_head * BLOCK_H + tl.arange(0, BLOCK_H) + Q = Q + pid_m * stride_x_seq + x_off = head_offsets[:, None] * stride_x_nheads + qk_head_dim + mask = head_offsets[:, None] < head_num + x_1_off = x_off + tl.arange(0, emb_dim // 2)[None, :] * 2 + x_2_off = x_1_off + 1 + x_1 = tl.load(Q + x_1_off, mask=mask) + x_2 = tl.load(Q + x_2_off, mask=mask) + x_left = x_1 * cos_left - x_2 * sin_left + x_right = x_2 * cos_right + x_1 * sin_right + x_left_off = x_off + tl.arange(0, emb_dim // 2)[None, :] + x_right_off = x_left_off + emb_dim // 2 + tl.store(Q + x_left_off, x_left, mask=mask) + tl.store(Q + x_right_off, x_right, mask=mask) + + @triton.autotune(configs=_AUTOTUNE_CONFIGS, key=["emb_dim", "head_num"], restore_value=["DO"]) + @triton.jit + def rotary_bwd_q_kernel( + DO, + COS, + SIN, + qk_head_dim, + emb_dim: tl.constexpr, + head_num: tl.constexpr, + batch_size, + seq_num, + cu_seqlens_q, + stride_x_seq, + stride_x_nheads, + cp_rank, + cp_size, + BLOCK_H: tl.constexpr, + ): + pid_m = tl.program_id(axis=0) + pid_head = tl.program_id(axis=1) + if cu_seqlens_q is None: + token_idx = pid_m // batch_size + else: + token_idx = _get_thd_token_idx(cu_seqlens_q, pid_m, seq_num, cp_rank, cp_size) + cos_left = tl.load(COS + token_idx * emb_dim + tl.arange(0, emb_dim // 2)) + sin_left = tl.load(SIN + token_idx * emb_dim + tl.arange(0, emb_dim // 2)) + cos_right = tl.load(COS + token_idx * emb_dim + emb_dim // 2 + tl.arange(0, emb_dim // 2)) + sin_right = tl.load(SIN + token_idx * emb_dim + emb_dim // 2 + tl.arange(0, emb_dim // 2)) + cos_left = cos_left.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) + sin_left = sin_left.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) + cos_right = cos_right.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) + sin_right = sin_right.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) + head_offsets = pid_head * BLOCK_H + tl.arange(0, BLOCK_H) + DO = DO + pid_m * stride_x_seq + x_off = head_offsets[:, None] * stride_x_nheads + qk_head_dim + mask = head_offsets[:, None] < head_num + x_left_off = x_off + tl.arange(0, emb_dim // 2)[None, :] + x_right_off = x_left_off + emb_dim // 2 + x_left = tl.load(DO + x_left_off, mask=mask) + x_right = tl.load(DO + x_right_off, mask=mask) + x_1 = x_left * cos_left + x_right * sin_right + x_2 = -x_left * sin_left + x_right * cos_right + x_1_off = x_off + tl.arange(0, emb_dim // 2)[None, :] * 2 + x_2_off = x_1_off + 1 + tl.store(DO + x_1_off, x_1, mask=mask) + tl.store(DO + x_2_off, x_2, mask=mask) + + @triton.autotune(configs=_AUTOTUNE_CONFIGS, key=["emb_dim", "k_dim", "v_dim", "head_num"]) + @triton.jit + def rotary_fwd_kv_kernel( + KV, + K_POS_EMB, + O_KEY, + O_VALUE, + COS, + SIN, + emb_dim: tl.constexpr, + k_dim: tl.constexpr, + v_dim: tl.constexpr, + head_num: tl.constexpr, + batch_size, + seq_num, + cu_seqlens_kv, + stride_kv_seq, + stride_kv_nheads, + stride_emb_seq, + stride_k_seq, + stride_k_nheads, + stride_v_seq, + stride_v_nheads, + cp_rank, + cp_size, + BLOCK_H: tl.constexpr, + ): + pid_m = tl.program_id(axis=0) + pid_head = tl.program_id(axis=1) + if cu_seqlens_kv is None: + token_idx = pid_m // batch_size + else: + token_idx = _get_thd_token_idx(cu_seqlens_kv, pid_m, seq_num, cp_rank, cp_size) + cos_left = tl.load(COS + token_idx * emb_dim + tl.arange(0, emb_dim // 2)) + sin_left = tl.load(SIN + token_idx * emb_dim + tl.arange(0, emb_dim // 2)) + cos_right = tl.load(COS + token_idx * emb_dim + emb_dim // 2 + tl.arange(0, emb_dim // 2)) + sin_right = tl.load(SIN + token_idx * emb_dim + emb_dim // 2 + tl.arange(0, emb_dim // 2)) + head_offsets = pid_head * BLOCK_H + tl.arange(0, BLOCK_H) + KV_ptr = KV + pid_m * stride_kv_seq + kv_off = head_offsets[:, None] * stride_kv_nheads + mask = head_offsets[:, None] < head_num + k_in_off = kv_off + tl.arange(0, k_dim)[None, :] + v_in_off = kv_off + k_dim + tl.arange(0, v_dim)[None, :] + k = tl.load(KV_ptr + k_in_off, mask=mask) + v = tl.load(KV_ptr + v_in_off, mask=mask) + K_ptr = O_KEY + pid_m * stride_k_seq + pid_head * BLOCK_H * stride_k_nheads + V_ptr = O_VALUE + pid_m * stride_v_seq + pid_head * BLOCK_H * stride_v_nheads + k_out_off = tl.arange(0, BLOCK_H)[:, None] * stride_k_nheads + tl.arange(0, k_dim)[None, :] + v_out_off = tl.arange(0, BLOCK_H)[:, None] * stride_v_nheads + tl.arange(0, v_dim)[None, :] + tl.store(K_ptr + k_out_off, k, mask=mask) + tl.store(V_ptr + v_out_off, v, mask=mask) + EMB = K_POS_EMB + pid_m * stride_emb_seq + x_1 = tl.load(EMB + tl.arange(0, emb_dim // 2) * 2) + x_2 = tl.load(EMB + tl.arange(0, emb_dim // 2) * 2 + 1) + x_left = x_1 * cos_left - x_2 * sin_left + x_right = x_2 * cos_right + x_1 * sin_right + x_left = x_left.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) + x_right = x_right.expand_dims(0).broadcast_to(BLOCK_H, emb_dim // 2) + x_left_off = ( + tl.arange(0, BLOCK_H)[:, None] * stride_k_nheads + + k_dim + + tl.arange(0, emb_dim // 2)[None, :] + ) + x_right_off = x_left_off + emb_dim // 2 + tl.store(K_ptr + x_left_off, x_left, mask=mask) + tl.store(K_ptr + x_right_off, x_right, mask=mask) + + @triton.autotune(configs=_AUTOTUNE_CONFIGS, key=["emb_dim", "k_dim", "v_dim", "head_num"]) + @triton.jit + def rotary_bwd_kv_kernel( + dK, + dV, + dKV, + dEMB, + COS, + SIN, + emb_dim: tl.constexpr, + k_dim: tl.constexpr, + v_dim: tl.constexpr, + head_num: tl.constexpr, + batch_size, + seq_num, + cu_seqlens_kv, + stride_dk_seq, + stride_dk_nheads, + stride_dv_seq, + stride_dv_nheads, + stride_dkv_seq, + stride_dkv_nheads, + stride_demb_seq, + cp_rank, + cp_size, + BLOCK_H: tl.constexpr, + ): + pid_m = tl.program_id(axis=0) + pid_head = tl.program_id(axis=1) + if cu_seqlens_kv is None: + token_idx = pid_m // batch_size + else: + token_idx = _get_thd_token_idx(cu_seqlens_kv, pid_m, seq_num, cp_rank, cp_size) + head_offsets = pid_head * BLOCK_H + tl.arange(0, BLOCK_H) + dKV_ptr = dKV + pid_m * stride_dkv_seq + dkv_off = head_offsets[:, None] * stride_dkv_nheads + mask = head_offsets[:, None] < head_num + dk_out_off = dkv_off + tl.arange(0, k_dim)[None, :] + dv_out_off = dkv_off + k_dim + tl.arange(0, v_dim)[None, :] + dK_ptr = dK + pid_m * stride_dk_seq + pid_head * BLOCK_H * stride_dk_nheads + dV_ptr = dV + pid_m * stride_dv_seq + pid_head * BLOCK_H * stride_dv_nheads + dk_in_off = tl.arange(0, BLOCK_H)[:, None] * stride_dk_nheads + tl.arange(0, k_dim)[None, :] + dv_in_off = tl.arange(0, BLOCK_H)[:, None] * stride_dv_nheads + tl.arange(0, v_dim)[None, :] + dk = tl.load(dK_ptr + dk_in_off, mask=mask) + dv = tl.load(dV_ptr + dv_in_off, mask=mask) + tl.store(dKV_ptr + dk_out_off, dk, mask=mask) + tl.store(dKV_ptr + dv_out_off, dv, mask=mask) + if pid_head == 0: + x_left_accum = tl.zeros((BLOCK_H, emb_dim // 2), dtype=tl.float32) + x_right_accum = tl.zeros((BLOCK_H, emb_dim // 2), dtype=tl.float32) + for i in tl.static_range(triton.cdiv(head_num, BLOCK_H)): + head_offsets_i = i * BLOCK_H + tl.arange(0, BLOCK_H) + dK_ptr_i = dK + pid_m * stride_dk_seq + x_off = head_offsets_i[:, None] * stride_dk_nheads + k_dim + mask_i = head_offsets_i[:, None] < head_num + x_left_off = x_off + tl.arange(0, emb_dim // 2)[None, :] + x_right_off = x_left_off + emb_dim // 2 + x_left_accum += tl.load(dK_ptr_i + x_left_off, mask=mask_i) + x_right_accum += tl.load(dK_ptr_i + x_right_off, mask=mask_i) + x_left_accum = tl.sum(x_left_accum, axis=0) + x_right_accum = tl.sum(x_right_accum, axis=0) + x_left_accum = x_left_accum.to(dEMB.dtype.element_ty) + x_right_accum = x_right_accum.to(dEMB.dtype.element_ty) + cos_left = tl.load(COS + token_idx * emb_dim + tl.arange(0, emb_dim // 2)) + sin_left = tl.load(SIN + token_idx * emb_dim + tl.arange(0, emb_dim // 2)) + cos_right = tl.load( + COS + token_idx * emb_dim + emb_dim // 2 + tl.arange(0, emb_dim // 2) + ) + sin_right = tl.load( + SIN + token_idx * emb_dim + emb_dim // 2 + tl.arange(0, emb_dim // 2) + ) + x_1 = x_left_accum * cos_left + x_right_accum * sin_right + x_2 = -x_left_accum * sin_left + x_right_accum * cos_right + dEMB_ptr = dEMB + pid_m * stride_demb_seq + tl.store(dEMB_ptr + tl.arange(0, emb_dim // 2) * 2, x_1) + tl.store(dEMB_ptr + tl.arange(0, emb_dim // 2) * 2 + 1, x_2) + + def _token_stride(tensor: torch.Tensor) -> int: + return tensor.stride(1) if tensor.dim() == 4 else tensor.stride(0) + + class _MLARoPEQTriton(torch.autograd.Function): + """In-place RoPE on the trailing rope slice of q [s, b, h, nope+rope].""" + + @staticmethod + def forward(ctx, q, cos, sin, head_dim_nope, head_dim_rope): + if not q.is_contiguous(): + q = q.contiguous() + s, b, nheads, _ = q.shape + grid = lambda META: (s * b, triton.cdiv(nheads, META["BLOCK_H"])) + rotary_fwd_q_kernel[grid]( + q, + cos, + sin, + head_dim_nope, + head_dim_rope, + nheads, + b, + None, + None, + _token_stride(q), + q.stride(2), + 0, + 1, + ) + ctx.save_for_backward(cos, sin) + ctx.dims = (s, b, nheads, head_dim_nope, head_dim_rope) + return q + + @staticmethod + def backward(ctx, dq): + cos, sin = ctx.saved_tensors + # attention backward may hand over a strided grad; the kernel + # assumes a contiguous [s, b, h, d] layout + dq = dq.contiguous() + s, b, nheads, head_dim_nope, head_dim_rope = ctx.dims + grid = lambda META: (s * b, triton.cdiv(nheads, META["BLOCK_H"])) + rotary_bwd_q_kernel[grid]( + dq, + cos, + sin, + head_dim_nope, + head_dim_rope, + nheads, + b, + None, + None, + _token_stride(dq), + dq.stride(2), + 0, + 1, + ) + return dq, None, None, None, None + + class _MLARoPEKVTriton(torch.autograd.Function): + """kv [s, b, h, nope+v] + shared rope head [s, b, 1, rope] -> (k, v).""" + + @staticmethod + def forward(ctx, kv, k_pos_emb, cos, sin, head_dim_nope, head_dim_rope, head_dim_v): + if not kv.is_contiguous(): + kv = kv.contiguous() + s, b, nheads, _ = kv.shape + o_key = kv.new_empty(s, b, nheads, head_dim_nope + head_dim_rope) + o_value = kv.new_empty(s, b, nheads, head_dim_v) + grid = lambda META: (s * b, triton.cdiv(nheads, META["BLOCK_H"])) + rotary_fwd_kv_kernel[grid]( + kv, + k_pos_emb, + o_key, + o_value, + cos, + sin, + head_dim_rope, + head_dim_nope, + head_dim_v, + nheads, + b, + None, + None, + _token_stride(kv), + kv.stride(2), + _token_stride(k_pos_emb), + _token_stride(o_key), + o_key.stride(2), + _token_stride(o_value), + o_value.stride(2), + 0, + 1, + ) + ctx.save_for_backward(cos, sin) + ctx.dims = (s, b, nheads, head_dim_nope, head_dim_rope, head_dim_v) + return o_key, o_value + + @staticmethod + def backward(ctx, dk_out, dv_out): + cos, sin = ctx.saved_tensors + s, b, nheads, ndp, ndr, ndv = ctx.dims + dk_out = dk_out.contiguous() + dv_out = dv_out.contiguous() + d_kv = dk_out.new_empty(s, b, nheads, ndp + ndv) + d_emb = dk_out.new_empty(s, b, 1, ndr) + grid = lambda META: (s * b, triton.cdiv(nheads, META["BLOCK_H"])) + rotary_bwd_kv_kernel[grid]( + dk_out, + dv_out, + d_kv, + d_emb, + cos, + sin, + ndr, + ndp, + ndv, + nheads, + b, + None, + None, + _token_stride(dk_out), + dk_out.stride(2), + _token_stride(dv_out), + dv_out.stride(2), + _token_stride(d_kv), + d_kv.stride(2), + _token_stride(d_emb), + 0, + 1, + ) + return d_kv, d_emb, None, None, None, None, None + + +def _rotate_interleaved_to_neox(x, cos_table, sin_table, seq_dim): + shape = [1, 1, 1, cos_table.shape[-1]] + shape[seq_dim] = cos_table.shape[0] + cos_ = cos_table.view(shape).to(x.dtype) + sin_ = sin_table.view(shape).to(x.dtype) + half = x.shape[-1] // 2 + x_1 = x[..., 0::2] + x_2 = x[..., 1::2] + x_left = x_1 * cos_[..., :half] - x_2 * sin_[..., :half] + x_right = x_2 * cos_[..., half:] + x_1 * sin_[..., half:] + return torch.cat((x_left, x_right), dim=-1) + + +def apply_mla_rope_q( + q: torch.Tensor, + cos_table: torch.Tensor, + sin_table: torch.Tensor, + head_dim_nope: int, + head_dim_rope: int, + tensor_format: str = "sbhd", +) -> torch.Tensor: + """RoPE on the trailing ``head_dim_rope`` slice of q; in place on the Triton path.""" + if HAVE_TRITON and tensor_format == "sbhd": + return _MLARoPEQTriton.apply(q, cos_table, sin_table, head_dim_nope, head_dim_rope) + seq_dim = 0 if tensor_format == "sbhd" else 1 + q_rope = _rotate_interleaved_to_neox(q[..., head_dim_nope:], cos_table, sin_table, seq_dim) + return torch.cat((q[..., :head_dim_nope], q_rope), dim=-1) + + +def apply_mla_rope_kv( + kv: torch.Tensor, + k_pos_emb: torch.Tensor, + cos_table: torch.Tensor, + sin_table: torch.Tensor, + head_dim_nope: int, + head_dim_rope: int, + head_dim_v: int, + tensor_format: str = "sbhd", +) -> Tuple[torch.Tensor, torch.Tensor]: + """Build (k, v) from kv ``[.., h, nope+v]`` and the shared rope head ``[.., 1, rope]``.""" + if HAVE_TRITON and tensor_format == "sbhd": + return _MLARoPEKVTriton.apply( + kv, k_pos_emb, cos_table, sin_table, head_dim_nope, head_dim_rope, head_dim_v + ) + seq_dim = 0 if tensor_format == "sbhd" else 1 + k_nope = kv[..., :head_dim_nope] + v = kv[..., head_dim_nope : head_dim_nope + head_dim_v] + k_rope = _rotate_interleaved_to_neox(k_pos_emb, cos_table, sin_table, seq_dim) + k_rope = k_rope.expand(*k_nope.shape[:-1], -1) + return torch.cat((k_nope, k_rope), dim=-1), v.contiguous() diff --git a/transformer_engine/pytorch/models/deepseek_v3/multi_latent_attention.py b/transformer_engine/pytorch/models/deepseek_v3/multi_latent_attention.py index a36075f2c5..e5840a812f 100644 --- a/transformer_engine/pytorch/models/deepseek_v3/multi_latent_attention.py +++ b/transformer_engine/pytorch/models/deepseek_v3/multi_latent_attention.py @@ -9,8 +9,12 @@ import torch from transformer_engine.pytorch.module import Linear, LayerNormLinear -from transformer_engine.pytorch.attention import DotProductAttention, RotaryPositionEmbedding -from transformer_engine.pytorch.attention.rope import apply_rotary_pos_emb +from transformer_engine.pytorch.attention import DotProductAttention +from transformer_engine.pytorch.models.deepseek_v3.mla_rope import ( + apply_mla_rope_kv, + apply_mla_rope_q, + build_rope_tables, +) __all__ = ["MultiLatentAttention"] @@ -29,6 +33,10 @@ class MultiLatentAttention(torch.nn.Module): ``kv_channels=(qk_nope_head_dim + qk_rope_head_dim, v_head_dim)``, which supports the cuDNN fused attention backend. + RoPE uses the fused MLA kernels from :mod:`.mla_rope` (in-place on the + query rope slice, single-pass key/value assembly); the rope slice follows + the HF/Megatron DeepSeekV3 convention (interleaved weights, NeoX output). + Parameters ---------- hidden_size : int @@ -126,8 +134,8 @@ def __init__( **common, ) - self.rope = RotaryPositionEmbedding(qk_rope_head_dim, rotary_base=rotary_base) - self._rope_freqs: Optional[torch.Tensor] = None + self.rotary_base = rotary_base + self._rope_tables: Optional[tuple] = None self.core_attention = DotProductAttention( num_attention_heads, @@ -140,10 +148,13 @@ def __init__( tp_size=tp_size, ) - def _rope_freqs_for(self, seq_len: int, device: torch.device) -> torch.Tensor: - if self._rope_freqs is None or self._rope_freqs.shape[0] < seq_len: - self._rope_freqs = self.rope(seq_len).to(device) - return self._rope_freqs[:seq_len] + def _rope_tables_for(self, seq_len: int, device: torch.device): + if self._rope_tables is None or self._rope_tables[0].shape[0] < seq_len: + self._rope_tables = build_rope_tables( + seq_len, self.qk_rope_head_dim, base=self.rotary_base, device=device + ) + cos, sin = self._rope_tables + return cos[:seq_len], sin[:seq_len] def forward( self, @@ -175,26 +186,26 @@ def forward( kv_latent, k_pos = torch.split(kv_down, [self.kv_lora_rank, self.qk_rope_head_dim], dim=-1) kv = self.kv_up_proj(kv_latent) kv = kv.view(*kv.shape[:-1], heads, self.qk_nope_head_dim + self.v_head_dim) - k_nope, v = torch.split(kv, [self.qk_nope_head_dim, self.v_head_dim], dim=-1) - - freqs = self._rope_freqs_for(seq_len, hidden_states.device) - q_rope = apply_rotary_pos_emb( - q[..., self.qk_nope_head_dim :].contiguous(), - freqs, - tensor_format=self.qkv_format, - fused=True, + + cos, sin = self._rope_tables_for(seq_len, hidden_states.device) + q = apply_mla_rope_q( + q, cos, sin, self.qk_nope_head_dim, self.qk_rope_head_dim, self.qkv_format ) - k_rope = apply_rotary_pos_emb( - k_pos.unsqueeze(-2), freqs, tensor_format=self.qkv_format, fused=True + k, v = apply_mla_rope_kv( + kv, + k_pos.unsqueeze(-2), + cos, + sin, + self.qk_nope_head_dim, + self.qk_rope_head_dim, + self.v_head_dim, + self.qkv_format, ) - q = torch.cat([q[..., : self.qk_nope_head_dim], q_rope], dim=-1) - k = torch.cat([k_nope, k_rope.expand(*k_nope.shape[:-1], -1)], dim=-1) - context = self.core_attention( q, k, - v.contiguous(), + v, attention_mask=attention_mask, qkv_format=self.qkv_format, attn_mask_type=attn_mask_type, From c713af7bf06fca793e48c86be6fd203b4836bf14 Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Fri, 21 Aug 2026 11:42:40 +0200 Subject: [PATCH 08/12] Add HF transformers numeric reference test for DeepSeekV3Layer Maps HF DeepseekV3DecoderLayer weights into DeepSeekV3Layer (GLU interleave for routed experts, fused latent norms) and checks forward and input grads match within bf16 tolerance. Expose layernorm_epsilon on MultiLatentAttention (HF latent RMSNorms use 1e-6). Co-Authored-By: Claude Fable 5 Signed-off-by: Pawel Gadzinski --- tests/pytorch/test_deepseek_hf.py | 144 ++++++++++++++++++ .../deepseek_v3/multi_latent_attention.py | 5 + 2 files changed, 149 insertions(+) create mode 100644 tests/pytorch/test_deepseek_hf.py diff --git a/tests/pytorch/test_deepseek_hf.py b/tests/pytorch/test_deepseek_hf.py new file mode 100644 index 0000000000..08f1c06ccd --- /dev/null +++ b/tests/pytorch/test_deepseek_hf.py @@ -0,0 +1,144 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +"""Numeric comparison of DeepSeekV3Layer against the HF transformers reference.""" + +import pytest +import torch + +transformers = pytest.importorskip("transformers") +from transformers.models.deepseek_v3.configuration_deepseek_v3 import DeepseekV3Config +from transformers.models.deepseek_v3.modeling_deepseek_v3 import ( + DeepseekV3DecoderLayer, + DeepseekV3RotaryEmbedding, +) + +from transformer_engine.pytorch.models import DeepSeekV3Layer +from transformer_engine.pytorch.utils import interleave_glu_tensor + +SEQ, BATCH = 64, 2 +HIDDEN, HEADS = 256, 4 +Q_LORA, KV_LORA = 96, 64 +NOPE, ROPE, VDIM = 64, 32, 64 +NUM_EXPERTS, TOPK, N_GROUP, TOPK_GROUP = 16, 4, 4, 2 +MOE_FFN, N_SHARED = 128, 1 +DTYPE = torch.bfloat16 + + +def _hf_config(): + return DeepseekV3Config( + hidden_size=HIDDEN, + intermediate_size=4 * HIDDEN, + moe_intermediate_size=MOE_FFN, + num_hidden_layers=1, + num_attention_heads=HEADS, + num_key_value_heads=HEADS, + n_shared_experts=N_SHARED, + n_routed_experts=NUM_EXPERTS, + routed_scaling_factor=2.5, + kv_lora_rank=KV_LORA, + q_lora_rank=Q_LORA, + qk_rope_head_dim=ROPE, + v_head_dim=VDIM, + qk_nope_head_dim=NOPE, + n_group=N_GROUP, + topk_group=TOPK_GROUP, + num_experts_per_tok=TOPK, + first_k_dense_replace=0, + norm_topk_prob=True, + rms_norm_eps=1e-5, + attention_bias=False, + attention_dropout=0.0, + rope_interleave=True, + _attn_implementation="eager", + ) + + +def _init_hf_layer(config): + torch.manual_seed(0) + layer = DeepseekV3DecoderLayer(config, layer_idx=0).to(device="cuda", dtype=DTYPE) + with torch.no_grad(): + for name, p in layer.named_parameters(): + if "layernorm" in name or "norm" in name: + p.copy_(1.0 + 0.1 * torch.randn_like(p)) + else: + p.normal_(0.0, 0.02) + bias = layer.mlp.gate.e_score_correction_bias + bias.copy_(0.1 * torch.randn_like(bias)) + return layer + + +def _build_te_layer(hf): + te_layer = DeepSeekV3Layer( + HIDDEN, + HEADS, + num_experts=NUM_EXPERTS, + moe_ffn_hidden_size=MOE_FFN, + topk=TOPK, + num_groups=N_GROUP, + group_topk=TOPK_GROUP, + routed_scaling_factor=2.5, + shared_expert_ffn_hidden_size=MOE_FFN * N_SHARED, + q_lora_rank=Q_LORA, + kv_lora_rank=KV_LORA, + qk_nope_head_dim=NOPE, + qk_rope_head_dim=ROPE, + v_head_dim=VDIM, + params_dtype=DTYPE, + ) + attn, mla = hf.self_attn, te_layer.self_attention + with torch.no_grad(): + te_layer.input_layernorm.weight.copy_(hf.input_layernorm.weight) + te_layer.pre_mlp_layernorm.weight.copy_(hf.post_attention_layernorm.weight) + + mla.q_down_proj.weight.copy_(attn.q_a_proj.weight) + mla.q_up_proj.layer_norm_weight.copy_(attn.q_a_layernorm.weight) + mla.q_up_proj.weight.copy_(attn.q_b_proj.weight) + mla.kv_down_proj.weight.copy_(attn.kv_a_proj_with_mqa.weight) + mla.kv_up_proj.layer_norm_weight.copy_(attn.kv_a_layernorm.weight) + mla.kv_up_proj.weight.copy_(attn.kv_b_proj.weight) + mla.out_proj.weight.copy_(attn.o_proj.weight) + + moe = te_layer.mlp + moe.gate.weight.copy_(hf.mlp.gate.weight) + moe.expert_bias.copy_(hf.mlp.gate.e_score_correction_bias) + fc1, _, fc2 = moe.experts + for e in range(NUM_EXPERTS): + getattr(fc1, f"weight{e}").copy_( + interleave_glu_tensor(hf.mlp.experts.gate_up_proj[e], 32) + ) + getattr(fc2, f"weight{e}").copy_(hf.mlp.experts.down_proj[e]) + shared = hf.mlp.shared_experts + moe.shared_expert[0].weight.copy_( + torch.cat([shared.gate_proj.weight, shared.up_proj.weight], dim=0) + ) + moe.shared_expert[2].weight.copy_(shared.down_proj.weight) + return te_layer + + +def test_layer_matches_hf(): + config = _hf_config() + hf = _init_hf_layer(config) + te_layer = _build_te_layer(hf) + + torch.manual_seed(1) + x = torch.randn(BATCH, SEQ, HIDDEN, dtype=DTYPE, device="cuda") + x_hf = x.clone().requires_grad_(True) + x_te = x.transpose(0, 1).contiguous().requires_grad_(True) # sbhd + + rotary = DeepseekV3RotaryEmbedding(config).to("cuda") + position_ids = torch.arange(SEQ, device="cuda").unsqueeze(0).expand(BATCH, -1) + cos, sin = rotary(x_hf, position_ids) + causal = torch.full((SEQ, SEQ), float("-inf"), device="cuda", dtype=DTYPE).triu(1) + causal = causal[None, None].expand(BATCH, 1, SEQ, SEQ) + + out_hf = hf(x_hf, attention_mask=causal, position_embeddings=(cos, sin)) + out_te = te_layer(x_te) + + torch.testing.assert_close(out_te.transpose(0, 1), out_hf, rtol=5e-2, atol=5e-2) + + grad = torch.randn_like(out_hf) + out_hf.backward(grad) + out_te.backward(grad.transpose(0, 1).contiguous()) + torch.testing.assert_close(x_te.grad.transpose(0, 1), x_hf.grad, rtol=5e-2, atol=5e-2) diff --git a/transformer_engine/pytorch/models/deepseek_v3/multi_latent_attention.py b/transformer_engine/pytorch/models/deepseek_v3/multi_latent_attention.py index e5840a812f..0ddf2f75b2 100644 --- a/transformer_engine/pytorch/models/deepseek_v3/multi_latent_attention.py +++ b/transformer_engine/pytorch/models/deepseek_v3/multi_latent_attention.py @@ -57,6 +57,8 @@ class MultiLatentAttention(torch.nn.Module): dropout probability on attention scores. attn_mask_type : str, default = "causal" attention mask type passed to :class:`DotProductAttention`. + layernorm_epsilon : float, default = 1e-6 + epsilon of the latent RMSNorms (matches DeepSeekV3). rotary_base : float, default = 10000.0 RoPE base. softmax_scale : float, optional @@ -83,6 +85,7 @@ def __init__( v_head_dim: int = 128, attention_dropout: float = 0.0, attn_mask_type: str = "causal", + layernorm_epsilon: float = 1e-6, rotary_base: float = 10000.0, softmax_scale: Optional[float] = None, qkv_format: str = "sbhd", @@ -113,6 +116,7 @@ def __init__( q_lora_rank, num_attention_heads * self.qk_head_dim, normalization="RMSNorm", + eps=layernorm_epsilon, parallel_mode="column" if tp_size > 1 else None, **tp, **common, @@ -122,6 +126,7 @@ def __init__( kv_lora_rank, num_attention_heads * (qk_nope_head_dim + v_head_dim), normalization="RMSNorm", + eps=layernorm_epsilon, parallel_mode="column" if tp_size > 1 else None, **tp, **common, From 88028a94f061e0ce075cfd788cf01e580d571ad4 Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Fri, 21 Aug 2026 12:08:12 +0200 Subject: [PATCH 09/12] Docstring cleanups for lint and docs build Co-Authored-By: Claude Fable 5 Signed-off-by: Pawel Gadzinski --- .../pytorch/models/deepseek_v3/mla_rope.py | 28 ++++++++++++++++--- .../models/deepseek_v3/transformer_layer.py | 13 +++++---- 2 files changed, 31 insertions(+), 10 deletions(-) diff --git a/transformer_engine/pytorch/models/deepseek_v3/mla_rope.py b/transformer_engine/pytorch/models/deepseek_v3/mla_rope.py index 350bedb69b..f8c01f75d6 100644 --- a/transformer_engine/pytorch/models/deepseek_v3/mla_rope.py +++ b/transformer_engine/pytorch/models/deepseek_v3/mla_rope.py @@ -92,6 +92,7 @@ def rotary_fwd_q_kernel( cp_size, BLOCK_H: tl.constexpr, ): + """In-place RoPE fwd on the trailing rope slice of q.""" pid_m = tl.program_id(axis=0) pid_head = tl.program_id(axis=1) if cu_seqlens_q is None: @@ -139,6 +140,7 @@ def rotary_bwd_q_kernel( cp_size, BLOCK_H: tl.constexpr, ): + """In-place RoPE bwd on the trailing rope slice of dq.""" pid_m = tl.program_id(axis=0) pid_head = tl.program_id(axis=1) if cu_seqlens_q is None: @@ -195,6 +197,7 @@ def rotary_fwd_kv_kernel( cp_size, BLOCK_H: tl.constexpr, ): + """Fwd: build (key, value) from kv and the shared rotated rope head.""" pid_m = tl.program_id(axis=0) pid_head = tl.program_id(axis=1) if cu_seqlens_kv is None: @@ -262,6 +265,7 @@ def rotary_bwd_kv_kernel( cp_size, BLOCK_H: tl.constexpr, ): + """Bwd: scatter (dk, dv) into dkv and reduce rope-slice grads into demb.""" pid_m = tl.program_id(axis=0) pid_head = tl.program_id(axis=1) if cu_seqlens_kv is None: @@ -320,10 +324,14 @@ class _MLARoPEQTriton(torch.autograd.Function): @staticmethod def forward(ctx, q, cos, sin, head_dim_nope, head_dim_rope): + """Rotate the rope slice of q in place.""" if not q.is_contiguous(): q = q.contiguous() s, b, nheads, _ = q.shape - grid = lambda META: (s * b, triton.cdiv(nheads, META["BLOCK_H"])) + + def grid(meta): + return (s * b, triton.cdiv(nheads, meta["BLOCK_H"])) + rotary_fwd_q_kernel[grid]( q, cos, @@ -345,12 +353,16 @@ def forward(ctx, q, cos, sin, head_dim_nope, head_dim_rope): @staticmethod def backward(ctx, dq): + """Counter-rotate the rope slice of dq (in place on the copy).""" cos, sin = ctx.saved_tensors # attention backward may hand over a strided grad; the kernel # assumes a contiguous [s, b, h, d] layout dq = dq.contiguous() s, b, nheads, head_dim_nope, head_dim_rope = ctx.dims - grid = lambda META: (s * b, triton.cdiv(nheads, META["BLOCK_H"])) + + def grid(meta): + return (s * b, triton.cdiv(nheads, meta["BLOCK_H"])) + rotary_bwd_q_kernel[grid]( dq, cos, @@ -373,12 +385,16 @@ class _MLARoPEKVTriton(torch.autograd.Function): @staticmethod def forward(ctx, kv, k_pos_emb, cos, sin, head_dim_nope, head_dim_rope, head_dim_v): + """Build (k, v) from kv and the shared rope head.""" if not kv.is_contiguous(): kv = kv.contiguous() s, b, nheads, _ = kv.shape o_key = kv.new_empty(s, b, nheads, head_dim_nope + head_dim_rope) o_value = kv.new_empty(s, b, nheads, head_dim_v) - grid = lambda META: (s * b, triton.cdiv(nheads, META["BLOCK_H"])) + + def grid(meta): + return (s * b, triton.cdiv(nheads, meta["BLOCK_H"])) + rotary_fwd_kv_kernel[grid]( kv, k_pos_emb, @@ -409,13 +425,17 @@ def forward(ctx, kv, k_pos_emb, cos, sin, head_dim_nope, head_dim_rope, head_dim @staticmethod def backward(ctx, dk_out, dv_out): + """Gradients for (kv, k_pos_emb) from (dk, dv).""" cos, sin = ctx.saved_tensors s, b, nheads, ndp, ndr, ndv = ctx.dims dk_out = dk_out.contiguous() dv_out = dv_out.contiguous() d_kv = dk_out.new_empty(s, b, nheads, ndp + ndv) d_emb = dk_out.new_empty(s, b, 1, ndr) - grid = lambda META: (s * b, triton.cdiv(nheads, META["BLOCK_H"])) + + def grid(meta): + return (s * b, triton.cdiv(nheads, meta["BLOCK_H"])) + rotary_bwd_kv_kernel[grid]( dk_out, dv_out, diff --git a/transformer_engine/pytorch/models/deepseek_v3/transformer_layer.py b/transformer_engine/pytorch/models/deepseek_v3/transformer_layer.py index af1eeb1a95..f41ec0061c 100644 --- a/transformer_engine/pytorch/models/deepseek_v3/transformer_layer.py +++ b/transformer_engine/pytorch/models/deepseek_v3/transformer_layer.py @@ -40,12 +40,13 @@ class DeepSeekV3Layer(torch.nn.Module): ffn size of each routed expert (required with MoE). hidden_dropout : float, default = 0.0 dropout probability on the residual branches. - kwargs common to the submodules (``q_lora_rank``, ``kv_lora_rank``, - ``qk_nope_head_dim``, ``qk_rope_head_dim``, ``v_head_dim``, - ``attention_dropout``, ``attn_mask_type``, ``qkv_format``, ``topk``, - ``num_groups``, ``group_topk``, ``routed_scaling_factor``, - ``shared_expert_ffn_hidden_size``, EP options, ...) are forwarded to - :class:`MultiLatentAttention` and :class:`DeepSeekV3MoE`. + **kwargs + kwargs common to the submodules (``q_lora_rank``, ``kv_lora_rank``, + ``qk_nope_head_dim``, ``qk_rope_head_dim``, ``v_head_dim``, + ``attention_dropout``, ``attn_mask_type``, ``qkv_format``, ``topk``, + ``num_groups``, ``group_topk``, ``routed_scaling_factor``, + ``shared_expert_ffn_hidden_size``, EP options, ...), forwarded to + :class:`MultiLatentAttention` and :class:`DeepSeekV3MoE`. """ _MLA_KWARGS = frozenset( From 5f68c9bbdaa984abb35ba071f44c0271a8fce986 Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Fri, 21 Aug 2026 14:09:24 +0200 Subject: [PATCH 10/12] Move model-specific layers to a dedicated docs page docs/api/pytorch_models.rst: usage (local and EP), fused-path notes, HF checkpoint weight mapping, and the class API; linked from the PyTorch API page via a toctree entry. Co-Authored-By: Claude Fable 5 Signed-off-by: Pawel Gadzinski --- docs/api/pytorch.rst | 8 ++- docs/api/pytorch_models.rst | 113 ++++++++++++++++++++++++++++++++++++ 2 files changed, 118 insertions(+), 3 deletions(-) create mode 100644 docs/api/pytorch_models.rst diff --git a/docs/api/pytorch.rst b/docs/api/pytorch.rst index bd3099b590..4fa279cfbc 100644 --- a/docs/api/pytorch.rst +++ b/docs/api/pytorch.rst @@ -62,11 +62,13 @@ PyTorch Model-specific layers --------------------- -.. autoapiclass:: transformer_engine.pytorch.models.DeepSeekV3Layer(**kwargs) +Full transformer layers for specific model families live in +``transformer_engine.pytorch.models``: -.. autoapiclass:: transformer_engine.pytorch.models.DeepSeekV3MoE(**kwargs) +.. toctree:: + :maxdepth: 1 -.. autoapiclass:: transformer_engine.pytorch.models.MultiLatentAttention(**kwargs) + pytorch_models Data types ---------- diff --git a/docs/api/pytorch_models.rst b/docs/api/pytorch_models.rst new file mode 100644 index 0000000000..0534acb282 --- /dev/null +++ b/docs/api/pytorch_models.rst @@ -0,0 +1,113 @@ +.. + Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + + See LICENSE for license information. + +Model-specific layers (te.models) +================================= + +The ``transformer_engine.pytorch.models`` namespace holds full transformer +layers for specific model families, composed from Transformer Engine modules +and fused kernels. Each family lives in its own subpackage. + +DeepSeek-V3 +----------- + +A DeepSeek-V3 transformer layer analogous to +:class:`transformer_engine.pytorch.TransformerLayer`: Multi-Latent Attention +(MLA) with low-rank q/kv latents and decoupled RoPE/NoPE heads, plus a +DeepSeek-style Mixture of Experts block (fused sigmoid router with +aux-loss-free expert bias and node-limited grouped top-k, grouped-GEMM SwiGLU +experts, optional shared expert). The same architecture is used by other +model families (e.g. GLM-5, Kimi K2), which can reuse these modules. + +Basic usage (single GPU, all experts local): + +.. code-block:: python + + import torch + import transformer_engine.pytorch as te + + layer = te.models.DeepSeekV3Layer( + hidden_size=7168, + num_attention_heads=128, + num_experts=64, + moe_ffn_hidden_size=2048, + topk=8, + shared_expert_ffn_hidden_size=2048, + params_dtype=torch.bfloat16, + ) + x = torch.randn(seq_len, batch, 7168, dtype=torch.bfloat16, device="cuda") + y = layer(x) # sbhd layout + +Expert parallelism routes tokens between GPUs with the NCCL EP backend +(``transformer_engine.pytorch.ep``). Call ``ep_bootstrap`` once per process +before the first forward; EP requires bfloat16 inputs and NCCL >= 2.30.4: + +.. code-block:: python + + from transformer_engine.pytorch.ep import ep_bootstrap + + ep_bootstrap(ep_group, num_experts=64, max_tokens_per_rank=tokens, + hidden_dim=7168, num_topk=8, recv_capacity_per_rank=capacity) + layer = te.models.DeepSeekV3Layer( + ..., + ep_group=ep_group, + ep_max_tokens_per_rank=tokens, + ) + +On SM100-class GPUs the routed experts fuse into a single CuTe grouped-GEMM +MLP when running under ``te.autocast`` with an MXFP8/NVFP4 recipe and +``NVTE_CUTEDSL_FUSED_GROUPED_MLP=1``; elsewhere the same modules run unfused +with an identical checkpoint layout. + +Loading HuggingFace checkpoints +^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ + +The layer follows the HF/Megatron DeepSeek-V3 conventions (interleaved rope +weights, sigmoid router bias used for selection only). Weights map from +``transformers`` ``DeepseekV3DecoderLayer`` as follows (latent RMSNorms are +fused into the up-projections): + +.. list-table:: + :header-rows: 1 + + * - Transformer Engine + - HuggingFace + * - ``input_layernorm.weight`` + - ``input_layernorm.weight`` + * - ``pre_mlp_layernorm.weight`` + - ``post_attention_layernorm.weight`` + * - ``self_attention.q_down_proj.weight`` + - ``self_attn.q_a_proj.weight`` + * - ``self_attention.q_up_proj.{layer_norm_weight, weight}`` + - ``self_attn.{q_a_layernorm, q_b_proj}.weight`` + * - ``self_attention.kv_down_proj.weight`` + - ``self_attn.kv_a_proj_with_mqa.weight`` + * - ``self_attention.kv_up_proj.{layer_norm_weight, weight}`` + - ``self_attn.{kv_a_layernorm, kv_b_proj}.weight`` + * - ``self_attention.out_proj.weight`` + - ``self_attn.o_proj.weight`` + * - ``mlp.gate.weight`` / ``mlp.expert_bias`` + - ``mlp.gate.weight`` / ``mlp.gate.e_score_correction_bias`` + * - ``mlp.experts[0].weight{i}`` + - ``interleave_glu_tensor(cat([gate_proj, up_proj]), 32)`` of expert *i* + * - ``mlp.experts[2].weight{i}`` + - ``mlp.experts.down_proj[i]`` + * - ``mlp.shared_expert[0].weight`` / ``[2].weight`` + - ``cat([gate_proj, up_proj])`` / ``down_proj`` of ``shared_experts`` + +See ``tests/pytorch/test_deepseek_hf.py`` for a complete, numerically +verified mapping. + +API +^^^ + +.. autoapiclass:: transformer_engine.pytorch.models.DeepSeekV3Layer(hidden_size, num_attention_heads, **kwargs) + :members: forward + +.. autoapiclass:: transformer_engine.pytorch.models.DeepSeekV3MoE(hidden_size, moe_ffn_hidden_size, num_experts, **kwargs) + :members: forward, update_expert_bias + +.. autoapiclass:: transformer_engine.pytorch.models.MultiLatentAttention(hidden_size, num_attention_heads, **kwargs) + :members: forward From 9db495a6577397f898f8d7a812fff557e1102e6e Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Fri, 21 Aug 2026 14:53:21 +0200 Subject: [PATCH 11/12] Drop HF-transformers comparison test from the repo Keep the verified weight-mapping table in the docs; the comparison itself stays as an out-of-tree script. Co-Authored-By: Claude Fable 5 Signed-off-by: Pawel Gadzinski --- docs/api/pytorch_models.rst | 4 +- tests/pytorch/test_deepseek_hf.py | 144 ------------------------------ 2 files changed, 2 insertions(+), 146 deletions(-) delete mode 100644 tests/pytorch/test_deepseek_hf.py diff --git a/docs/api/pytorch_models.rst b/docs/api/pytorch_models.rst index 0534acb282..665d547cc8 100644 --- a/docs/api/pytorch_models.rst +++ b/docs/api/pytorch_models.rst @@ -97,8 +97,8 @@ fused into the up-projections): * - ``mlp.shared_expert[0].weight`` / ``[2].weight`` - ``cat([gate_proj, up_proj])`` / ``down_proj`` of ``shared_experts`` -See ``tests/pytorch/test_deepseek_hf.py`` for a complete, numerically -verified mapping. +The routed-expert fc1 layout can be produced with +:func:`transformer_engine.pytorch.interleave_glu_tensor`. API ^^^ diff --git a/tests/pytorch/test_deepseek_hf.py b/tests/pytorch/test_deepseek_hf.py deleted file mode 100644 index 08f1c06ccd..0000000000 --- a/tests/pytorch/test_deepseek_hf.py +++ /dev/null @@ -1,144 +0,0 @@ -# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. -# -# See LICENSE for license information. - -"""Numeric comparison of DeepSeekV3Layer against the HF transformers reference.""" - -import pytest -import torch - -transformers = pytest.importorskip("transformers") -from transformers.models.deepseek_v3.configuration_deepseek_v3 import DeepseekV3Config -from transformers.models.deepseek_v3.modeling_deepseek_v3 import ( - DeepseekV3DecoderLayer, - DeepseekV3RotaryEmbedding, -) - -from transformer_engine.pytorch.models import DeepSeekV3Layer -from transformer_engine.pytorch.utils import interleave_glu_tensor - -SEQ, BATCH = 64, 2 -HIDDEN, HEADS = 256, 4 -Q_LORA, KV_LORA = 96, 64 -NOPE, ROPE, VDIM = 64, 32, 64 -NUM_EXPERTS, TOPK, N_GROUP, TOPK_GROUP = 16, 4, 4, 2 -MOE_FFN, N_SHARED = 128, 1 -DTYPE = torch.bfloat16 - - -def _hf_config(): - return DeepseekV3Config( - hidden_size=HIDDEN, - intermediate_size=4 * HIDDEN, - moe_intermediate_size=MOE_FFN, - num_hidden_layers=1, - num_attention_heads=HEADS, - num_key_value_heads=HEADS, - n_shared_experts=N_SHARED, - n_routed_experts=NUM_EXPERTS, - routed_scaling_factor=2.5, - kv_lora_rank=KV_LORA, - q_lora_rank=Q_LORA, - qk_rope_head_dim=ROPE, - v_head_dim=VDIM, - qk_nope_head_dim=NOPE, - n_group=N_GROUP, - topk_group=TOPK_GROUP, - num_experts_per_tok=TOPK, - first_k_dense_replace=0, - norm_topk_prob=True, - rms_norm_eps=1e-5, - attention_bias=False, - attention_dropout=0.0, - rope_interleave=True, - _attn_implementation="eager", - ) - - -def _init_hf_layer(config): - torch.manual_seed(0) - layer = DeepseekV3DecoderLayer(config, layer_idx=0).to(device="cuda", dtype=DTYPE) - with torch.no_grad(): - for name, p in layer.named_parameters(): - if "layernorm" in name or "norm" in name: - p.copy_(1.0 + 0.1 * torch.randn_like(p)) - else: - p.normal_(0.0, 0.02) - bias = layer.mlp.gate.e_score_correction_bias - bias.copy_(0.1 * torch.randn_like(bias)) - return layer - - -def _build_te_layer(hf): - te_layer = DeepSeekV3Layer( - HIDDEN, - HEADS, - num_experts=NUM_EXPERTS, - moe_ffn_hidden_size=MOE_FFN, - topk=TOPK, - num_groups=N_GROUP, - group_topk=TOPK_GROUP, - routed_scaling_factor=2.5, - shared_expert_ffn_hidden_size=MOE_FFN * N_SHARED, - q_lora_rank=Q_LORA, - kv_lora_rank=KV_LORA, - qk_nope_head_dim=NOPE, - qk_rope_head_dim=ROPE, - v_head_dim=VDIM, - params_dtype=DTYPE, - ) - attn, mla = hf.self_attn, te_layer.self_attention - with torch.no_grad(): - te_layer.input_layernorm.weight.copy_(hf.input_layernorm.weight) - te_layer.pre_mlp_layernorm.weight.copy_(hf.post_attention_layernorm.weight) - - mla.q_down_proj.weight.copy_(attn.q_a_proj.weight) - mla.q_up_proj.layer_norm_weight.copy_(attn.q_a_layernorm.weight) - mla.q_up_proj.weight.copy_(attn.q_b_proj.weight) - mla.kv_down_proj.weight.copy_(attn.kv_a_proj_with_mqa.weight) - mla.kv_up_proj.layer_norm_weight.copy_(attn.kv_a_layernorm.weight) - mla.kv_up_proj.weight.copy_(attn.kv_b_proj.weight) - mla.out_proj.weight.copy_(attn.o_proj.weight) - - moe = te_layer.mlp - moe.gate.weight.copy_(hf.mlp.gate.weight) - moe.expert_bias.copy_(hf.mlp.gate.e_score_correction_bias) - fc1, _, fc2 = moe.experts - for e in range(NUM_EXPERTS): - getattr(fc1, f"weight{e}").copy_( - interleave_glu_tensor(hf.mlp.experts.gate_up_proj[e], 32) - ) - getattr(fc2, f"weight{e}").copy_(hf.mlp.experts.down_proj[e]) - shared = hf.mlp.shared_experts - moe.shared_expert[0].weight.copy_( - torch.cat([shared.gate_proj.weight, shared.up_proj.weight], dim=0) - ) - moe.shared_expert[2].weight.copy_(shared.down_proj.weight) - return te_layer - - -def test_layer_matches_hf(): - config = _hf_config() - hf = _init_hf_layer(config) - te_layer = _build_te_layer(hf) - - torch.manual_seed(1) - x = torch.randn(BATCH, SEQ, HIDDEN, dtype=DTYPE, device="cuda") - x_hf = x.clone().requires_grad_(True) - x_te = x.transpose(0, 1).contiguous().requires_grad_(True) # sbhd - - rotary = DeepseekV3RotaryEmbedding(config).to("cuda") - position_ids = torch.arange(SEQ, device="cuda").unsqueeze(0).expand(BATCH, -1) - cos, sin = rotary(x_hf, position_ids) - causal = torch.full((SEQ, SEQ), float("-inf"), device="cuda", dtype=DTYPE).triu(1) - causal = causal[None, None].expand(BATCH, 1, SEQ, SEQ) - - out_hf = hf(x_hf, attention_mask=causal, position_embeddings=(cos, sin)) - out_te = te_layer(x_te) - - torch.testing.assert_close(out_te.transpose(0, 1), out_hf, rtol=5e-2, atol=5e-2) - - grad = torch.randn_like(out_hf) - out_hf.backward(grad) - out_te.backward(grad.transpose(0, 1).contiguous()) - torch.testing.assert_close(x_te.grad.transpose(0, 1), x_hf.grad, rtol=5e-2, atol=5e-2) From 4883b1728218c11cba4c14d77f46f98a66d57f5a Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Fri, 21 Aug 2026 14:54:35 +0200 Subject: [PATCH 12/12] Docs: reduce models page to a plain API listing Co-Authored-By: Claude Fable 5 Signed-off-by: Pawel Gadzinski --- docs/api/pytorch.rst | 7 +-- docs/api/pytorch_models.rst | 98 +------------------------------------ 2 files changed, 4 insertions(+), 101 deletions(-) diff --git a/docs/api/pytorch.rst b/docs/api/pytorch.rst index 4fa279cfbc..1c39469f2c 100644 --- a/docs/api/pytorch.rst +++ b/docs/api/pytorch.rst @@ -59,11 +59,8 @@ PyTorch .. autoapifunction:: transformer_engine.pytorch.deinterleave_glu_tensor -Model-specific layers ---------------------- - -Full transformer layers for specific model families live in -``transformer_engine.pytorch.models``: +Models +------ .. toctree:: :maxdepth: 1 diff --git a/docs/api/pytorch_models.rst b/docs/api/pytorch_models.rst index 665d547cc8..2cde879ffb 100644 --- a/docs/api/pytorch_models.rst +++ b/docs/api/pytorch_models.rst @@ -3,106 +3,12 @@ See LICENSE for license information. -Model-specific layers (te.models) -================================= - -The ``transformer_engine.pytorch.models`` namespace holds full transformer -layers for specific model families, composed from Transformer Engine modules -and fused kernels. Each family lives in its own subpackage. +Models +====== DeepSeek-V3 ----------- -A DeepSeek-V3 transformer layer analogous to -:class:`transformer_engine.pytorch.TransformerLayer`: Multi-Latent Attention -(MLA) with low-rank q/kv latents and decoupled RoPE/NoPE heads, plus a -DeepSeek-style Mixture of Experts block (fused sigmoid router with -aux-loss-free expert bias and node-limited grouped top-k, grouped-GEMM SwiGLU -experts, optional shared expert). The same architecture is used by other -model families (e.g. GLM-5, Kimi K2), which can reuse these modules. - -Basic usage (single GPU, all experts local): - -.. code-block:: python - - import torch - import transformer_engine.pytorch as te - - layer = te.models.DeepSeekV3Layer( - hidden_size=7168, - num_attention_heads=128, - num_experts=64, - moe_ffn_hidden_size=2048, - topk=8, - shared_expert_ffn_hidden_size=2048, - params_dtype=torch.bfloat16, - ) - x = torch.randn(seq_len, batch, 7168, dtype=torch.bfloat16, device="cuda") - y = layer(x) # sbhd layout - -Expert parallelism routes tokens between GPUs with the NCCL EP backend -(``transformer_engine.pytorch.ep``). Call ``ep_bootstrap`` once per process -before the first forward; EP requires bfloat16 inputs and NCCL >= 2.30.4: - -.. code-block:: python - - from transformer_engine.pytorch.ep import ep_bootstrap - - ep_bootstrap(ep_group, num_experts=64, max_tokens_per_rank=tokens, - hidden_dim=7168, num_topk=8, recv_capacity_per_rank=capacity) - layer = te.models.DeepSeekV3Layer( - ..., - ep_group=ep_group, - ep_max_tokens_per_rank=tokens, - ) - -On SM100-class GPUs the routed experts fuse into a single CuTe grouped-GEMM -MLP when running under ``te.autocast`` with an MXFP8/NVFP4 recipe and -``NVTE_CUTEDSL_FUSED_GROUPED_MLP=1``; elsewhere the same modules run unfused -with an identical checkpoint layout. - -Loading HuggingFace checkpoints -^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^ - -The layer follows the HF/Megatron DeepSeek-V3 conventions (interleaved rope -weights, sigmoid router bias used for selection only). Weights map from -``transformers`` ``DeepseekV3DecoderLayer`` as follows (latent RMSNorms are -fused into the up-projections): - -.. list-table:: - :header-rows: 1 - - * - Transformer Engine - - HuggingFace - * - ``input_layernorm.weight`` - - ``input_layernorm.weight`` - * - ``pre_mlp_layernorm.weight`` - - ``post_attention_layernorm.weight`` - * - ``self_attention.q_down_proj.weight`` - - ``self_attn.q_a_proj.weight`` - * - ``self_attention.q_up_proj.{layer_norm_weight, weight}`` - - ``self_attn.{q_a_layernorm, q_b_proj}.weight`` - * - ``self_attention.kv_down_proj.weight`` - - ``self_attn.kv_a_proj_with_mqa.weight`` - * - ``self_attention.kv_up_proj.{layer_norm_weight, weight}`` - - ``self_attn.{kv_a_layernorm, kv_b_proj}.weight`` - * - ``self_attention.out_proj.weight`` - - ``self_attn.o_proj.weight`` - * - ``mlp.gate.weight`` / ``mlp.expert_bias`` - - ``mlp.gate.weight`` / ``mlp.gate.e_score_correction_bias`` - * - ``mlp.experts[0].weight{i}`` - - ``interleave_glu_tensor(cat([gate_proj, up_proj]), 32)`` of expert *i* - * - ``mlp.experts[2].weight{i}`` - - ``mlp.experts.down_proj[i]`` - * - ``mlp.shared_expert[0].weight`` / ``[2].weight`` - - ``cat([gate_proj, up_proj])`` / ``down_proj`` of ``shared_experts`` - -The routed-expert fc1 layout can be produced with -:func:`transformer_engine.pytorch.interleave_glu_tensor`. - -API -^^^ - .. autoapiclass:: transformer_engine.pytorch.models.DeepSeekV3Layer(hidden_size, num_attention_heads, **kwargs) :members: forward