From 45ed767a7af61a64dc41e29be1b6c0cc14c21d9c Mon Sep 17 00:00:00 2001 From: Pawel Gadzinski Date: Tue, 18 Aug 2026 12:37:21 +0200 Subject: [PATCH 1/7] [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 2/7] 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 3/7] 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 4/7] [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 5/7] 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 6/7] 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 7/7] 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,