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(