diff --git a/docs/api/pytorch.rst b/docs/api/pytorch.rst
index 73a08974c8b..a9bc32c6aa3 100644
--- a/docs/api/pytorch.rst
+++ b/docs/api/pytorch.rst
@@ -6,6 +6,11 @@
PyTorch
=======
+.. autoapiclass:: transformer_engine.pytorch.autocast(enabled=True, calibrating=False, recipe=None, amax_reduction_group=None)
+
+General-purpose layers
+----------------------
+
.. autoapiclass:: transformer_engine.pytorch.Linear(in_features, out_features, **kwargs)
:members: forward, set_tensor_parallel_group
@@ -40,20 +45,34 @@ PyTorch
.. autoapiclass:: transformer_engine.pytorch.TransformerLayer(hidden_size, ffn_hidden_size, num_attention_heads, **kwargs)
:members: forward, set_context_parallel_group, set_tensor_parallel_group
+Model layers
+------------
+
+DeepSeek-V3
+^^^^^^^^^^^
+
+.. autoapiclass:: transformer_engine.pytorch.models.DeepSeekV3Layer(hidden_size, num_attention_heads, **kwargs)
+ :members: forward
+
+.. autoapiclass:: transformer_engine.pytorch.models.DeepSeekV3MoE(hidden_size, moe_ffn_hidden_size, num_experts, **kwargs)
+ :members: forward, update_expert_bias
+
+.. autoapiclass:: transformer_engine.pytorch.models.MultiLatentAttention(hidden_size, num_attention_heads, **kwargs)
+ :members: forward
+
+Other
+-----
+
.. autoapiclass:: transformer_engine.pytorch.dot_product_attention.inference.InferenceParams(max_batch_size, max_sequence_length)
:members: reset, allocate_memory, pre_step, get_seqlens_pre_step, convert_paged_to_nonpaged, step
.. autoapiclass:: transformer_engine.pytorch.CudaRNGStatesTracker()
:members: reset, get_states, set_states, add, fork
-
-.. autoapiclass:: transformer_engine.pytorch.autocast(enabled=True, calibrating=False, recipe=None, amax_reduction_group=None)
-
.. autoapifunction:: transformer_engine.pytorch.quantized_model_init
.. autoapifunction:: transformer_engine.pytorch.checkpoint
-
.. autoapifunction:: transformer_engine.pytorch.make_graphed_callables
.. autoapifunction:: transformer_engine.pytorch.get_cpu_offload_context
@@ -68,6 +87,21 @@ PyTorch
.. autoapifunction:: transformer_engine.pytorch.deinterleave_glu_tensor
+MLA rotary position embeddings
+^^^^^^^^^^^^^^^^^^^^^^^^^^^^^^
+
+These functions apply the decoupled RoPE used by multi-latent attention (MLA).
+The rotary input channels contain interleaved pairs; rotated outputs use
+NeoX half-split order. ``sbhd`` denotes sequence, batch, head and channel axes;
+``bshd`` exchanges the first two axes. Build the cosine/sine tables on the
+same device as the tensors, with one row per sequence position.
+
+.. autoapifunction:: transformer_engine.pytorch.attention.mla_rope.build_rope_tables
+
+.. autoapifunction:: transformer_engine.pytorch.attention.mla_rope.apply_mla_rope_q
+
+.. autoapifunction:: transformer_engine.pytorch.attention.mla_rope.apply_mla_rope_kv
+
Data types
----------
diff --git a/examples/pytorch/deepseek_v3/README.md b/examples/pytorch/deepseek_v3/README.md
new file mode 100644
index 00000000000..5bf185c6895
--- /dev/null
+++ b/examples/pytorch/deepseek_v3/README.md
@@ -0,0 +1,223 @@
+# DeepSeekV3Layer with expert parallelism
+
+A forward + backward benchmark of one `DeepSeekV3Layer`: RMSNorm, multi-latent attention,
+and MoE with a shared expert. Routed experts are sharded across GPUs using NCCL EP.
+
+## Results
+
+GB200 and GB300 GPUs, 4096 tokens per rank, top-k 8, 8 local experts per GPU.
+Times cover one layer's forward + backward; throughput is global, in millions of tokens/s.
+Every number is the median of three independent runs (spread within 0.3 ms).
+See [benchmark configuration](#c-benchmark-configuration) for the full setup.
+
+### TE vs. plain PyTorch (`--dsv3`)
+
+Same layer and dimensions, BF16 unless noted.
+
+All variants run the router backward and clear parameter gradients before each step.
+
+| MoE implementation | 4 GB200 (32 experts) | 4 GB300 (32 experts) | 8 GB300 (64 experts) |
+|---|---:|---:|---:|
+| `naive`: all_to_all + loop over experts | 28.61 ms | 27.21 ms | 27.70 ms |
+| `naive_grouped`: all_to_all + TE grouped GEMM | 17.79 ms | 16.72 ms | 17.13 ms |
+| `te`: NCCL EP + grouped GEMM | 13.68 ms | 12.37 ms | 14.09 ms |
+| `te`, mxfp8 (unfused grouped GEMM) | 10.99 ms | 10.32 ms | 12.11 ms |
+| `te`, mxfp8 fused | 10.11 ms | 9.49 ms | 10.45 ms |
+
+### TE throughput (`--dsv3`)
+
+| Precision | 4 GB200 · Mtok/s | 4 GB300 · Mtok/s | 8 GB300 · Mtok/s |
+|---|---:|---:|---:|
+| BF16 | 1.20 | 1.33 | 2.32 |
+| MXFP8 | 1.49 | 1.59 | 2.71 |
+| MXFP8 fused | 1.62 | 1.73 | 3.14 |
+
+4 GPUs = 1 node / 32 experts; 8 GPUs = 2 nodes / 64 experts.
+Both nodes share one NVLink domain (MNNVL).
+
+“Fused” enables `NVTE_CUTEDSL_FUSED_GROUPED_MLP=1`.
+
+Paired runs on the same hardware and software compared these results with the previous
+MoE alignment policy and its extra 1024-row EP tail margin. The full-layer TE medians
+decreased by 0.1–0.8% across the nine configurations, with no observed slowdown.
+
+## Quick start
+
+Requires SM90+ GPUs with NVLink and an NCCL EP-enabled TE build;
+see [full requirements](#a-requirements). From this directory:
+
+```bash
+bash run_deepseek_v3_layer_ep.sh --dsv3
+```
+
+For MXFP8 with the fused grouped MLP (SM100-class GPUs):
+
+```bash
+NVTE_CUTEDSL_FUSED_GROUPED_MLP=1 bash run_deepseek_v3_layer_ep.sh --dsv3 --recipe mxfp8
+```
+
+## Appendix
+
+[Requirements](#a-requirements) · [Running](#b-running) ·
+[Configuration](#c-benchmark-configuration) · [Profiling](#d-profiling-with-nsys) ·
+[TE kernels](#e-te-kernel-breakdown) ·
+[Naive kernel profiles](#f-naive-kernel-profiles) ·
+[EP internals](#g-ep-implementation-notes)
+
+### A. Requirements
+
+The following NCCL EP requirements apply to `--impl te`. The naive variants use
+ordinary NCCL collectives and can also run on older GPUs without NVLink.
+
+- SM90 or newer GPUs connected with NVLink (NCCL EP falls back to the network transport and
+ deadlocks on PCIe-only nodes).
+- NCCL >= 2.30.4, PyTorch >= 2.11 (symmetric memory), Transformer Engine built with the
+ `3rdparty/nccl-extensions` submodule.
+- `NVTE_CUTEDSL_FUSED_GROUPED_MLP=1` additionally needs SM100-class GPUs and the CuTe DSL
+ (`nvidia-cutlass-dsl`) for the fused MXFP8 grouped MLP.
+
+### B. Running
+
+Run these commands from this directory. Single node, all local GPUs:
+
+```bash
+bash run_deepseek_v3_layer_ep.sh # small dims, bf16
+bash run_deepseek_v3_layer_ep.sh --dsv3 # DeepSeek-V3 layer dims, bf16
+bash run_deepseek_v3_layer_ep.sh --dsv3 --recipe mxfp8 # MXFP8 experts (unfused grouped GEMM)
+NVTE_CUTEDSL_FUSED_GROUPED_MLP=1 bash run_deepseek_v3_layer_ep.sh --dsv3 --recipe mxfp8
+bash run_deepseek_v3_layer_ep.sh --dsv3 --impl naive # all_to_all + loop over experts
+bash run_deepseek_v3_layer_ep.sh --dsv3 --impl naive_grouped # all_to_all + TE grouped GEMM
+```
+
+`--impl` selects the MoE block inside the same layer (attention and norms are identical):
+
+- `te` (default): `DeepSeekV3MoE`, NCCL EP dispatch/combine, experts as one grouped GEMM.
+- `naive`: MoE written with plain PyTorch, no TE MoE code: sigmoid top-k router with expert
+ bias, `all_to_all_single` dispatch and combine (two host syncs per layer for the split sizes),
+ a Python loop over the local experts with dense `F.linear` SwiGLU MLPs, `index_copy` /
+ `index_add` to place results, and a shared expert. This is what an EP MoE looks like before
+ any fused kernels.
+- `naive_grouped`: the same all_to_all dispatch and combine, but the received rows are sorted by
+ local expert and run through one `te.ops.GroupedLinear` / `ScaledSwiGLU` / `GroupedLinear`
+ stack. Isolates the cost of the Python loop from the cost of the communication path.
+
+Multi-node: launch `torchrun` yourself, EP spans every rank:
+
+```bash
+torchrun --nnodes=2 --nproc-per-node=4 --rdzv-backend=c10d --rdzv-endpoint=
:29500 \
+ deepseek_v3_layer_ep.py --dsv3 --recipe mxfp8
+```
+
+Every rank owns `--num-local-experts` experts (default 8), so the expert count is
+`8 * world_size`. Other knobs: `--tokens-per-rank`, `--topk`, `--hidden`, `--num-heads`,
+`--moe-ffn`, the MLA dims (`--q-lora-rank`, `--kv-lora-rank`, `--qk-nope-head-dim`,
+`--qk-rope-head-dim`, `--v-head-dim`), `--warmup`, `--iters`.
+`--warmup 0` is supported; `--iters` must be positive and `--tokens-per-rank` must be a
+positive multiple of four. MXFP8 requires `--impl te`.
+
+### C. Benchmark configuration
+
+| Setting | Value |
+|---|---|
+| hardware | GB200 (SM100), one node; GB300 (SM103), one or two nodes; 4 GPUs per node, both GB300 nodes in one NVLink domain (MNNVL) |
+| software | CUDA 13.3, NCCL 2.30.7, PyTorch 2.13.0a0+9186a08b2c.nv26.07, cuDNN 9.26, cuDNN Frontend 1.29.0, CUTLASS DSL 4.6.2 |
+| NVIDIA driver | GB200: 580.173.02; GB300: 580.173.10 |
+| `--dsv3` dims | hidden 7168, 128 heads, MLA q_lora 1536 / kv_lora 512 / nope 128 / rope 64 / v 128, expert ffn 2048, shared expert ffn 2048 |
+| MoE | 8 local experts per rank (32 on 4 GPUs, 64 on 8), top-k 8, 4096 tokens per rank |
+| precision | bf16 params and activations; `--recipe mxfp8` = MXFP8 block scaling for the expert GEMMs and dense projections |
+| timing | fwd + bwd, 5 warmup, 10 timed iterations per run, median of 3 runs; wall-clock timing with CUDA synchronization, no CUDA graphs or profiler |
+| environment | `OMP_NUM_THREADS=8`, `NVTE_GROUPED_LINEAR_SINGLE_PARAM=0`, `NVTE_ALLOW_NONDETERMINISTIC_ALGO=1`, `NVTE_FLASH_ATTN_V2=1`, `NVTE_FLASH_ATTN_V3=0`, `NVTE_FLASH_ATTN_V4=0`; EP buffer reused |
+
+Each run reuses a normally distributed random input and backpropagates an all-ones
+output gradient. Parameter and input gradients are cleared before each iteration.
+
+### D. Profiling with nsys
+
+The timed iterations run inside a `torch.cuda.profiler.start()` / `stop()` window, so
+`-c cudaProfilerApi` records only them, one NVTX range per iteration:
+
+```bash
+NSYS=1 NVTE_CUTEDSL_FUSED_GROUPED_MLP=1 bash run_deepseek_v3_layer_ep.sh --dsv3 --recipe mxfp8
+# -> results/deepseek_v3_layer_ep_.nsys-rep
+nsys stats --report cuda_gpu_kern_sum results/deepseek_v3_layer_ep_.nsys-rep
+```
+
+The launcher wraps `torchrun`, so all local ranks land in one report. For multi-node runs put
+the same `nsys profile ... -o _%q{SLURM_NODEID}` in front of `torchrun` on each node.
+
+Running under `nsys` adds about 1.5 ms per iteration to these numbers.
+
+### E. TE kernel breakdown
+
+The profiles in sections E and F predate the September 30 measurements and were not
+regenerated for the results above.
+
+8 GPUs, `--dsv3`, MXFP8 fused. Per GPU and iteration, from `nsys stats --report cuda_gpu_kern_sum`
+on one node (kernel time 10.8 ms; the iteration takes 10.5 ms without the profiler). Profiling adds
+skew between nodes, so the wait row is taken from the node that was not slowed down by nsys.
+
+| Group | ms | Kernels |
+|---|---:|---|
+| NCCL EP all-to-all | 2.3 | `nccl_ep_jit_ht_dispatch_kernel` (1.05), `nccl_ep_jit_ht_combine_kernel` (1.29), each twice per iteration (fwd + bwd) |
+| NCCL EP local permute | 0.6 | `local_permute_dup/reduce`: staging buffer to expert-major layout, zero-filled padding |
+| rank skew wait | 0.5 | `ncclDevKernel_AllGather_RING_LL` (routing-map all-gather in prepare, ~0.06 ms of transfer); the first collective of the layer absorbs load imbalance between ranks |
+| fused grouped MLP (cuDNN, MXFP8) | 2.4 | fc1+SwiGLU fwd (0.53), fc2 fwd (0.82), dGLU bwd (0.31), wgrad (0.78) |
+| MXFP8 quantization | 1.1 | `group_quantize_mxfp8` on the recv buffer (0.4), `quantize_mxfp8_kernel_cast_only` for dense GEMM inputs (0.7) |
+| dense MXFP8 GEMMs (MLA projections, shared expert) | 1.4 | `nvjet_sm103_qqtst_*` |
+| attention (cuDNN SDPA) | 0.9 | flash fprop (0.18) + bprop (0.51) + dq / dO helpers |
+| RMSNorm, RoPE, adds | 0.8 | `rmsnorm_fwd/bwd` (0.39), `rotary_*_kv` (0.18), residual adds (0.2) |
+
+Both nodes sit in one NVLink domain, so dispatch and combine move roughly 470 MB per GPU per
+call over NVLink at close to link bandwidth; on an InfiniBand-connected pair of nodes the
+all-to-all share would be much larger.
+
+### F. Naive kernel profiles
+
+Per GPU and iteration, the `naive` MoE spends (8 GPUs, kernel time 24.0 ms of a 27.5 ms
+iteration; the rest is host syncs and launch gaps):
+
+| Group | ms | Details |
+|---|---:|---|
+| expert and dense GEMMs (`nvjet_*`) | 6.5 | 8 separate GEMM pairs per rank instead of one grouped GEMM, plus the MLA projections |
+| elementwise adds | 3.6 | `index_add` and its backward, residuals |
+| `ncclDevKernel_SendRecv` | 3.4 | 8 all_to_all launches per iteration (tokens and probs fwd/bwd, results fwd/bwd, counts, indices) |
+| indexing kernels | 3.3 | `x[tok]`, `nonzero` masks, `index_copy`, `index_fill`, `indexing_backward`, sorts |
+| `FillFunctor` (zeros) | 2.5 | `zeros_like` for the per-expert output buffer and `index_add` targets |
+| device-to-device copies | 2.3 | `index_copy` and gathers materialising per-expert slices |
+| attention, norms | 1.1 | same as in the TE variant |
+
+#### Historical `naive_grouped` profile
+
+8 GPUs, kernel time 17.1 ms of a 17.2 ms iteration. This profile predates the
+router-gradient fix and omits probability-gradient communication. It is retained for
+reference, not for direct comparison with the other profiles or headline timings.
+
+| Group | ms | Details |
+|---|---:|---|
+| `ncclDevKernel_SendRecv` | 3.8 | 7 all_to_all launches in this historical run; the current implementation has 8 |
+| grouped GEMMs (`nvjet_*_ptrGroup_*`) | 4.1 | fc1 / fc2 forward, dgrad, wgrad as grouped GEMMs, same as in `te` |
+| sorting rows by expert and back | 2.6 | `argsort`, gathers (`x[tok]`, `x_recv[by_expert]`), `index_copy`, `indexing_backward` |
+| dense GEMMs (MLA projections, shared expert) | 2.4 | same as in `te` |
+| elementwise adds | 0.7 | `index_add`, residuals |
+| attention, norms | 1.1 | same as in `te` |
+
+#### Interpreting the current results
+
+In the headline GB300 BF16 timings, `naive` -> `naive_grouped` reduces iteration time by
+10.5 ms on 4 GPUs and 10.6 ms on 8 GPUs. `naive_grouped` -> `te` saves another
+4.4 ms and 3.0 ms, respectively. These are end-to-end differences between
+implementations, not isolated measurements of Python-loop or communication overhead.
+
+NCCL EP dispatch and combine write directly into the expert-major layout and zero-fill
+the padding. Both naive variants synchronise with the host to obtain the all_to_all
+split sizes; TE avoids those synchronizations. MXFP8 is only supported by `te`: the naive
+variants would need per-expert row counts padded to the MXFP8 block size.
+
+### G. EP implementation notes
+
+- `ep_bootstrap` must be given the same recv capacity the layer uses:
+ `DeepSeekV3MoE.ep_recv_capacity(ep_size, tokens_per_rank, topk, num_local_experts)`.
+- Per-expert zones in the recv buffer are aligned to 256 rows; the fused grouped MLP requires
+ that alignment.
+- The recv and grad buffers are allocated uninitialized: NCCL EP zero-fills the alignment
+ padding between experts itself.
diff --git a/examples/pytorch/deepseek_v3/deepseek_v3_layer_ep.py b/examples/pytorch/deepseek_v3/deepseek_v3_layer_ep.py
new file mode 100644
index 00000000000..e95a7ecba6d
--- /dev/null
+++ b/examples/pytorch/deepseek_v3/deepseek_v3_layer_ep.py
@@ -0,0 +1,283 @@
+# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+#
+# See LICENSE for license information.
+"""DeepSeekV3Layer with expert parallelism over all ranks: forward + backward timing.
+
+One process per GPU, launched via run_deepseek_v3_layer_ep.sh (torchrun). Every rank
+holds ``--num-local-experts`` routed experts. ``--impl te`` exchanges tokens with NCCL EP
+and runs the experts as one grouped GEMM (DeepSeekV3MoE); ``--impl naive`` is a plain
+PyTorch MoE (all_to_all_single + a Python loop over experts) dropped into the same layer;
+``--impl naive_grouped`` keeps the all_to_all but runs the experts as one TE grouped GEMM.
+Timed iterations run inside a ``torch.cuda.profiler`` window, so
+``nsys profile -c cudaProfilerApi --capture-range-end=stop torchrun ...`` records
+only them.
+"""
+
+import argparse
+import os
+import sys
+import time
+from contextlib import nullcontext
+
+import torch
+import torch.distributed as dist
+import torch.nn.functional as F
+from torch.distributed.nn.functional import all_to_all_single
+
+import transformer_engine.pytorch as te
+from transformer_engine.common import recipe as te_recipe
+from transformer_engine.pytorch.ep import ep_bootstrap, ep_finalize, release_symm_mem_pool
+from transformer_engine.pytorch.models import DeepSeekV3Layer, DeepSeekV3MoE
+
+
+def _parse_args(argv=None):
+ p = argparse.ArgumentParser(description="DeepSeekV3Layer EP example (fwd + bwd)")
+ p.add_argument("--tokens-per-rank", type=int, default=4096)
+ p.add_argument("--hidden", type=int, default=2048)
+ p.add_argument("--num-heads", type=int, default=16)
+ p.add_argument("--moe-ffn", type=int, default=1024)
+ p.add_argument("--num-local-experts", type=int, default=8)
+ p.add_argument("--topk", type=int, default=8)
+ p.add_argument("--q-lora-rank", type=int, default=512)
+ p.add_argument("--kv-lora-rank", type=int, default=256)
+ p.add_argument("--qk-nope-head-dim", type=int, default=64)
+ p.add_argument("--qk-rope-head-dim", type=int, default=32)
+ p.add_argument("--v-head-dim", type=int, default=64)
+ p.add_argument(
+ "--dsv3",
+ action="store_true",
+ help=(
+ "DeepSeek-V3 layer dims (hidden 7168, 128 heads, MLA 1536/512/128/64/128, expert ffn"
+ " 2048)."
+ ),
+ )
+ p.add_argument("--impl", choices=["te", "naive", "naive_grouped"], default="te")
+ p.add_argument("--recipe", choices=["none", "mxfp8"], default="none")
+ p.add_argument(
+ "--ep-buffer-per-call",
+ action="store_true",
+ help="Let the layer create a new EpBuffer per forward instead of reusing one.",
+ )
+ p.add_argument("--warmup", type=int, default=5)
+ p.add_argument("--iters", type=int, default=10)
+ args = p.parse_args(argv)
+ if args.warmup < 0:
+ p.error("--warmup must be non-negative")
+ if args.iters <= 0:
+ p.error("--iters must be positive")
+ if args.tokens_per_rank <= 0 or args.tokens_per_rank % 4:
+ p.error("--tokens-per-rank must be a positive multiple of 4")
+ if args.impl != "te" and args.recipe != "none":
+ p.error("--recipe mxfp8 is only supported with --impl te")
+ if args.dsv3:
+ args.hidden, args.num_heads, args.moe_ffn = 7168, 128, 2048
+ args.q_lora_rank, args.kv_lora_rank = 1536, 512
+ args.qk_nope_head_dim, args.qk_rope_head_dim, args.v_head_dim = 128, 64, 128
+ return args
+
+
+def _autocast(name):
+ if name == "none":
+ return nullcontext()
+ return te.autocast(enabled=True, recipe=te_recipe.MXFP8BlockScaling())
+
+
+class NaiveMoE(torch.nn.Module):
+ """DeepSeek-style MoE with torch all_to_all dispatch/combine: sigmoid top-k router with
+ expert bias, a shared expert, and experts either as a Python loop of dense SwiGLU MLPs or
+ (``grouped=True``) as one TE grouped GEMM stack."""
+
+ def __init__(self, hidden, ffn, num_experts, topk, ep_group, shared_ffn, dtype, grouped=False):
+ super().__init__()
+ self.grouped = grouped
+ self.hidden, self.topk, self.group = hidden, topk, ep_group
+ self.ws, self.rank = dist.get_world_size(ep_group), dist.get_rank(ep_group)
+ self.num_experts, self.local = num_experts, num_experts // self.ws
+ self.gate = torch.nn.Linear(hidden, num_experts, bias=False, dtype=dtype, device="cuda")
+ self.register_buffer("expert_bias", torch.zeros(num_experts, device="cuda"))
+ std = hidden**-0.5
+ if grouped:
+ self.experts = te.ops.Sequential(
+ te.ops.GroupedLinear(self.local, hidden, 2 * ffn, bias=False, dtype=dtype),
+ te.ops.ScaledSwiGLU(glu_interleave_size=32),
+ te.ops.GroupedLinear(self.local, ffn, hidden, bias=False, dtype=dtype),
+ )
+ else:
+ self.w1 = torch.nn.Parameter(
+ torch.randn(self.local, 2 * ffn, hidden, dtype=dtype, device="cuda") * std
+ )
+ self.w2 = torch.nn.Parameter(
+ torch.randn(self.local, hidden, ffn, dtype=dtype, device="cuda") * ffn**-0.5
+ )
+ self.shared_w1 = torch.nn.Linear(
+ hidden, 2 * shared_ffn, bias=False, dtype=dtype, device="cuda"
+ )
+ self.shared_w2 = torch.nn.Linear(shared_ffn, hidden, bias=False, dtype=dtype, device="cuda")
+
+ @staticmethod
+ def _swiglu(h):
+ a, g = h.chunk(2, dim=-1)
+ return F.silu(a) * g
+
+ def forward(self, hidden_states):
+ x = hidden_states.reshape(-1, self.hidden)
+ scores = torch.sigmoid(self.gate(x).float())
+ _, idx = torch.topk(scores + self.expert_bias, self.topk, dim=-1)
+ probs = scores.gather(1, idx)
+ probs = probs / probs.sum(-1, keepdim=True) * 2.5
+ # Dispatch: sort (token, expert) pairs by destination rank, exchange counts, all_to_all.
+ flat_e, flat_p = idx.reshape(-1), probs.reshape(-1)
+ tok = torch.arange(x.shape[0], device=x.device).repeat_interleave(self.topk)
+ order = torch.argsort(flat_e // self.local, stable=True)
+ flat_e, flat_p, tok = flat_e[order], flat_p[order], tok[order]
+ send = torch.bincount(flat_e // self.local, minlength=self.ws)
+ recv = torch.empty_like(send)
+ dist.all_to_all_single(recv, send, group=self.group)
+ send, recv = send.tolist(), recv.tolist()
+ n_recv = sum(recv)
+ x_recv = all_to_all_single(
+ torch.empty(n_recv, self.hidden, dtype=x.dtype, device=x.device),
+ x[tok],
+ recv,
+ send,
+ group=self.group,
+ )
+ e_recv = torch.empty(n_recv, dtype=flat_e.dtype, device=x.device)
+ p_recv = torch.empty(n_recv, dtype=flat_p.dtype, device=x.device)
+ dist.all_to_all_single(e_recv, flat_e.contiguous(), recv, send, group=self.group)
+ p_recv = all_to_all_single(p_recv, flat_p.contiguous(), recv, send, group=self.group)
+ local_e = e_recv - self.rank * self.local
+ if self.grouped:
+ # Experts: sort received rows by local expert, one grouped GEMM stack.
+ by_expert = torch.argsort(local_e, stable=True)
+ counts = torch.bincount(local_e, minlength=self.local)
+ y_sorted = self.experts(
+ x_recv[by_expert], counts, p_recv[by_expert].to(x.dtype), counts
+ )
+ y_recv = torch.empty_like(x_recv).index_copy(0, by_expert, y_sorted)
+ else:
+ # Experts: one dense SwiGLU MLP per local expert.
+ y_recv = torch.zeros_like(x_recv)
+ for e in range(self.local):
+ sel = (local_e == e).nonzero().squeeze(1)
+ if sel.numel() == 0:
+ continue
+ h = self._swiglu(F.linear(x_recv[sel], self.w1[e])) * p_recv[sel, None].to(x.dtype)
+ y_recv = y_recv.index_copy(0, sel, F.linear(h, self.w2[e]))
+ # Combine: reverse all_to_all, sum the top-k contributions per token.
+ y = all_to_all_single(torch.empty_like(x[tok]), y_recv, send, recv, group=self.group)
+ out = torch.zeros_like(x).index_add(0, tok, y)
+ out = out + self.shared_w2(self._swiglu(self.shared_w1(x)))
+ return out.view_as(hidden_states)
+
+
+def main():
+ """Build the layer, run warmup + timed fwd/bwd iterations, print throughput on rank 0."""
+ args = _parse_args()
+ local_rank = int(os.environ["LOCAL_RANK"])
+ torch.cuda.set_device(local_rank)
+ dist.init_process_group("nccl", device_id=torch.device("cuda", local_rank))
+ rank, world_size = dist.get_rank(), dist.get_world_size()
+
+ major, minor = torch.cuda.get_device_capability()
+ if args.impl == "te" and major * 10 + minor < 90:
+ if rank == 0:
+ print(f"SKIPPED: NCCL EP requires SM>=90 (got SM{major}{minor})")
+ dist.destroy_process_group()
+ return 0
+
+ ep_group = dist.new_group(ranks=list(range(world_size)), backend="nccl")
+ dist.all_reduce(torch.zeros(1, device="cuda"), group=ep_group)
+ num_experts = args.num_local_experts * world_size
+ if args.impl == "te":
+ ep_bootstrap(
+ ep_group,
+ num_experts=num_experts,
+ max_tokens_per_rank=args.tokens_per_rank,
+ hidden_dim=args.hidden,
+ num_topk=args.topk,
+ recv_capacity_per_rank=DeepSeekV3MoE.ep_recv_capacity(
+ world_size, args.tokens_per_rank, args.topk, args.num_local_experts
+ ),
+ )
+
+ torch.manual_seed(0)
+ mlp_kwargs = dict(
+ num_experts=num_experts,
+ moe_ffn_hidden_size=args.moe_ffn,
+ shared_expert_ffn_hidden_size=args.moe_ffn,
+ topk=args.topk,
+ ep_group=ep_group,
+ ep_max_tokens_per_rank=args.tokens_per_rank,
+ )
+ if args.impl != "te":
+ # Build the TE MoE without EP (replaced below); keeps the pre-MLP RMSNorm.
+ mlp_kwargs.pop("ep_group"), mlp_kwargs.pop("ep_max_tokens_per_rank")
+ layer = DeepSeekV3Layer(
+ args.hidden,
+ args.num_heads,
+ params_dtype=torch.bfloat16,
+ **mlp_kwargs,
+ q_lora_rank=args.q_lora_rank,
+ kv_lora_rank=args.kv_lora_rank,
+ qk_nope_head_dim=args.qk_nope_head_dim,
+ qk_rope_head_dim=args.qk_rope_head_dim,
+ v_head_dim=args.v_head_dim,
+ )
+ if args.impl != "te":
+ layer.mlp = NaiveMoE(
+ args.hidden,
+ args.moe_ffn,
+ num_experts,
+ args.topk,
+ ep_group,
+ args.moe_ffn,
+ torch.bfloat16,
+ grouped=args.impl == "naive_grouped",
+ )
+ seq = args.tokens_per_rank // 4
+ x = torch.randn(seq, 4, args.hidden, dtype=torch.bfloat16, device="cuda", requires_grad=True)
+
+ ep_buffer = None
+ if args.impl == "te" and not args.ep_buffer_per_call:
+ ep_buffer = layer.mlp.make_ep_buffer()
+
+ def step():
+ layer.zero_grad(set_to_none=True)
+ x.grad = None
+ with _autocast(args.recipe):
+ out = layer(x, ep_buffer=ep_buffer)
+ out.backward(torch.ones_like(out))
+
+ for _ in range(args.warmup):
+ step()
+ torch.cuda.synchronize()
+ dist.barrier()
+
+ torch.cuda.profiler.start()
+ start = time.perf_counter()
+ for i in range(args.iters):
+ with torch.cuda.nvtx.range(f"iter{i}"):
+ step()
+ torch.cuda.synchronize()
+ ms = (time.perf_counter() - start) / args.iters * 1e3
+ torch.cuda.profiler.stop()
+ dist.barrier()
+
+ if rank == 0:
+ tok_s = args.tokens_per_rank * world_size / (ms / 1e3)
+ print(
+ f"DeepSeekV3Layer impl={args.impl}:"
+ f" ranks={world_size} experts={num_experts} topk={args.topk} tokens/rank={args.tokens_per_rank} hidden={args.hidden} recipe={args.recipe} fused_mlp={os.environ.get('NVTE_CUTEDSL_FUSED_GROUPED_MLP', '0')} ep_buffer={'per_call' if ep_buffer is None else 'reused'} fwd+bwd"
+ f" {ms:.3f} ms/iter ({tok_s / 1e6:.2f} Mtok/s)",
+ flush=True,
+ )
+ if args.impl == "te":
+ ep_finalize()
+ release_symm_mem_pool()
+ dist.destroy_process_group()
+ return 0
+
+
+if __name__ == "__main__":
+ sys.exit(main())
diff --git a/examples/pytorch/deepseek_v3/run_deepseek_v3_layer_ep.sh b/examples/pytorch/deepseek_v3/run_deepseek_v3_layer_ep.sh
new file mode 100755
index 00000000000..9341b7f14a6
--- /dev/null
+++ b/examples/pytorch/deepseek_v3/run_deepseek_v3_layer_ep.sh
@@ -0,0 +1,35 @@
+#!/bin/bash
+# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+#
+# See LICENSE for license information.
+#
+# Launcher for deepseek_v3_layer_ep.py on all local GPUs. Extra args go to the script:
+# bash run_deepseek_v3_layer_ep.sh # bf16, small dims
+# bash run_deepseek_v3_layer_ep.sh --dsv3 --recipe mxfp8 # DeepSeek-V3 dims, MXFP8 experts
+# NVTE_CUTEDSL_FUSED_GROUPED_MLP=1 bash run_deepseek_v3_layer_ep.sh --dsv3 --recipe mxfp8
+# NSYS=1 bash run_deepseek_v3_layer_ep.sh --dsv3 # nsys report in results/
+# Multi-node: run torchrun yourself with --nnodes/--rdzv-endpoint; EP spans all ranks.
+
+set -uo pipefail
+
+DETECTED_GPUS=$(nvidia-smi -L 2>/dev/null | wc -l)
+NUM_GPUS="${NUM_GPUS:-${DETECTED_GPUS}}"
+if [ "${NUM_GPUS}" -lt 2 ]; then
+ echo "EP requires >= 2 GPUs (found ${NUM_GPUS}); SKIPPING."
+ exit 0
+fi
+
+SCRIPT_DIR="$(cd "$(dirname "${BASH_SOURCE[0]}")" && pwd)"
+: ${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"
+
+PREFIX=()
+if [ "${NSYS:-0}" = "1" ]; then
+ mkdir -p "${SCRIPT_DIR}/results"
+ PREFIX=(nsys profile -t cuda,nvtx,nccl -c cudaProfilerApi --capture-range-end=stop
+ --cuda-graph-trace=node -o "${SCRIPT_DIR}/results/deepseek_v3_layer_ep_%h")
+fi
+
+"${PREFIX[@]}" torchrun --standalone --nnodes=1 --nproc-per-node="${NUM_GPUS}" \
+ "${SCRIPT_DIR}/deepseek_v3_layer_ep.py" "$@"
diff --git a/qa/L0_pytorch_unittest/test.sh b/qa/L0_pytorch_unittest/test.sh
index b3b6ccacac7..bcf453e3fde 100644
--- a/qa/L0_pytorch_unittest/test.sh
+++ b/qa/L0_pytorch_unittest/test.sh
@@ -38,6 +38,7 @@ PYTORCH_JIT=0 NVTE_TORCH_COMPILE=0 NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 NVTE_FUSED
PYTORCH_JIT=0 NVTE_TORCH_COMPILE=0 NVTE_ALLOW_NONDETERMINISTIC_ALGO=0 NVTE_FUSED_ATTN=0 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_cuda_graphs.xml $TE_PATH/tests/pytorch/test_cuda_graphs.py || test_fail "test_cuda_graphs.py"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_jit.xml $TE_PATH/tests/pytorch/test_jit.py || test_fail "test_jit.py"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_fused_rope.xml $TE_PATH/tests/pytorch/test_fused_rope.py || test_fail "test_fused_rope.py"
+python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_mla_rope.xml $TE_PATH/tests/pytorch/attention/test_mla_rope.py || test_fail "test_mla_rope.py"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_nvfp4.xml $TE_PATH/tests/pytorch/nvfp4 || test_fail "test_nvfp4"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_mxfp8.xml $TE_PATH/tests/pytorch/mxfp8 || test_fail "test_mxfp8"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_weight_swizzle_in_layers.xml $TE_PATH/tests/pytorch/test_weight_swizzle_in_layers.py || test_fail "test_weight_swizzle_in_layers.py"
@@ -59,6 +60,7 @@ python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_backward_overrid
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_permutation.xml $TE_PATH/tests/pytorch/test_permutation.py || test_fail "test_permutation.py"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_cross_entropy.xml $TE_PATH/tests/pytorch/test_cross_entropy.py || test_fail "test_cross_entropy.py"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_cpu_offloading.xml $TE_PATH/tests/pytorch/test_cpu_offloading.py || test_fail "test_cpu_offloading.py"
+python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_models.xml $TE_PATH/tests/pytorch/test_models.py || test_fail "test_models.py"
NVTE_FLASH_ATTN=0 NVTE_CPU_OFFLOAD_V1=1 python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_cpu_offloading_v1.xml $TE_PATH/tests/pytorch/test_cpu_offloading_v1.py || test_fail "test_cpu_offloading_v1.py"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_hybrid_quantization.xml $TE_PATH/tests/pytorch/test_hybrid_quantization.py || test_fail "test_hybrid_quantization.py"
python3 -m pytest --tb=auto --junitxml=$XML_LOG_DIR/pytest_test_identity_quantizer.xml $TE_PATH/tests/pytorch/test_identity_quantizer.py || test_fail "test_identity_quantizer.py"
diff --git a/qa/L1_pytorch_distributed_unittest/test.sh b/qa/L1_pytorch_distributed_unittest/test.sh
index f1de313fdce..68f242870c1 100644
--- a/qa/L1_pytorch_distributed_unittest/test.sh
+++ b/qa/L1_pytorch_distributed_unittest/test.sh
@@ -56,6 +56,7 @@ python3 -m pytest -v -s --junitxml=$XML_LOG_DIR/pytest_test_cu_seqlens_cache.xml
python3 -m pytest -v -s --junitxml=$XML_LOG_DIR/pytest_test_cast_master_weights_to_fp8.xml $TE_PATH/tests/pytorch/distributed/test_cast_master_weights_to_fp8.py || test_fail "test_cast_master_weights_to_fp8.py"
python3 -m pytest -v -s --junitxml=$XML_LOG_DIR/pytest_test_newton_schulz.xml $TE_PATH/tests/pytorch/distributed/test_newton_schulz.py || test_fail "test_newton_schulz.py"
python3 -m pytest -v -s --junitxml=$XML_LOG_DIR/pytest_test_ep.xml $TE_PATH/tests/pytorch/distributed/test_ep.py || test_fail "test_ep.py"
+python3 -m pytest -v -s --junitxml=$XML_LOG_DIR/pytest_test_models.xml $TE_PATH/tests/pytorch/distributed/test_models.py || test_fail "distributed/test_models.py"
# debug tests
diff --git a/tests/pytorch/attention/test_linear_mxfp8_attention.py b/tests/pytorch/attention/test_linear_mxfp8_attention.py
index e5cb1c8a655..f1926814f3a 100644
--- a/tests/pytorch/attention/test_linear_mxfp8_attention.py
+++ b/tests/pytorch/attention/test_linear_mxfp8_attention.py
@@ -35,7 +35,11 @@
_current_file = pathlib.Path(__file__).resolve()
sys.path = [str(_current_file.parent.parent)] + sys.path
from utils import ModelConfig, compare_and_assert, get_available_attention_backends
-from mla_rope_utils import apply_mla_rope, build_rope_tables
+from transformer_engine.pytorch.attention.mla_rope import (
+ apply_mla_rope_kv,
+ apply_mla_rope_q,
+ build_rope_tables,
+)
try:
@@ -179,6 +183,13 @@ def _run_projections(
return q_flat, kv_flat, q, kv, k_pos_emb
+def _apply_rope(q, kv, k_pos_emb, rope_tables):
+ cos, sin = rope_tables
+ q = apply_mla_rope_q(q, cos, sin, HEAD_DIM_NOPE, HEAD_DIM_ROPE)
+ k, v = apply_mla_rope_kv(kv, k_pos_emb, cos, sin, HEAD_DIM_NOPE, HEAD_DIM_ROPE, HEAD_DIM_V)
+ return q, k, v
+
+
def _run_forward_bf16(
modules: tuple,
x: torch.Tensor,
@@ -186,7 +197,7 @@ def _run_forward_bf16(
) -> tuple[torch.Tensor, torch.Tensor, torch.Tensor, torch.Tensor]:
q_proj, kv_proj, dpa, out_linear = modules
_, _, q, kv, k_pos_emb = _run_projections(q_proj, kv_proj, x)
- q, k, v = apply_mla_rope(q, kv, k_pos_emb, cos_table=rope_tables[0], sin_table=rope_tables[1])
+ q, k, v = _apply_rope(q, kv, k_pos_emb, rope_tables)
attn_out = dpa(q, k, v, qkv_format="sbhd")
return q, k, v, out_linear(attn_out.view(x.shape[0], x.shape[1], HIDDEN_SIZE))
@@ -208,13 +219,7 @@ def _run_forward_mxfp8(
x,
is_first_microbatch,
)
- q, k, v = apply_mla_rope(
- q,
- kv,
- k_pos_emb,
- cos_table=rope_tables[0],
- sin_table=rope_tables[1],
- )
+ q, k, v = _apply_rope(q, kv, k_pos_emb, rope_tables)
attn_out = dpa(q, k, v, qkv_format="sbhd")
out = out_linear(
attn_out.view(x.shape[0], x.shape[1], HIDDEN_SIZE),
@@ -288,7 +293,7 @@ def test_accuracy(self, batch_size: int, seq_len: int) -> None:
_set_seed()
baseline_modules, mxfp8_modules = _build_modules()
x = torch.randn(seq_len, batch_size, HIDDEN_SIZE, dtype=torch.bfloat16, device="cuda")
- rope_tables = build_rope_tables(seq_len, device=x.device)
+ rope_tables = build_rope_tables(seq_len, HEAD_DIM_ROPE, device=x.device)
q_bf16, k_bf16, v_bf16, out_bf16 = _run_forward_bf16(baseline_modules, x, rope_tables)
q_mxfp8, k_mxfp8, v_mxfp8, out_mxfp8 = _run_forward_mxfp8(
@@ -374,7 +379,7 @@ def test_backward(self, batch_size: int, seq_len: int) -> None:
device="cuda",
requires_grad=True,
)
- rope_tables = build_rope_tables(seq_len, device=x.device)
+ rope_tables = build_rope_tables(seq_len, HEAD_DIM_ROPE, device=x.device)
*_, out_mxfp8 = _run_forward_mxfp8(mxfp8_modules, x, fp8_recipe, rope_tables)
out_mxfp8.sum().backward()
@@ -408,7 +413,7 @@ def test_performance(self, batch_size: int, seq_len: int) -> None:
device="cuda",
requires_grad=True,
)
- rope_tables = build_rope_tables(seq_len, device=x.device)
+ rope_tables = build_rope_tables(seq_len, HEAD_DIM_ROPE, device=x.device)
mxfp8_fprop_ms, mxfp8_bprop_ms = _benchmark_training_step(
_run_forward_mxfp8, mxfp8_modules, x, fp8_recipe, rope_tables
diff --git a/tests/pytorch/attention/test_mla_rope.py b/tests/pytorch/attention/test_mla_rope.py
new file mode 100644
index 00000000000..2b35cec5cef
--- /dev/null
+++ b/tests/pytorch/attention/test_mla_rope.py
@@ -0,0 +1,184 @@
+# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+#
+# See LICENSE for license information.
+
+import math
+
+import pytest
+import torch
+
+from transformer_engine.pytorch.attention import mla_rope
+from transformer_engine.pytorch.module import LayerNormLinear
+
+
+@pytest.mark.parametrize("nope,rope,vdim", [(64, 32, 64), (48, 32, 64), (64, 48, 64), (64, 32, 48)])
+def test_mla_rope_matches_pytorch(nope, rope, vdim):
+ if not mla_rope.HAVE_TRITON:
+ pytest.skip("Triton unavailable")
+ s, b, h = 64, 2, 4
+ 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):
+ 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)
+ 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)
+
+
+@pytest.mark.parametrize("compiled", [False, True])
+def test_mla_rope_q_preserves_input_and_gradient(compiled):
+ if not mla_rope.HAVE_TRITON:
+ pytest.skip("Triton unavailable")
+ s, b, h, nope, rope = 8, 2, 4, 64, 32
+ cos, sin = mla_rope.build_rope_tables(s, rope, device="cuda")
+ q = torch.randn(s, b, h, nope + rope, device="cuda", requires_grad=True)
+ q_before = q.detach().clone()
+ incoming_grad = torch.randn_like(q)
+ grad_before = incoming_grad.clone()
+ apply_rope = lambda x: mla_rope.apply_mla_rope_q(x, cos, sin, nope, rope)
+ if compiled:
+ apply_rope = torch.compile(apply_rope, fullgraph=True)
+
+ aux = (q * q).sum()
+ rotated = apply_rope(q)
+ assert rotated.data_ptr() != q.data_ptr()
+ torch.autograd.backward((rotated, aux), (incoming_grad, torch.ones_like(aux)))
+
+ q_ref = q_before.requires_grad_()
+ rotated_ref = torch.cat(
+ (q_ref[..., :nope], mla_rope._rotate_interleaved_to_neox(q_ref[..., nope:], cos, sin, 0)),
+ dim=-1,
+ )
+ torch.autograd.backward(
+ (rotated_ref, (q_ref * q_ref).sum()), (grad_before, torch.ones_like(aux))
+ )
+ torch.testing.assert_close(q, q_before)
+ torch.testing.assert_close(incoming_grad, grad_before)
+ torch.testing.assert_close(rotated, rotated_ref)
+ torch.testing.assert_close(q.grad, q_ref.grad)
+
+
+def test_mla_rope_q_in_place_eager():
+ if not mla_rope.HAVE_TRITON:
+ pytest.skip("Triton unavailable")
+ s, b, h, nope, rope = 8, 2, 4, 64, 32
+ cos, sin = mla_rope.build_rope_tables(s, rope, device="cuda")
+ leaf = torch.randn(s, b, h, nope + rope, device="cuda", requires_grad=True)
+ q = (leaf * 1).view(s, b, h, nope + rope)
+ q_before = q.detach().clone()
+ incoming_grad = torch.randn_like(q)
+ grad_before = incoming_grad.clone()
+ rotated = mla_rope.apply_mla_rope_q(q, cos, sin, nope, rope, in_place=True)
+ assert rotated.data_ptr() == q.data_ptr()
+ torch.autograd.backward(rotated, incoming_grad)
+
+ q_ref = q_before.requires_grad_()
+ rotated_ref = torch.cat(
+ (q_ref[..., :nope], mla_rope._rotate_interleaved_to_neox(q_ref[..., nope:], cos, sin, 0)),
+ dim=-1,
+ )
+ torch.autograd.backward(rotated_ref, grad_before)
+ torch.testing.assert_close(rotated, rotated_ref)
+ torch.testing.assert_close(incoming_grad, grad_before)
+ torch.testing.assert_close(leaf.grad, q_ref.grad)
+
+
+def test_mla_rope_q_in_place_rejects_compile():
+ if not mla_rope.HAVE_TRITON:
+ pytest.skip("Triton unavailable")
+ s, b, h, nope, rope = 8, 2, 4, 64, 32
+ cos, sin = mla_rope.build_rope_tables(s, rope, device="cuda")
+ q = torch.randn(s, b, h, nope + rope, device="cuda")
+ compiled = torch.compile(
+ lambda x: mla_rope.apply_mla_rope_q(x, cos, sin, nope, rope, in_place=True),
+ fullgraph=True,
+ )
+ with pytest.raises(RuntimeError, match="in_place=True is not supported under torch.compile"):
+ compiled(q)
+
+
+def test_mla_rope_q_layernormlinear_view():
+ if not mla_rope.HAVE_TRITON:
+ pytest.skip("Triton unavailable")
+ s, b, h, nope, rope = 8, 2, 4, 64, 32
+ cos, sin = mla_rope.build_rope_tables(s, rope, device="cuda")
+ projection = LayerNormLinear(
+ 64, h * (nope + rope), normalization="RMSNorm", params_dtype=torch.float32, device="cuda"
+ )
+ x = torch.randn(s, b, 64, device="cuda", requires_grad=True)
+ q = projection(x).view(s, b, h, nope + rope)
+ q_before = q.detach().clone()
+ rotated = mla_rope.apply_mla_rope_q(q, cos, sin, nope, rope)
+ rotated.sum().backward()
+
+ rotated_ref = torch.cat(
+ (
+ q_before[..., :nope],
+ mla_rope._rotate_interleaved_to_neox(q_before[..., nope:], cos, sin, 0),
+ ),
+ dim=-1,
+ )
+ torch.testing.assert_close(q, q_before)
+ torch.testing.assert_close(rotated, rotated_ref)
+ assert x.grad is not None
+
+
+def test_rope_tables_yarn():
+ s, rope = 8192, 64
+ cos, sin = mla_rope.build_rope_tables(s, rope, device="cuda")
+ cos_none, sin_none = mla_rope.build_rope_tables(s, rope, device="cuda", scaling_factor=None)
+ assert torch.equal(cos, cos_none) and torch.equal(sin, sin_none)
+
+ yarn = dict(scaling_factor=40.0, original_max_position_embeddings=4096)
+ cos_y, sin_y = mla_rope.build_rope_tables(s, rope, device="cuda", **yarn)
+ factor = mla_rope.yarn_concentration_factor(40.0, 1.0, 0.0)
+ assert factor == pytest.approx(0.1 * math.log(40.0) + 1.0)
+ # amplitude scaled by the concentration factor
+ torch.testing.assert_close(cos_y**2 + sin_y**2, torch.full_like(cos_y, factor**2))
+ # high-frequency dims untouched, low-frequency dims interpolated by 1/scaling_factor
+ torch.testing.assert_close(cos_y[:, 0] / factor, cos[:, 0])
+ angle_y = torch.atan2(sin_y[:, rope // 2 - 1], cos_y[:, rope // 2 - 1])
+ angle = torch.atan2(sin[:, rope // 2 - 1], cos[:, rope // 2 - 1])
+ torch.testing.assert_close(angle_y[:64], angle[:64] / 40.0, atol=1e-4, rtol=0)
diff --git a/tests/pytorch/distributed/run_models.py b/tests/pytorch/distributed/run_models.py
new file mode 100644
index 00000000000..e38a45f2672
--- /dev/null
+++ b/tests/pytorch/distributed/run_models.py
@@ -0,0 +1,241 @@
+# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+#
+# See LICENSE for license information.
+"""Multi-process tests for model-specific layers (te.models), launched via torchrun."""
+
+import os
+import sys
+
+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, MultiLatentAttention
+
+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:
+ return DeepSeekV3MoE.ep_recv_capacity(ep_size, TOKENS_PER_RANK, TOP_K, NUM_LOCAL_EXPERTS)
+
+
+def _broadcast_params(module: torch.nn.Module) -> None:
+ for t in list(module.parameters()) + list(module.buffers()):
+ dist.broadcast(t.detach(), src=0)
+
+
+def _make_layer(ep_group, ep_size: int, num_experts: int) -> DeepSeekV3Layer:
+ ep = ep_group is not None
+ return DeepSeekV3Layer(
+ HIDDEN,
+ HEADS,
+ num_experts=num_experts,
+ moe_ffn_hidden_size=MOE_FFN,
+ topk=TOP_K,
+ shared_expert_ffn_hidden_size=SHARED_FFN,
+ params_dtype=DTYPE,
+ ep_group=ep_group,
+ ep_max_tokens_per_rank=TOKENS_PER_RANK if ep else None,
+ **MLA_KWARGS,
+ )
+
+
+def _copy_weights(ep_layer: DeepSeekV3Layer, ref: DeepSeekV3Layer, rank: int) -> None:
+ ref_params = dict(ref.named_parameters())
+ ref_bufs = dict(ref.named_buffers())
+ with torch.no_grad():
+ for name, p in ep_layer.named_parameters():
+ if not name.startswith("mlp.experts."):
+ p.copy_(ref_params[name])
+ for name, b in ep_layer.named_buffers():
+ if name in ref_bufs and b.shape == ref_bufs[name].shape:
+ b.copy_(ref_bufs[name])
+ ep_fc1, _, ep_fc2 = ep_layer.mlp.experts[1:4]
+ ref_fc1, _, ref_fc2 = ref.mlp.experts
+ for local_e in range(NUM_LOCAL_EXPERTS):
+ global_e = 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_layer_ep_matches_local(
+ rank: int, ep_size: int, ep_group, num_microbatches: int = 1
+) -> None:
+ """Full DeepSeekV3Layer with EP must match the all-experts-local layer numerically."""
+ num_experts = NUM_LOCAL_EXPERTS * ep_size
+ torch.manual_seed(0)
+ ref = _make_layer(None, ep_size, num_experts)
+ _broadcast_params(ref)
+ ep_layer = _make_layer(ep_group, ep_size, num_experts)
+ _copy_weights(ep_layer, ref, rank)
+
+ torch.manual_seed(1234 + rank)
+ microbatches = []
+ for mb in range(num_microbatches):
+ seq_len = TOKENS_PER_RANK // 2 - 8 * mb
+ x = torch.randn(seq_len, 2, HIDDEN, dtype=DTYPE, device="cuda")
+ x_ep = x.clone().requires_grad_(True)
+ x_ref = x.clone().requires_grad_(True)
+ out_ep = ep_layer(x_ep)
+ out_ref = ref(x_ref)
+ assert out_ep.shape == x.shape
+ torch.testing.assert_close(out_ep, out_ref, rtol=0.05, atol=0.05)
+ microbatches.append((x_ep, x_ref, out_ep, out_ref))
+
+ # All microbatches must retain their routing until their own backward.
+ for x_ep, x_ref, out_ep, out_ref in microbatches:
+ grad_out = torch.randn_like(out_ep)
+ out_ep.backward(grad_out.clone())
+ out_ref.backward(grad_out.clone())
+ torch.testing.assert_close(x_ep.grad, x_ref.grad, rtol=0.05, atol=0.05)
+
+ ref_params = dict(ref.named_parameters())
+ for name, p in ep_layer.named_parameters():
+ if name.startswith("mlp.experts.") or p.grad is None:
+ continue
+ torch.testing.assert_close(p.grad, ref_params[name].grad, rtol=0.1, atol=0.1, msg=name)
+
+ # A local expert's wgrad on its owner rank equals the sum of the
+ # 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_layer.mlp.experts[1:4]
+ ref_fc1, _, ref_fc2 = ref.mlp.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(num_experts)]
+ for g in ref_grads:
+ dist.all_reduce(g)
+ for local_e in range(NUM_LOCAL_EXPERTS):
+ global_e = rank * NUM_LOCAL_EXPERTS + local_e
+ ep_grad = getattr(ep_fc, f"weight{local_e}").grad.float()
+ torch.testing.assert_close(ep_grad, ref_grads[global_e], rtol=0.1, atol=0.1)
+
+ counts = ep_layer.mlp._last_tokens_per_expert.clone()
+ dist.all_reduce(counts, group=ep_group)
+ last_num_tokens = microbatches[-1][0].numel() // HIDDEN
+ assert counts.sum().item() == ep_size * last_num_tokens * TOP_K
+
+ expected_bias = ep_layer.mlp.expert_bias.clone()
+ expected_bias += ep_layer.mlp.expert_bias_update_rate * torch.sign(
+ counts.float().mean() - counts.float()
+ )
+ ep_layer.mlp._last_tokens_per_expert = counts
+ ep_layer.mlp.update_expert_bias()
+ torch.testing.assert_close(ep_layer.mlp.expert_bias, expected_bias, rtol=0, atol=0)
+ biases = [torch.empty_like(expected_bias) for _ in range(ep_size)]
+ dist.all_gather(biases, ep_layer.mlp.expert_bias, group=ep_group)
+ for bias in biases:
+ torch.testing.assert_close(bias, expected_bias, rtol=0, atol=0)
+
+
+def test_mla_tp_matches_local(rank: int, tp_size: int, tp_group) -> None:
+ """Compare sharded MLA outputs and gradients with an unsharded reference."""
+ torch.backends.cuda.matmul.allow_tf32 = False
+ for fmt in ("sbhd", "bshd"):
+ for explicit_size in (False, True):
+ torch.manual_seed(0)
+ common = dict(params_dtype=torch.float32, qkv_format=fmt, **MLA_KWARGS)
+ ref = MultiLatentAttention(HIDDEN, 2 * tp_size, **common)
+ _broadcast_params(ref)
+ tp = MultiLatentAttention(
+ HIDDEN,
+ 2 * tp_size,
+ tp_group=tp_group,
+ **({"tp_size": tp_size} if explicit_size else {}),
+ **common,
+ )
+ ref_params = dict(ref.named_parameters())
+
+ def shard(name, tensor):
+ if name in ("q_up_proj.weight", "kv_up_proj.weight"):
+ return tensor.chunk(tp_size, dim=0)[rank]
+ if name == "out_proj.weight":
+ return tensor.chunk(tp_size, dim=1)[rank]
+ return tensor
+
+ with torch.no_grad():
+ for name, param in tp.named_parameters():
+ param.copy_(shard(name, ref_params[name]))
+
+ torch.manual_seed(1234)
+ shape = (16, 2, HIDDEN) if fmt == "sbhd" else (2, 16, HIDDEN)
+ x = torch.randn(shape, device="cuda", requires_grad=True)
+ x_ref = x.detach().clone().requires_grad_()
+ out = tp(x)
+ out_ref = ref(x_ref)
+ torch.testing.assert_close(out, out_ref, rtol=1e-3, atol=1e-3)
+ grad = torch.randn_like(out)
+ out.backward(grad.clone())
+ out_ref.backward(grad.clone())
+ torch.testing.assert_close(x.grad, x_ref.grad, rtol=1e-3, atol=1e-3)
+ for name, param in tp.named_parameters():
+ torch.testing.assert_close(
+ param.grad,
+ shard(name, ref_params[name].grad),
+ rtol=1e-3,
+ atol=1e-3,
+ msg=name,
+ )
+
+
+def main() -> int:
+ dist.init_process_group(backend="nccl")
+ torch.cuda.set_device(int(os.environ["LOCAL_RANK"]))
+ if "--tp" in sys.argv:
+ test_mla_tp_matches_local(dist.get_rank(), dist.get_world_size(), dist.group.WORLD)
+ print(f"[rank {dist.get_rank()}] TP PASSED")
+ dist.destroy_process_group()
+ return 0
+ from torch.distributed import _symmetric_memory as _symm_mem
+
+ _symm_mem.set_backend("NCCL")
+
+ rank = dist.get_rank()
+ ep_size = dist.get_world_size()
+ if _device_sm() < 90:
+ if rank == 0:
+ print(f"NCCL EP requires SM>=90 (got SM{_device_sm()}); skipping.")
+ dist.destroy_process_group()
+ return 0
+
+ ep_group = dist.new_group(ranks=list(range(ep_size)), backend="nccl")
+ ep_bootstrap(
+ ep_group,
+ num_experts=NUM_LOCAL_EXPERTS * ep_size,
+ max_tokens_per_rank=TOKENS_PER_RANK,
+ hidden_dim=HIDDEN,
+ num_topk=TOP_K,
+ recv_capacity_per_rank=None if "--eager" in sys.argv else _recv_capacity(ep_size),
+ )
+ for num_microbatches in (1, 3):
+ test_layer_ep_matches_local(rank, ep_size, ep_group, num_microbatches)
+ print(f"[rank {rank}] PASSED")
+
+ dist.barrier()
+ ep_finalize()
+ release_symm_mem_pool()
+ dist.destroy_process_group()
+ return 0
+
+
+if __name__ == "__main__":
+ sys.exit(main())
diff --git a/tests/pytorch/distributed/test_models.py b/tests/pytorch/distributed/test_models.py
new file mode 100644
index 00000000000..d748b6bac76
--- /dev/null
+++ b/tests/pytorch/distributed/test_models.py
@@ -0,0 +1,46 @@
+# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+#
+# See LICENSE for license information.
+
+import os
+import subprocess
+from pathlib import Path
+
+import pytest
+import torch
+
+TEST_ROOT = Path(__file__).parent.resolve()
+NUM_PROCS = min(8, torch.cuda.device_count())
+LAUNCH_CMD = ["torchrun", f"--nproc_per_node={NUM_PROCS}"]
+
+
+def _has_nvlink() -> bool:
+ # NCCL EP falls back to the network transport and deadlocks on PCIe-only nodes.
+ out = subprocess.run(
+ ["nvidia-smi", "nvlink", "--status"], capture_output=True, text=True, check=False
+ ).stdout
+ return "GB/s" in out
+
+
+@pytest.mark.skipif(NUM_PROCS < 2, reason="EP requires >= 2 GPUs")
+@pytest.mark.skipif(not _has_nvlink(), reason="NCCL EP requires NVLink")
+@pytest.mark.parametrize("eager", [False, True], ids=["fixed", "eager"])
+def test_deepseek_layer_ep(eager):
+ result = subprocess.run(
+ LAUNCH_CMD + [str(TEST_ROOT / "run_models.py")] + (["--eager"] if eager else []),
+ env=os.environ,
+ check=False,
+ timeout=300,
+ )
+ assert result.returncode == 0
+
+
+@pytest.mark.skipif(NUM_PROCS < 2, reason="TP requires >= 2 GPUs")
+def test_mla_tp():
+ result = subprocess.run(
+ LAUNCH_CMD + [str(TEST_ROOT / "run_models.py"), "--tp"],
+ env=os.environ,
+ check=False,
+ timeout=300,
+ )
+ assert result.returncode == 0
diff --git a/tests/pytorch/test_float8_blockwise_gemm_exact.py b/tests/pytorch/test_float8_blockwise_gemm_exact.py
index eff571b5cd2..b00847c7602 100644
--- a/tests/pytorch/test_float8_blockwise_gemm_exact.py
+++ b/tests/pytorch/test_float8_blockwise_gemm_exact.py
@@ -8,6 +8,7 @@
import transformer_engine_torch as tex
from transformer_engine.pytorch.constants import TE_DType
+from transformer_engine.pytorch.cpp_extensions import general_grouped_gemm
from transformer_engine.pytorch import (
Float8BlockQuantizer,
get_device_compute_capability,
@@ -22,6 +23,55 @@ def fp8_blockwise_gemm_supported() -> bool:
return supported and not emulated
+@pytest.mark.parametrize("dtype", [torch.float32, torch.bfloat16], ids=str)
+@pytest.mark.parametrize("single_output", [False, True])
+def test_grouped_split_accumulator_enforced(dtype, single_output):
+ available, reason = te.is_fp8_block_scaling_available(return_reason=True)
+ if not available:
+ pytest.skip(reason)
+
+ torch.manual_seed(0)
+ quantizers = [
+ Float8BlockQuantizer(
+ fp8_dtype=tex.DType.kFloat8E4M3,
+ rowwise=True,
+ columnwise=False,
+ force_pow_2_scales=True,
+ block_scaling_dim=dim,
+ )
+ for dim in (1, 2)
+ ]
+ splits = [128, 256]
+ # Exactly representable operands isolate accumulation-mode selection from rounding.
+ inputs = [
+ quantizers[0](torch.randint(-4, 5, (rows, 128), device="cuda") / 8) for rows in splits
+ ]
+ weights = [quantizers[1](torch.randint(-4, 5, (128, 128), device="cuda") / 8) for _ in splits]
+ outputs = (
+ [torch.empty(sum(splits), 128, device="cuda", dtype=dtype)]
+ if single_output
+ else [torch.empty(rows, 128, device="cuda", dtype=dtype) for rows in splits]
+ )
+ general_grouped_gemm(
+ weights,
+ inputs,
+ outputs,
+ [None] * len(splits),
+ dtype,
+ m_splits=splits,
+ single_output=single_output,
+ use_split_accumulator=False,
+ )
+ actual = outputs[0] if single_output else torch.cat(outputs)
+ expected = torch.cat(
+ [
+ x.dequantize(dtype=torch.float32) @ weight.dequantize(dtype=torch.float32).T
+ for x, weight in zip(inputs, weights)
+ ]
+ ).to(dtype)
+ torch.testing.assert_close(actual, expected)
+
+
def cublas_gemm_fp8_blockwise_case(
x_dtype,
w_dtype,
diff --git a/tests/pytorch/test_grouped_tensor.py b/tests/pytorch/test_grouped_tensor.py
index 388b3bb7041..a38a80c9c00 100644
--- a/tests/pytorch/test_grouped_tensor.py
+++ b/tests/pytorch/test_grouped_tensor.py
@@ -5,6 +5,9 @@
"""Tests for GroupedTensor class"""
import os
+import subprocess
+import sys
+import textwrap
from types import SimpleNamespace
from typing import List, Optional, Tuple
@@ -50,6 +53,39 @@
)
+@pytest.mark.parametrize("block_scaling_dim", [1, 2])
+@pytest.mark.skipif(
+ not fp8_block_scaling_grouped_available, reason=reason_for_no_fp8_block_scaling_grouped
+)
+def test_group_quantize_fp8_blockwise_rejects_unaligned_rows(block_scaling_dim):
+ # A device error invalidates the CUDA context, so run it in a separate process.
+ code = textwrap.dedent(
+ f"""
+ import torch
+ from transformer_engine.pytorch import Float8BlockQuantizer
+ import transformer_engine_torch as tex
+
+ x = torch.ones(128, 128, dtype=torch.bfloat16, device="cuda")
+ splits = torch.tensor([16, 112], dtype=torch.int64, device="cuda")
+ quantizer = Float8BlockQuantizer(
+ fp8_dtype=tex.DType.kFloat8E4M3,
+ rowwise=True,
+ columnwise=False,
+ block_scaling_dim={block_scaling_dim},
+ )
+ tex.group_quantize(x, quantizer, 2, splits)
+ torch.cuda.synchronize()
+ """
+ )
+ result = subprocess.run(
+ [sys.executable, "-c", code], capture_output=True, text=True, check=False, timeout=120
+ )
+ output = result.stdout + result.stderr
+ assert result.returncode != 0, output
+ assert "multiple of 128" in output, output
+ assert "CUDA error" in output, output
+
+
def test_mark_grouped_tensor_supports_plain_tensor():
tensor = torch.empty(16)
diff --git a/tests/pytorch/test_models.py b/tests/pytorch/test_models.py
new file mode 100644
index 00000000000..9d677be2fe2
--- /dev/null
+++ b/tests/pytorch/test_models.py
@@ -0,0 +1,194 @@
+# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+#
+# See LICENSE for license information.
+
+import math
+
+import pytest
+import torch
+
+import transformer_engine.pytorch as te
+from transformer_engine.pytorch.ops.fused.grouped_mlp import (
+ GroupedMLP_CuTeGEMMGLU,
+ fuse_glu_ops,
+)
+from transformer_engine.pytorch.ops.fuser import OperationFuser
+from transformer_engine.pytorch.utils import deinterleave_glu_tensor
+from transformer_engine.pytorch.models import DeepSeekV3MoE, MultiLatentAttention
+from utils import make_recipe, quantization_tols
+
+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
+ )
+
+
+@pytest.mark.parametrize("mscale_all_dim", [0.0, 1.0])
+def test_mla_yarn_softmax_scale(mscale_all_dim):
+ mla = MultiLatentAttention(
+ HIDDEN,
+ HEADS,
+ params_dtype=DTYPE,
+ rope_scaling_factor=40.0,
+ original_max_position_embeddings=64,
+ mscale_all_dim=mscale_all_dim,
+ **MLA_KWARGS,
+ )
+ m = 0.1 * mscale_all_dim * math.log(40.0) + 1.0
+ qk_head_dim = MLA_KWARGS["qk_nope_head_dim"] + MLA_KWARGS["qk_rope_head_dim"]
+ assert mla.softmax_scale == pytest.approx(m * m / math.sqrt(qk_head_dim))
+
+
+@pytest.mark.parametrize("shared", [False, True], ids=["no_shared", "shared"])
+@pytest.mark.parametrize("grouped", [False, True], ids=["ungrouped", "grouped"])
+@pytest.mark.parametrize("topk", [2, 4])
+@pytest.mark.parametrize("quantization", [None, "fp8_block_scaling"])
+def test_moe_matches_dense_reference(shared, grouped, topk, quantization):
+ """Routed output must equal the prob-weighted sum of the selected expert MLPs."""
+ if quantization is not None:
+ available, reason = te.is_fp8_block_scaling_available(return_reason=True)
+ if not available:
+ pytest.skip(reason)
+ torch.manual_seed(0)
+ num_experts = 4
+ moe = DeepSeekV3MoE(
+ HIDDEN,
+ moe_ffn_hidden_size=128,
+ num_experts=num_experts,
+ topk=topk,
+ num_groups=2 if grouped else None,
+ group_topk=topk // 2 if grouped else None,
+ shared_expert_ffn_hidden_size=128 if shared else None,
+ params_dtype=DTYPE,
+ )
+ x = _input()
+ with te.autocast(enabled=quantization is not None, recipe=make_recipe(quantization)):
+ out = moe(x)
+ assert out.shape == x.shape
+ assert torch.isfinite(out).all()
+ out.sum().backward()
+ assert torch.isfinite(x.grad).all()
+ for name, parameter in moe.named_parameters():
+ if parameter.requires_grad:
+ assert parameter.grad is not None, name
+ assert torch.isfinite(parameter.grad).all(), name
+
+ tokens = x.detach().reshape(-1, HIDDEN)
+ probs, _ = moe._route(moe.gate(tokens).float())
+ assert (probs > 0).sum(dim=1).eq(topk).all()
+ assert moe._last_tokens_per_expert.sum().item() == tokens.shape[0] * topk
+
+ fc1, _, fc2 = moe.experts
+ quantizers = [
+ te.Float8BlockQuantizer(
+ fp8_dtype=te.DType.kFloat8E4M3,
+ rowwise=True,
+ columnwise=False,
+ block_scaling_dim=dim,
+ )
+ for dim in (1, 2)
+ ]
+
+ def qdq(tensor, block_dim):
+ if quantization is None:
+ return tensor
+ return quantizers[block_dim - 1](tensor).dequantize(dtype=DTYPE)
+
+ ref_tokens = qdq(tokens, 1)
+ ref = torch.zeros_like(tokens)
+ for e in range(num_experts):
+ w1 = deinterleave_glu_tensor(qdq(getattr(fc1, f"weight{e}"), 2), 32)
+ w2 = qdq(getattr(fc2, f"weight{e}"), 2)
+ gate_part, lin_part = (ref_tokens @ w1.t()).chunk(2, dim=-1)
+ act = torch.nn.functional.silu(gate_part.float()) * lin_part.float()
+ act = qdq(act.to(DTYPE) * probs[:, e : e + 1].to(DTYPE), 1)
+ ref += act @ w2.t()
+ if shared:
+ with te.autocast(enabled=quantization is not None, recipe=make_recipe(quantization)):
+ ref += moe.shared_expert(tokens)
+ torch.testing.assert_close(out.reshape(-1, HIDDEN), ref, rtol=0.05, atol=0.05)
+
+ bias_before = moe.expert_bias.clone()
+ moe.update_expert_bias()
+ assert torch.isfinite(moe.expert_bias).all()
+ if topk < num_experts:
+ assert not torch.equal(bias_before, moe.expert_bias)
+
+
+@pytest.mark.parametrize("quantization", ["mxfp8", "nvfp4_rht"])
+def test_moe_fused_quantized_uneven_expert_rows(monkeypatch, quantization):
+ available_fn = te.is_mxfp8_available if quantization == "mxfp8" else te.is_nvfp4_available
+ available, reason = available_fn(return_reason=True)
+ if not available:
+ pytest.skip(reason)
+
+ recipe = make_recipe(quantization)
+ torch.manual_seed(0)
+ moe = DeepSeekV3MoE(HIDDEN, 128, num_experts=2, topk=1, params_dtype=DTYPE)
+ fused_ops = fuse_glu_ops(list(moe.experts), recipe=recipe)
+ if (
+ fuse_glu_ops not in OperationFuser.forward_backward_fusion_functions
+ or len(fused_ops) != 1
+ or not isinstance(fused_ops[0], GroupedMLP_CuTeGEMMGLU)
+ ):
+ pytest.skip("requires the fused grouped MLP")
+ reference = DeepSeekV3MoE(HIDDEN, 128, num_experts=2, topk=1, params_dtype=DTYPE)
+ reference.load_state_dict(moe.state_dict())
+ with torch.no_grad():
+ for module in (moe, reference):
+ module.gate.weight.zero_()
+ module.gate.weight[0, 0] = 1
+ module.gate.weight[1, 0] = -1
+
+ tokens = torch.empty(512, HIDDEN, device="cuda", dtype=DTYPE).uniform_(-0.25, 0.25)
+ tokens[:128, 0] = 2
+ tokens[128:, 0] = -2
+ x = tokens.detach().requires_grad_()
+ x_ref = tokens.detach().clone().requires_grad_()
+ grad = torch.empty_like(tokens).uniform_(-0.25, 0.25)
+ splits = []
+ moe.experts.register_forward_pre_hook(lambda _module, args: splits.append(args[1].clone()))
+
+ with te.autocast(enabled=True, recipe=recipe):
+ out = moe(x)
+ out.backward(grad.clone())
+ assert torch.equal(splits[0], torch.tensor([256, 512], device="cuda"))
+ fused_op = moe.experts._module_groups[0]._forward_ops[0][0]
+ assert isinstance(fused_op, GroupedMLP_CuTeGEMMGLU)
+
+ monkeypatch.setattr(
+ OperationFuser,
+ "forward_backward_fusion_functions",
+ [fn for fn in OperationFuser.forward_backward_fusion_functions if fn is not fuse_glu_ops],
+ )
+ ref_splits = []
+ reference.experts.register_forward_pre_hook(
+ lambda _module, args: ref_splits.append(args[1].clone())
+ )
+ with te.autocast(enabled=True, recipe=make_recipe(quantization)):
+ ref_out = reference(x_ref)
+ ref_out.backward(grad.clone())
+ assert torch.equal(ref_splits[0], splits[0])
+
+ tols = quantization_tols(quantization)
+ torch.testing.assert_close(out, ref_out, **tols)
+ torch.testing.assert_close(x.grad, x_ref.grad, **tols)
+ ref_params = dict(reference.named_parameters())
+ for name, param in moe.named_parameters():
+ torch.testing.assert_close(param.grad, ref_params[name].grad, **tols)
diff --git a/tests/pytorch/test_sanity.py b/tests/pytorch/test_sanity.py
index 7db3002427c..f60687821b5 100644
--- a/tests/pytorch/test_sanity.py
+++ b/tests/pytorch/test_sanity.py
@@ -35,6 +35,7 @@
is_bf16_available,
)
from transformer_engine.common import recipe
+from transformer_engine.pytorch.models import DeepSeekV3Layer
from transformer_engine.pytorch.cpp_extensions import general_gemm
from transformer_engine.pytorch.tensor.utils import replace_raw_data
from transformer_engine.pytorch.module import (
@@ -285,7 +286,7 @@ def _test_sanity_e2e_gradient_accumulation_fusion(block, dtype, config, fp8_reci
), f"grad_added_to_main_grad not set to True for {failed_grad_added_flags}."
-def _test_sanity_e2e(block, dtype, config, fp8_recipe, skip_wgrad):
+def _test_sanity_e2e(block, dtype, config, fp8_recipe, skip_wgrad, *, check_finite=False):
te_inp_hidden_states = torch.randn(
(config.max_seqlen_q, config.batch_size, config.hidden_size),
dtype=dtype,
@@ -303,6 +304,14 @@ def _test_sanity_e2e(block, dtype, config, fp8_recipe, skip_wgrad):
loss.backward()
torch.cuda.synchronize()
+ if check_finite:
+ assert torch.isfinite(te_out).all(), "Non-finite output"
+ assert torch.isfinite(te_inp_hidden_states.grad).all(), "Non-finite input gradient"
+ for name, parameter in block.named_parameters():
+ if parameter.requires_grad:
+ assert parameter.grad is not None, name
+ assert torch.isfinite(parameter.grad).all(), name
+
def _test_sanity_e2e_bert(block, dtype, config, fp8_recipe, skip_wgrad):
te_inp_hidden_states = torch.randn(
@@ -860,6 +869,41 @@ def checked_norm(*args, _norm=norm_, _suffix=suffix, **kwargs):
assert set(seen_norm_stages) == {"fwd", "bwd"}
+@pytest.mark.parametrize("dtype", param_types)
+@pytest.mark.parametrize("fp8_recipe", fp8_recipes, ids=recipe_id)
+@pytest.mark.parametrize("moe", all_boolean, ids=["moe", "dense"])
+def test_sanity_deepseek_v3_layer(dtype, fp8_recipe, moe):
+ config = model_configs["small"]
+
+ if fp8_recipe is not None:
+ if not is_fp8_supported(config):
+ pytest.skip("Model config does not support FP8")
+ if fp8_recipe.nvfp4() and dtype == torch.float16:
+ pytest.skip("FP16 output for NVFP4 not supported")
+ if fp8_recipe.nvfp4() and moe and dtype != torch.bfloat16:
+ pytest.skip("NVFP4 GroupedLinear requires BF16")
+
+ mlp_kwargs = (
+ dict(num_experts=4, topk=2, moe_ffn_hidden_size=64, shared_expert_ffn_hidden_size=64)
+ if moe
+ else dict(ffn_hidden_size=4 * config.hidden_size)
+ )
+ block = DeepSeekV3Layer(
+ config.hidden_size,
+ config.num_heads,
+ q_lora_rank=64,
+ kv_lora_rank=64,
+ qk_nope_head_dim=32,
+ qk_rope_head_dim=32,
+ v_head_dim=32,
+ params_dtype=dtype,
+ device="cuda",
+ **mlp_kwargs,
+ )
+
+ _test_sanity_e2e(block, dtype, config, fp8_recipe, skip_wgrad=False, check_finite=True)
+
+
@pytest.mark.parametrize("dtype", param_types)
@pytest.mark.parametrize("fp8_recipe", fp8_recipes, ids=recipe_id)
@pytest.mark.parametrize("model", ["small"])
diff --git a/transformer_engine/common/cast/fp8_blockwise/group_quantize_fp8_blockwise.cuh b/transformer_engine/common/cast/fp8_blockwise/group_quantize_fp8_blockwise.cuh
index 31feaf833dc..68eeed1f09c 100644
--- a/transformer_engine/common/cast/fp8_blockwise/group_quantize_fp8_blockwise.cuh
+++ b/transformer_engine/common/cast/fp8_blockwise/group_quantize_fp8_blockwise.cuh
@@ -145,6 +145,8 @@ __device__ __forceinline__ size_t find_tensor_id_by_block_y(
NVTE_DEVICE_ERROR(
"Grouped FP8 block-scaling quantize: each tensor's first dimension must be a "
"multiple of 128 (VARYING_FIRST_DIM).");
+ // Device asserts may be disabled in release builds.
+ __trap();
}
}
}
diff --git a/transformer_engine/pytorch/__init__.py b/transformer_engine/pytorch/__init__.py
index 09f2a068332..8be5e2dc36a 100644
--- a/transformer_engine/pytorch/__init__.py
+++ b/transformer_engine/pytorch/__init__.py
@@ -37,6 +37,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/tests/pytorch/attention/mla_rope_utils.py b/transformer_engine/pytorch/attention/mla_rope.py
similarity index 55%
rename from tests/pytorch/attention/mla_rope_utils.py
rename to transformer_engine/pytorch/attention/mla_rope.py
index 90eebfc66ae..fb44692f327 100644
--- a/tests/pytorch/attention/mla_rope_utils.py
+++ b/transformer_engine/pytorch/attention/mla_rope.py
@@ -2,16 +2,18 @@
#
# See LICENSE for license information.
-"""MLA RoPE for DSv3 671B - Triton forward and backward kernels.
+"""Fused MLA RoPE kernels (DeepSeekV3-style decoupled RoPE/NoPE).
-Source: Megatron-LM megatron/core/fusions/fused_mla_yarn_rope_apply.py
-Falls back to pure PyTorch when Triton is unavailable.
+The query kernel rotates the trailing ``head_dim_rope`` slice; the KV
+kernel builds the 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.
-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.
-"""
+The rope slice is read interleaved (checkpoint layout) and written in NeoX
+half-split layout."""
+
+import math
+from typing import Optional, Tuple
import torch
@@ -23,25 +25,101 @@
except ImportError:
HAVE_TRITON = False
-HEAD_DIM_ROPE = 64
-HEAD_DIM_NOPE = 128
-HEAD_DIM_V = 128
-ROTARY_BASE = 10000
+__all__ = [
+ "build_rope_tables",
+ "apply_mla_rope_q",
+ "apply_mla_rope_kv",
+ "yarn_mscale",
+ "yarn_concentration_factor",
+]
+
+
+def _yarn_correction_dim(num_rotations, dim, base, max_pos):
+ return (dim * math.log(max_pos / (num_rotations * 2 * math.pi))) / (2 * math.log(base))
+
+
+def _yarn_correction_range(beta_fast, beta_slow, dim, base, max_pos, round_to_int=True):
+ low = _yarn_correction_dim(beta_fast, dim, base, max_pos)
+ high = _yarn_correction_dim(beta_slow, dim, base, max_pos)
+ if round_to_int:
+ low, high = math.floor(low), math.ceil(high)
+ return max(low, 0), min(high, dim - 1)
+
+
+def _yarn_linear_ramp(low, high, dim, device):
+ if low == high:
+ high += 0.001
+ ramp = (torch.arange(dim, dtype=torch.float32, device=device) - low) / (high - low)
+ return torch.clamp(ramp, 0, 1)
+
+
+def yarn_mscale(scale: float, mscale: float = 1.0) -> float:
+ """YaRN attention temperature factor ``0.1 * mscale * ln(scale) + 1`` (1 for scale <= 1)."""
+ if scale <= 1:
+ return 1.0
+ return 0.1 * mscale * math.log(scale) + 1.0
+
+
+def yarn_concentration_factor(scaling_factor: float, mscale: float, mscale_all_dim: float) -> float:
+ """Factor multiplied into cos/sin tables."""
+ return yarn_mscale(scaling_factor, mscale) / yarn_mscale(scaling_factor, mscale_all_dim)
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)
- )
+ emb_dim: int,
+ base: float = 10000.0,
+ device: Optional[torch.device] = None,
+ scaling_factor: Optional[float] = None,
+ original_max_position_embeddings: int = 4096,
+ beta_fast: float = 32.0,
+ beta_slow: float = 1.0,
+ mscale: float = 1.0,
+ mscale_all_dim: float = 0.0,
+) -> Tuple[torch.Tensor, torch.Tensor]:
+ """Build cosine/sine tables for :func:`apply_mla_rope_q` and :func:`apply_mla_rope_kv`.
+
+ Parameters
+ ----------
+ seq_len : int
+ Number of sequence positions.
+ emb_dim : int
+ Positive, even number of rotary channels; excludes the non-rotary channels.
+ base : float, default = 10000.0
+ Base for the rotary frequencies.
+ device : torch.device, optional
+ Table device; ``None`` uses the default PyTorch device.
+ scaling_factor : float, optional
+ YaRN context-extension factor; ``None`` selects unscaled RoPE.
+ original_max_position_embeddings : int, default = 4096
+ Original context length used to determine the YaRN frequency ramp.
+ beta_fast, beta_slow : float, default = 32.0, 1.0
+ Rotation-count boundaries of the YaRN interpolation ramp.
+ mscale, mscale_all_dim : float, default = 1.0, 0.0
+ YaRN magnitude coefficients. Tables are multiplied by the ratio of
+ ``yarn_mscale(scaling_factor, mscale)`` to
+ ``yarn_mscale(scaling_factor, mscale_all_dim)``.
+
+ Returns
+ -------
+ tuple of torch.Tensor
+ Contiguous FP32 ``(cos_table, sin_table)``, each of shape
+ ``[seq_len, emb_dim]`` with duplicated halves for NeoX layout.
+ """
+ exponent = torch.arange(0, emb_dim, 2, dtype=torch.float32, device=device) / emb_dim
+ inv_freq = 1.0 / (base**exponent)
+ factor = 1.0
+ if scaling_factor is not None:
+ low, high = _yarn_correction_range(
+ beta_fast, beta_slow, emb_dim, base, original_max_position_embeddings
+ )
+ extra_mask = 1.0 - _yarn_linear_ramp(low, high, emb_dim // 2, device)
+ inv_freq = (inv_freq / scaling_factor) * (1 - extra_mask) + inv_freq * extra_mask
+ factor = yarn_concentration_factor(scaling_factor, mscale, mscale_all_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()
+ return (torch.cos(freqs) * factor).contiguous(), (torch.sin(freqs) * factor).contiguous()
if HAVE_TRITON:
@@ -69,20 +147,9 @@ def _get_thd_token_idx(cu_seqlens, pid_m, seq_num, cp_rank, cp_size):
) * 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"],
- )
+ _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,
@@ -100,6 +167,7 @@ def rotary_fwd_q_kernel(
cp_size,
BLOCK_H: tl.constexpr,
):
+ """In-place RoPE fwd on the trailing rope slice of q."""
pid_m = tl.program_id(axis=0)
pid_head = tl.program_id(axis=1)
if cu_seqlens_q is None:
@@ -129,20 +197,7 @@ def rotary_fwd_q_kernel(
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.autotune(configs=_AUTOTUNE_CONFIGS, key=["emb_dim", "head_num"], restore_value=["DO"])
@triton.jit
def rotary_bwd_q_kernel(
DO,
@@ -160,6 +215,7 @@ def rotary_bwd_q_kernel(
cp_size,
BLOCK_H: tl.constexpr,
):
+ """In-place RoPE bwd on the trailing rope slice of dq."""
pid_m = tl.program_id(axis=0)
pid_head = tl.program_id(axis=1)
if cu_seqlens_q is None:
@@ -189,19 +245,7 @@ def rotary_bwd_q_kernel(
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.autotune(configs=_AUTOTUNE_CONFIGS, key=["emb_dim", "k_dim", "v_dim", "head_num"])
@triton.jit
def rotary_fwd_kv_kernel(
KV,
@@ -228,6 +272,7 @@ def rotary_fwd_kv_kernel(
cp_size,
BLOCK_H: tl.constexpr,
):
+ """Fwd: build (key, value) from kv and the shared rotated rope head."""
pid_m = tl.program_id(axis=0)
pid_head = tl.program_id(axis=1)
if cu_seqlens_kv is None:
@@ -268,19 +313,7 @@ def rotary_fwd_kv_kernel(
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.autotune(configs=_AUTOTUNE_CONFIGS, key=["emb_dim", "k_dim", "v_dim", "head_num"])
@triton.jit
def rotary_bwd_kv_kernel(
dK,
@@ -307,6 +340,7 @@ def rotary_bwd_kv_kernel(
cp_size,
BLOCK_H: tl.constexpr,
):
+ """Bwd: scatter (dk, dv) into dkv and reduce rope-slice grads into demb."""
pid_m = tl.program_id(axis=0)
pid_head = tl.program_id(axis=1)
if cu_seqlens_kv is None:
@@ -357,19 +391,25 @@ def rotary_bwd_kv_kernel(
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)
+ def _token_stride(tensor: torch.Tensor) -> int:
+ return tensor.stride(1) if tensor.dim() == 4 else tensor.stride(0)
class _MLARoPEQTriton(torch.autograd.Function):
+ """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):
+ def forward(ctx, q, cos, sin, head_dim_nope, head_dim_rope, in_place):
+ """Rotate the rope slice of q."""
+ if in_place:
+ ctx.mark_dirty(q)
+ else:
+ q = q.clone(memory_format=torch.contiguous_format)
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](
+ def grid(meta):
+ return (s * b, triton.cdiv(nheads, meta["BLOCK_H"]))
+
+ rotary_fwd_q_kernel[grid](
q,
cos,
sin,
@@ -379,54 +419,58 @@ def forward(ctx, q, cos, sin, head_dim_nope, head_dim_rope):
b,
None,
None,
- _flattened_token_stride(q),
+ _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
+ ctx.dims = (s, b, nheads, head_dim_nope, head_dim_rope)
return q
@staticmethod
def backward(ctx, dq):
+ """Counter-rotate the rope slice of dq."""
cos, sin = ctx.saved_tensors
- s, b, nheads = ctx.s, ctx.b, ctx.nheads
- total = s * b
+ dq = dq.clone(memory_format=torch.contiguous_format)
+ s, b, nheads, head_dim_nope, head_dim_rope = ctx.dims
+
+ def grid(meta):
+ return (s * b, triton.cdiv(nheads, meta["BLOCK_H"]))
- grid_q = lambda META: (total, triton.cdiv(nheads, META["BLOCK_H"]))
- rotary_bwd_q_kernel[grid_q](
+ rotary_bwd_q_kernel[grid](
dq,
cos,
sin,
- ctx.head_dim_nope,
- ctx.head_dim_rope,
+ head_dim_nope,
+ head_dim_rope,
nheads,
b,
None,
None,
- _flattened_token_stride(dq),
+ _token_stride(dq),
dq.stride(2),
0,
1,
)
- return dq, None, None, None, None
+ return dq, None, 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):
+ """Build (k, v) from kv and the shared rope head."""
+ if not kv.is_contiguous():
+ kv = kv.contiguous()
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](
+
+ def grid(meta):
+ return (s * b, triton.cdiv(nheads, meta["BLOCK_H"]))
+
+ rotary_fwd_kv_kernel[grid](
kv,
k_pos_emb,
o_key,
@@ -440,37 +484,34 @@ def forward(ctx, kv, k_pos_emb, cos, sin, head_dim_nope, head_dim_rope, head_dim
b,
None,
None,
- _flattened_token_stride(kv),
+ _token_stride(kv),
kv.stride(2),
- _flattened_token_stride(k_pos_emb),
- _flattened_token_stride(o_key),
+ _token_stride(k_pos_emb),
+ _token_stride(o_key),
o_key.stride(2),
- _flattened_token_stride(o_value),
+ _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
+ 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):
+ """Gradients for (kv, k_pos_emb) from (dk, dv)."""
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
-
+ 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_kv = lambda META: (total, triton.cdiv(nheads, META["BLOCK_H"]))
- rotary_bwd_kv_kernel[grid_kv](
+
+ def grid(meta):
+ return (s * b, triton.cdiv(nheads, meta["BLOCK_H"]))
+
+ rotary_bwd_kv_kernel[grid](
dk_out,
dv_out,
d_kv,
@@ -484,185 +525,156 @@ def backward(ctx, dk_out, dv_out):
b,
None,
None,
- _flattened_token_stride(dk_out),
+ _token_stride(dk_out),
dk_out.stride(2),
- _flattened_token_stride(dv_out),
+ _token_stride(dv_out),
dv_out.stride(2),
- _flattened_token_stride(d_kv),
+ _token_stride(d_kv),
d_kv.stride(2),
- _flattened_token_stride(d_emb),
+ _token_stride(d_emb),
0,
1,
)
return d_kv, d_emb, None, None, None, None, None
-def _apply_mla_rope_q_with_tables(
+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_NOPE,
- head_dim_rope: int = HEAD_DIM_ROPE,
+ head_dim_nope: int,
+ head_dim_rope: int,
+ tensor_format: str = "sbhd",
+ in_place: bool = False,
) -> torch.Tensor:
- if HAVE_TRITON:
+ """Rotate the trailing query channels, preserving the non-rotary prefix.
+
+ Parameters
+ ----------
+ q : torch.Tensor
+ CUDA query tensor of shape ``[s, b, h, head_dim_nope + head_dim_rope]``
+ for ``sbhd``, or ``[b, s, h, head_dim_nope + head_dim_rope]`` for ``bshd``.
+ cos_table, sin_table : torch.Tensor
+ Contiguous FP32 tables of shape ``[s, head_dim_rope]`` on the same device
+ as ``q``, as returned by :func:`build_rope_tables`.
+ head_dim_nope, head_dim_rope : int
+ Number of non-rotary and rotary channels per head.
+ tensor_format : {"sbhd", "bshd"}, default = "sbhd"
+ Input and output layout.
+ in_place : bool, default = False
+ Rotate ``q`` in place instead of allocating an output. Requires a contiguous
+ ``sbhd`` input and a power-of-two rotary dimension of at least two.
+ Supported in eager mode only.
+
+ Returns
+ -------
+ torch.Tensor
+ Rotated queries with the same shape and dtype as ``q``.
+
+ Notes
+ -----
+ Triton is used for ``sbhd`` when ``head_dim_rope`` is a power of two and
+ at least two; other supported layouts/dimensions use PyTorch operations.
+ The Triton path computes gradients for ``q`` only; treat the tables as constants.
+
+ The default path leaves ``q`` and incoming backward gradients unchanged.
+ With ``in_place=True``, ``q`` must have no other consumers that need its
+ original values, and PyTorch's usual in-place autograd rules apply. In
+ particular, views of custom autograd Function outputs cannot be mutated.
+ """
+ use_triton = (
+ HAVE_TRITON
+ and tensor_format == "sbhd"
+ and head_dim_rope >= 2
+ and _is_power_of_two(head_dim_rope)
+ )
+ if in_place and (not use_triton or not q.is_contiguous()):
+ raise ValueError("in_place=True requires contiguous sbhd input and Triton RoPE support")
+ if in_place and torch.compiler.is_compiling():
+ raise RuntimeError("in_place=True is not supported under torch.compile")
+ if use_triton:
return _MLARoPEQTriton.apply(
- q,
- cos_table,
- sin_table,
- head_dim_nope,
- head_dim_rope,
+ q, cos_table, sin_table, head_dim_nope, head_dim_rope, in_place
)
- return _apply_pytorch_q(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_with_tables(
+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_NOPE,
- head_dim_rope: int = HEAD_DIM_ROPE,
- head_dim_v: int = HEAD_DIM_V,
-) -> tuple[torch.Tensor, torch.Tensor]:
- if HAVE_TRITON:
+ head_dim_nope: int,
+ head_dim_rope: int,
+ head_dim_v: int,
+ tensor_format: str = "sbhd",
+) -> Tuple[torch.Tensor, torch.Tensor]:
+ """Assemble keys and values using a shared rotary key head.
+
+ Parameters
+ ----------
+ kv : torch.Tensor
+ CUDA tensor of shape ``[s, b, h, head_dim_nope + head_dim_v]`` for
+ ``sbhd``, or ``[b, s, h, head_dim_nope + head_dim_v]`` for ``bshd``.
+ Each head contains the non-rotary key channels followed by value channels.
+ k_pos_emb : torch.Tensor
+ Shared rotary key channels of shape ``[s, b, 1, head_dim_rope]`` for
+ ``sbhd``, or ``[b, s, 1, head_dim_rope]`` for ``bshd``, matching the
+ device and dtype of ``kv``. Triton requires a contiguous last dimension
+ and ``stride(0) == b * stride(1)``.
+ cos_table, sin_table : torch.Tensor
+ Contiguous FP32 tables of shape ``[s, head_dim_rope]`` on the same device
+ as ``kv``, as returned by :func:`build_rope_tables`.
+ head_dim_nope, head_dim_rope, head_dim_v : int
+ Number of non-rotary key, rotary key and value channels per head.
+ tensor_format : {"sbhd", "bshd"}, default = "sbhd"
+ Input and output layout.
+
+ Returns
+ -------
+ tuple of torch.Tensor
+ Contiguous ``(k, v)`` in the selected layout and input dtype. The last
+ dimensions are ``head_dim_nope + head_dim_rope`` for ``k`` and
+ ``head_dim_v`` for ``v``. The rotated shared key head is broadcast to all heads.
+
+ Notes
+ -----
+ Triton is used for ``sbhd`` when all three head dimensions are powers of two
+ and ``head_dim_rope >= 2``; other supported layouts/dimensions use PyTorch.
+ Inputs are preserved. The Triton path computes gradients for ``kv`` and
+ ``k_pos_emb`` only; treat the tables as constants.
+ """
+ if (
+ HAVE_TRITON
+ and tensor_format == "sbhd"
+ and head_dim_rope >= 2
+ and all(_is_power_of_two(dim) for dim in (head_dim_nope, head_dim_rope, head_dim_v))
+ ):
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,
- 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, 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,
+ kv, k_pos_emb, cos_table, sin_table, head_dim_nope, head_dim_rope, head_dim_v
)
- 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,
- )
- 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,
-):
+ 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).expand(
- -1, -1, kv.shape[2], -1
- )
- return torch.cat((k_nope, k_rope), dim=-1), 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()
+
+
+def _is_power_of_two(value: int) -> bool:
+ return value > 0 and value & (value - 1) == 0
diff --git a/transformer_engine/pytorch/cpp_extensions/gemm.py b/transformer_engine/pytorch/cpp_extensions/gemm.py
index ab51f9e5147..999f267c303 100644
--- a/transformer_engine/pytorch/cpp_extensions/gemm.py
+++ b/transformer_engine/pytorch/cpp_extensions/gemm.py
@@ -405,6 +405,9 @@ def general_grouped_gemm(
A = [_unwrap_tensor(a, "rowwise" if transa else "columnwise") for a in A]
B = [_unwrap_tensor(b, "columnwise" if transb else "rowwise") for b in B]
+ if any(isinstance(tensor, Float8BlockwiseQTensorStorage) for tensor in (*A, *B)):
+ use_split_accumulator = True
+
empty_tensor = _empty_tensor()
empty_tensors = [empty_tensor] * num_gemms
diff --git a/transformer_engine/pytorch/models/__init__.py b/transformer_engine/pytorch/models/__init__.py
new file mode 100644
index 00000000000..bee5474c816
--- /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/models/deepseek_v3/__init__.py b/transformer_engine/pytorch/models/deepseek_v3/__init__.py
new file mode 100644
index 00000000000..a7cbb50ae2f
--- /dev/null
+++ b/transformer_engine/pytorch/models/deepseek_v3/__init__.py
@@ -0,0 +1,13 @@
+# 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.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/models/deepseek_v3/moe.py b/transformer_engine/pytorch/models/deepseek_v3/moe.py
new file mode 100644
index 00000000000..9ad63ba6d42
--- /dev/null
+++ b/transformer_engine/pytorch/models/deepseek_v3/moe.py
@@ -0,0 +1,328 @@
+# 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."""
+
+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_and_pad_with_probs,
+ moe_permute_with_probs,
+ moe_unpermute,
+)
+from transformer_engine.pytorch.quantization import (
+ FP8GlobalStateManager,
+ get_align_size_for_quantization,
+)
+
+__all__ = ["DeepSeekV3MoE"]
+
+
+_EP_ALIGNMENT = 256
+_FUSED_MLP_ROWS = 256
+
+
+def _make_swiglu_mlp(hidden_size, ffn_hidden_size, dtype, device, num_experts=None):
+ """Dense SwiGLU MLP, or a grouped one (probs applied inside the activation) per expert.
+
+ The grouped variant fuses into a single CuTe grouped MLP on supported hardware.
+ """
+ common = {"bias": False, "dtype": dtype, "device": device}
+ if num_experts is None:
+ return te_ops.Sequential(
+ te_ops.Linear(hidden_size, 2 * ffn_hidden_size, **common),
+ te_ops.SwiGLU(),
+ te_ops.Linear(ffn_hidden_size, hidden_size, **common),
+ )
+ return te_ops.Sequential(
+ te_ops.GroupedLinear(num_experts, hidden_size, 2 * ffn_hidden_size, **common),
+ te_ops.ScaledSwiGLU(glu_interleave_size=32),
+ te_ops.GroupedLinear(num_experts, ffn_hidden_size, hidden_size, **common),
+ )
+
+
+class DeepSeekV3MoE(torch.nn.Module):
+ """
+ DeepSeekV3 Mixture-of-Experts block.
+
+ Each token is scored by a sigmoid router with a non-trainable expert bias
+ updated by ``update_expert_bias()`` (aux-loss-free load balancing) and,
+ optionally, group-limited routing: experts are split into ``num_groups``
+ groups, the top ``group_topk`` groups are selected by their summed scores,
+ and the final ``topk`` experts are chosen only from those groups. Selected
+ tokens run through the routed experts, a SwiGLU MLP shared across experts
+ as a grouped GEMM, with the routing probability applied inside the MLP. An
+ optional shared expert (dense SwiGLU MLP) is added to every token. On
+ hardware that supports it the expert MLP runs as a single fused
+ grouped-GEMM kernel. MXFP8 and NVFP4 pad each expert's token count to a
+ multiple of 256, including unfused execution; FP8 block scaling uses 128.
+
+ Without ``ep_group`` all experts live on the local device. With
+ ``ep_group`` the experts are split across the group and tokens are
+ exchanged over NCCL; this requires ``ep_bootstrap`` to be called once per
+ process before constructing the module, and bfloat16 inputs.
+
+ 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; requires ``group_topk``.
+ group_topk : int, optional
+ number of groups each token is limited to; requires ``num_groups``.
+ 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).
+ """
+
+ 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,
+ ) -> None:
+ super().__init__()
+
+ if num_experts <= 0:
+ raise ValueError("num_experts must be positive.")
+ if not 1 <= topk <= num_experts:
+ raise ValueError("topk must be in [1, num_experts].")
+ if (num_groups is None) != (group_topk is None):
+ raise ValueError("num_groups and group_topk must be provided together.")
+ if num_groups is not None:
+ if num_groups <= 0 or num_experts % num_groups != 0:
+ raise ValueError("num_groups must be positive and divide num_experts.")
+ if not 1 <= group_topk <= num_groups:
+ raise ValueError("group_topk must be in [1, num_groups].")
+ if topk % group_topk != 0:
+ raise ValueError("topk must be divisible by group_topk.")
+ if topk // group_topk > num_experts // num_groups:
+ raise ValueError("topk per group must not exceed the number of experts per group.")
+
+ 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
+
+ expert_mlp = _make_swiglu_mlp(
+ hidden_size, moe_ffn_hidden_size, dtype, device, num_experts=num_local_experts
+ )
+
+ self._ep_buffer_kwargs = None
+ if ep_group is not None:
+ from transformer_engine.pytorch.ep import EpConfig, get_ep_drop_on_overflow
+
+ assert ep_max_tokens_per_rank is not None, "EP requires ep_max_tokens_per_rank."
+ drop_on_overflow = get_ep_drop_on_overflow()
+ if drop_on_overflow is None:
+ raise RuntimeError("EP requires ep_bootstrap before constructing DeepSeekV3MoE.")
+ cap = self.ep_recv_capacity(
+ self.ep_size, ep_max_tokens_per_rank, topk, num_local_experts
+ )
+ self._ep_buffer_kwargs = {
+ "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": cap,
+ "alignment": _EP_ALIGNMENT,
+ }
+ config = EpConfig(
+ ep_group=ep_group,
+ drop_on_overflow=drop_on_overflow,
+ **self._ep_buffer_kwargs,
+ )
+ dispatch = te_ops.MoeDispatch(config)
+ combine = te_ops.MoeCombine(config)
+ fc1, activation, fc2 = expert_mlp
+ dispatch.set_extra_output_channel(0, "tokens_per_expert", output_to_caller=False)
+ dispatch.set_extra_output_channel(1, "routing_weights", output_to_caller=False)
+ fc1.set_extra_input_channel(0, "tokens_per_expert")
+ activation.set_extra_input_channel(0, "routing_weights")
+ fc2.set_extra_input_channel(0, "tokens_per_expert")
+ self.experts = te_ops.Sequential(
+ {"dispatch": dispatch, "0": fc1, "1": activation, "2": fc2, "combine": combine}
+ )
+ else:
+ self.experts = expert_mlp
+
+ self.shared_expert = None
+ if shared_expert_ffn_hidden_size is not None:
+ self.shared_expert = _make_swiglu_mlp(
+ hidden_size, shared_expert_ffn_hidden_size, dtype, device
+ )
+
+ @staticmethod
+ def ep_recv_capacity(
+ ep_size: int, max_tokens_per_rank: int, topk: int, num_local_experts: int
+ ) -> int:
+ """Worst-case receive capacity including per-expert alignment padding."""
+ cap = ep_size * max_tokens_per_rank * topk
+ cap += num_local_experts * _EP_ALIGNMENT
+ return -(-cap // _FUSED_MLP_ROWS) * _FUSED_MLP_ROWS
+
+ 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()
+
+ # Quantized grouped GEMMs need every expert's row count aligned.
+ align = 1
+ if FP8GlobalStateManager.is_fp8_enabled():
+ recipe = FP8GlobalStateManager.get_fp8_recipe()
+ align = get_align_size_for_quantization(recipe)
+ if recipe.mxfp8() or recipe.nvfp4():
+ align = max(align, _FUSED_MLP_ROWS)
+ elif recipe.float8_block_scaling():
+ align = max(align, 128)
+ if align > 1:
+ permuted, permuted_probs, row_id_map, pad_offsets, tokens_per_expert = (
+ moe_permute_and_pad_with_probs(tokens, probs, routing_map, tokens_per_expert, align)
+ )
+ else:
+ permuted, permuted_probs, row_id_map = moe_permute_with_probs(
+ tokens, probs, routing_map, num_out_tokens=tokens.shape[0] * self.topk
+ )
+ pad_offsets = None
+
+ # 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.
+ num_rows = permuted.shape[0]
+ pad = (-num_rows) % 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_rows], row_id_map, restore_shape=tokens.shape, pad_offsets=pad_offsets
+ )
+
+ def make_ep_buffer(self, device: Optional[torch.device] = None):
+ """EP buffer for this module's routing config. ``forward`` creates one per call unless
+ given one; a buffer holds one call's routing state until its backward runs, so reuse
+ it only across calls whose backward has completed."""
+ from transformer_engine.pytorch.ep import EpBuffer
+
+ assert self._ep_buffer_kwargs is not None, "make_ep_buffer requires ep_group."
+ if device is None:
+ device = torch.device("cuda", torch.cuda.current_device())
+ return EpBuffer(**self._ep_buffer_kwargs, device=device)
+
+ def _forward_ep(self, tokens: torch.Tensor, ep_buffer=None) -> torch.Tensor:
+ assert tokens.dtype == torch.bfloat16, "The EP path requires bfloat16 inputs."
+ buffer = ep_buffer if ep_buffer is not None else self.make_ep_buffer(tokens.device)
+ 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)
+ flat_idx = topk_idx.flatten()
+ self._last_tokens_per_expert = torch.zeros(
+ self.num_experts, dtype=torch.long, device=tokens.device
+ ).scatter_add_(0, flat_idx, torch.ones_like(flat_idx))
+ topk_weights = probs.gather(1, topk_idx)
+
+ return self.experts(
+ tokens,
+ topk_idx,
+ topk_weights,
+ topk_idx,
+ op_kwargs={self.experts[0]: {"buffer": buffer}, self.experts[-1]: {"buffer": buffer}},
+ )
+
+ def forward(self, hidden_states: torch.Tensor, ep_buffer=None) -> torch.Tensor:
+ """
+ Parameters
+ ----------
+ hidden_states : torch.Tensor
+ input of shape ``[..., hidden_size]``.
+ ep_buffer : EpBuffer, optional
+ buffer from :meth:`make_ep_buffer` to reuse instead of creating one
+ per call (EP only).
+ """
+ tokens = hidden_states.reshape(-1, self.hidden_size)
+ if self.ep_group is not None:
+ out = self._forward_ep(tokens, ep_buffer)
+ 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
new file mode 100644
index 00000000000..efc69f1594e
--- /dev/null
+++ b/transformer_engine/pytorch/models/deepseek_v3/multi_latent_attention.py
@@ -0,0 +1,281 @@
+# 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 math
+from typing import Optional, Union
+
+import torch
+
+from transformer_engine.pytorch.module import Linear, LayerNormLinear
+from transformer_engine.pytorch.attention import DotProductAttention
+from transformer_engine.pytorch.distributed import allreduce, get_distributed_world_size
+from transformer_engine.pytorch.attention.mla_rope import (
+ apply_mla_rope_kv,
+ apply_mla_rope_q,
+ build_rope_tables,
+ yarn_mscale,
+)
+
+__all__ = ["MultiLatentAttention"]
+
+
+class _ReduceGrad(torch.autograd.Function):
+ """Replicate an input across TP ranks and sum its gradients."""
+
+ @staticmethod
+ def forward(ctx, inp, tp_group):
+ """Return the replicated input unchanged."""
+ ctx.tp_group = tp_group
+ return inp
+
+ @staticmethod
+ def backward(ctx, grad_output):
+ """Sum gradients from each rank's attention heads."""
+ grad_input = grad_output.clone(memory_format=torch.contiguous_format)
+ grad_input, _ = allreduce(grad_input, ctx.tp_group)
+ return grad_input, None
+
+
+class MultiLatentAttention(torch.nn.Module):
+ """
+ Multi-Latent Attention as used in DeepSeekV3.
+
+ 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.
+
+ RoPE uses the fused MLA kernels from
+ :mod:`transformer_engine.pytorch.attention.mla_rope` (single-pass key/value
+ assembly); the rope slice follows
+ the DeepSeekV3 checkpoint convention (interleaved weights, NeoX output).
+
+ 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`.
+ layernorm_epsilon : float, default = 1e-6
+ epsilon of the latent RMSNorms (matches DeepSeekV3).
+ rotary_base : float, default = 10000.0
+ RoPE base.
+ rope_scaling_factor : float, optional
+ YaRN context-extension factor; ``None`` disables YaRN.
+ original_max_position_embeddings : int, default = 4096
+ pre-extension context length (YaRN).
+ beta_fast : float, default = 32.0
+ YaRN high-frequency rotation bound.
+ beta_slow : float, default = 1.0
+ YaRN low-frequency rotation bound.
+ mscale : float, default = 1.0
+ YaRN mscale of the rope part.
+ mscale_all_dim : float, default = 0.0
+ YaRN mscale of all dims; sets the default softmax scale to
+ ``m**2 / sqrt(qk head dim)`` with ``m = 0.1 * mscale_all_dim * ln(factor) + 1``.
+ softmax_scale : float, optional
+ softmax scale; defaults to ``1/sqrt(qk head dim)`` (times the YaRN
+ ``m**2`` when YaRN is enabled).
+ 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,
+ 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",
+ layernorm_epsilon: float = 1e-6,
+ rotary_base: float = 10000.0,
+ rope_scaling_factor: Optional[float] = None,
+ original_max_position_embeddings: int = 4096,
+ beta_fast: float = 32.0,
+ beta_slow: float = 1.0,
+ mscale: float = 1.0,
+ mscale_all_dim: float = 0.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__()
+
+ if tp_group is not None:
+ tp_size = get_distributed_world_size(tp_group)
+ 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",
+ eps=layernorm_epsilon,
+ 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",
+ eps=layernorm_epsilon,
+ 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.rotary_base = rotary_base
+ self._yarn_kwargs = {
+ "scaling_factor": rope_scaling_factor,
+ "original_max_position_embeddings": original_max_position_embeddings,
+ "beta_fast": beta_fast,
+ "beta_slow": beta_slow,
+ "mscale": mscale,
+ "mscale_all_dim": mscale_all_dim,
+ }
+ self._rope_tables: Optional[tuple] = None
+
+ if softmax_scale is None and rope_scaling_factor is not None:
+ m = yarn_mscale(rope_scaling_factor, mscale_all_dim)
+ softmax_scale = m * m / math.sqrt(self.qk_head_dim)
+ self.softmax_scale = softmax_scale
+
+ 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_tables_for(self, seq_len: int, device: torch.device):
+ if (
+ self._rope_tables is None
+ or self._rope_tables[0].shape[0] < seq_len
+ or self._rope_tables[0].device != device
+ ):
+ self._rope_tables = build_rope_tables(
+ seq_len,
+ self.qk_rope_head_dim,
+ base=self.rotary_base,
+ device=device,
+ **self._yarn_kwargs,
+ )
+ cos, sin = self._rope_tables
+ return cos[:seq_len], sin[:seq_len]
+
+ def forward(
+ self,
+ hidden_states: torch.Tensor,
+ attention_mask: Optional[torch.Tensor] = None,
+ attn_mask_type: Optional[str] = None,
+ ) -> 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.
+ """
+ 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)
+ if self.kv_up_proj.tp_size > 1:
+ # k_pos is shared across TP ranks; its attention gradients must be summed.
+ k_pos = _ReduceGrad.apply(k_pos, self.kv_up_proj.tp_group)
+ kv = self.kv_up_proj(kv_latent)
+ kv = kv.view(*kv.shape[:-1], heads, self.qk_nope_head_dim + self.v_head_dim)
+
+ 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, 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,
+ )
+
+ context = self.core_attention(
+ q,
+ k,
+ v,
+ attention_mask=attention_mask,
+ qkv_format=self.qkv_format,
+ attn_mask_type=attn_mask_type,
+ )
+ 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
new file mode 100644
index 00000000000..8ac8b93f8e9
--- /dev/null
+++ b/transformer_engine/pytorch/models/deepseek_v3/transformer_layer.py
@@ -0,0 +1,180 @@
+# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
+#
+# See LICENSE for license information.
+
+"""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`: 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.
+
+ 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
+ kwargs common to the submodules (``q_lora_rank``, ``kv_lora_rank``,
+ ``qk_nope_head_dim``, ``qk_rope_head_dim``, ``v_head_dim``,
+ ``attention_dropout``, ``attn_mask_type``, ``qkv_format``, ``topk``,
+ ``num_groups``, ``group_topk``, ``routed_scaling_factor``,
+ ``shared_expert_ffn_hidden_size``, EP options, ...), forwarded to
+ :class:`MultiLatentAttention` and :class:`DeepSeekV3MoE`.
+ """
+
+ _MLA_KWARGS = frozenset(
+ {
+ "q_lora_rank",
+ "kv_lora_rank",
+ "qk_nope_head_dim",
+ "qk_rope_head_dim",
+ "v_head_dim",
+ "attention_dropout",
+ "attn_mask_type",
+ "rotary_base",
+ "rope_scaling_factor",
+ "original_max_position_embeddings",
+ "beta_fast",
+ "beta_slow",
+ "mscale",
+ "mscale_all_dim",
+ "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",
+ }
+ )
+
+ 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__()
+
+ 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,
+ layernorm_epsilon=layernorm_epsilon,
+ 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,
+ ep_buffer=None,
+ ) -> 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.
+ ep_buffer : EpBuffer, optional
+ forwarded to :meth:`DeepSeekV3MoE.forward` (MoE layers with EP).
+ """
+ attention_out = self.self_attention(
+ self.input_layernorm(hidden_states),
+ attention_mask=attention_mask,
+ )
+ hidden_states = self._residual_add(attention_out, hidden_states)
+
+ mlp_kwargs = {"ep_buffer": ep_buffer} if ep_buffer is not None else {}
+ if self.pre_mlp_layernorm is not None:
+ mlp_out = self.mlp(self.pre_mlp_layernorm(hidden_states), **mlp_kwargs)
+ else:
+ mlp_out = self.mlp(hidden_states)
+ return self._residual_add(mlp_out, hidden_states)
diff --git a/transformer_engine/pytorch/ops/basic/combine.py b/transformer_engine/pytorch/ops/basic/combine.py
index 9f9825d8b28..bca6d49ee68 100644
--- a/transformer_engine/pytorch/ops/basic/combine.py
+++ b/transformer_engine/pytorch/ops/basic/combine.py
@@ -102,16 +102,12 @@ def fuser_forward(
basic_op_kwargs: list[dict[str, Any]],
) -> tuple[torch.Tensor, list[tuple[()]]]:
# NCCL EP reads the routing state from the EpBuffer, so topk_idx is unused.
- del (
- basic_op_extra_inputs,
- prev_op_grad_output_quantizer,
- next_op_input_quantizer,
- basic_op_kwargs,
- )
+ del basic_op_extra_inputs, prev_op_grad_output_quantizer, next_op_input_quantizer
# Only BF16 combine forward is supported for now.
input_ = maybe_dequantize(input_, torch.bfloat16)
ctx = basic_op_ctxs[0]
- buffer = validate_ep_buffer("MoeCombine", self.config, self.buffer)
+ kwargs = basic_op_kwargs[0]
+ buffer = validate_ep_buffer("MoeCombine", self.config, kwargs.get("buffer", self.buffer))
_validate_combine_inputs(input_, buffer)
result, combine_state = _ep_combine_fwd(
input_,
diff --git a/transformer_engine/pytorch/ops/basic/dispatch.py b/transformer_engine/pytorch/ops/basic/dispatch.py
index 4613be2c689..b9eee731346 100644
--- a/transformer_engine/pytorch/ops/basic/dispatch.py
+++ b/transformer_engine/pytorch/ops/basic/dispatch.py
@@ -103,10 +103,11 @@ def fuser_forward(
next_op_input_quantizer: Optional[Quantizer],
basic_op_kwargs: list[dict[str, Any]],
) -> tuple[torch.Tensor, Iterable[Iterable[torch.Tensor]]]:
- del next_op_input_quantizer, basic_op_kwargs
+ del next_op_input_quantizer
topk_idx, topk_weights = basic_op_extra_inputs[0]
ctx = basic_op_ctxs[0]
- buffer = validate_ep_buffer("MoeDispatch", self.config, self.buffer)
+ kwargs = basic_op_kwargs[0]
+ buffer = validate_ep_buffer("MoeDispatch", self.config, kwargs.get("buffer", self.buffer))
input_shape = _validate_dispatch_input(input_, buffer)
buffer.num_local_tokens = input_shape[0]
_validate_routing_inputs(