Skip to content
Open
Show file tree
Hide file tree
Changes from all commits
Commits
Show all changes
60 commits
Select commit Hold shift + click to select a range
45ed767
[PyTorch] Add DeepSeekV3Layer skeleton (MLA + MoE)
pggPL Aug 18, 2026
c306c6f
Move DeepSeekV3 skeleton to models/deepseek_v3 subpackage
pggPL Aug 18, 2026
f73d04e
Add DeepSeekV3 layer entries to PyTorch API docs
pggPL Aug 18, 2026
09f28a9
[PyTorch] Implement DeepSeekV3Layer: MLA + MoE from TE building blocks
pggPL Aug 18, 2026
e23100b
Add distributed EP test for DeepSeekV3 MoE/layer
pggPL Aug 18, 2026
4c6e1e8
Fix EP wgrad test collective + zero EP recv/grad buffers
pggPL Aug 18, 2026
aa17c37
Use fused MLA RoPE kernels in MultiLatentAttention
pggPL Aug 18, 2026
c713af7
Add HF transformers numeric reference test for DeepSeekV3Layer
pggPL Aug 21, 2026
88028a9
Docstring cleanups for lint and docs build
pggPL Aug 21, 2026
5f68c9b
Move model-specific layers to a dedicated docs page
pggPL Aug 21, 2026
9db495a
Drop HF-transformers comparison test from the repo
pggPL Aug 21, 2026
4883b17
Docs: reduce models page to a plain API listing
pggPL Aug 21, 2026
2f52af1
Rename distributed DeepSeek EP tests to generic test_models
pggPL Sep 3, 2026
e841f96
Rename test_deepseek.py to test_models.py and add models tests to QA …
pggPL Sep 3, 2026
935e475
Add YaRN RoPE scaling to DeepSeek V3 MLA
pggPL Sep 3, 2026
7eaefd9
Drop tests/pytorch/attention/mla_rope_utils.py shim; use models.deeps…
pggPL Sep 3, 2026
cef2e39
Distributed models test: single full DeepSeekV3Layer EP-vs-local nume…
pggPL Sep 3, 2026
5453fea
run_models.py: plain main() instead of unittest, simplify launcher
pggPL Sep 3, 2026
f24835b
run_models.py: fail hard instead of swallowing symm-mem/cleanup errors
pggPL Sep 3, 2026
d66fc1c
DeepSeekV3MoE docstring: ep_bootstrap must precede construction
pggPL Sep 3, 2026
80041fa
Distributed models test: launch torchrun directly from pytest, drop s…
pggPL Sep 3, 2026
2475a91
Add DeepSeekV3Layer to test_sanity; pad per-expert rows for quantized…
pggPL Sep 3, 2026
babc5e7
Tests: drop fwd/bwd smoke tests covered by sanity, trim sanity combos…
pggPL Sep 3, 2026
3354906
Rewrite DeepSeekV3MoE class docstring
pggPL Sep 3, 2026
397733e
DeepSeekV3MoE: drop ep_recv_capacity_per_rank and ep_alignment parame…
pggPL Sep 3, 2026
cc10354
DeepSeekV3MoE: build shared expert with the same SwiGLU MLP helper as…
pggPL Sep 3, 2026
2c1cd4c
DeepSeekV3MoE EP path: count tokens per expert with scatter_add inste…
pggPL Sep 3, 2026
46c065c
Merge remote-tracking branch 'origin/main' into deepseek_v3_layer
pggPL Sep 3, 2026
6bae1ba
Docs: list model-specific layers inline on the PyTorch API page; grou…
pggPL Sep 3, 2026
610a1e2
Lint: use dict literals in models.deepseek_v3
pggPL Sep 3, 2026
5c41a3b
Merge remote-tracking branch 'origin/main' into deepseek_v3_layer
pggPL Sep 8, 2026
b58a7ae
Merge remote-tracking branch 'upstream/main' into deepseek_v3_layer
pggPL Sep 8, 2026
93a40df
[PyTorch] DeepSeekV3MoE EP: drop full-buffer zeroing of recv/grad buf…
pggPL Sep 9, 2026
5e36ef1
[PyTorch] Add DeepSeekV3Layer expert-parallel example
pggPL Sep 9, 2026
3612879
[PyTorch] DeepSeekV3Layer example: naive PyTorch MoE and dense baselines
pggPL Sep 9, 2026
eb36d7d
[PyTorch] DeepSeekV3Layer example: two-level plain PyTorch MoE baseline
pggPL Sep 9, 2026
7a1693a
[PyTorch] DeepSeekV3: MLA TP gradient reduction, RoPE guards, per-for…
pggPL Sep 9, 2026
49ecaff
[PyTorch] DeepSeekV3Layer example: argument validation, gradient chec…
pggPL Sep 9, 2026
a131eb3
[PyTorch] Drop example-dependent tests from test_models.py
pggPL Sep 9, 2026
b5605bc
[Docs] Organize DeepSeekV3 example README around results
pggPL Sep 9, 2026
e9fac4f
[Docs] Keep naive MoE comparison in headline results
pggPL Sep 9, 2026
00bcbd9
[PyTorch] DeepSeekV3Layer example README: remeasure all variants on t…
pggPL Sep 9, 2026
4ffa639
[PyTorch] DeepSeekV3Layer example README: split NCCL EP permute from …
pggPL Sep 9, 2026
0367781
[PyTorch] DeepSeekV3MoE: optional caller-owned EpBuffer
pggPL Sep 9, 2026
da81d1d
[PyTorch] DeepSeekV3Layer example README: medians of three runs, brea…
pggPL Sep 9, 2026
2f7859a
[PyTorch] Validate MoE grouping and check synchronized expert bias up…
pggPL Sep 10, 2026
5371619
Merge upstream/main into deepseek_v3_layer
pggPL Sep 17, 2026
db03bd8
[PyTorch] Move MLA RoPE helpers into attention
pggPL Sep 17, 2026
e90cbd7
[PyTorch] Document MLA RoPE helpers
pggPL Sep 17, 2026
7030c91
[Docs] Nest MLA RoPE helpers under Other
pggPL Sep 17, 2026
793b46b
Merge upstream/main into deepseek_v3_layer
pggPL Sep 28, 2026
939c9db
Fix DeepSeekV3 layer review findings
pggPL Sep 30, 2026
00dbae6
Use bounded inputs and RHT in quantized MoE regression test
pggPL Sep 30, 2026
271a1be
Merge remote-tracking branch 'upstream/main' into deepseek_v3_layer
pggPL Sep 30, 2026
16eb75d
Refresh DeepSeekV3 benchmarks on GB200 and GB300
pggPL Sep 30, 2026
b32a4f8
Remove benchmark commit provenance from DeepSeek README
pggPL Oct 1, 2026
85612d4
Merge upstream main for MoE Sequential ops
pggPL Oct 1, 2026
0343fab
Compose DeepSeek expert parallel path with MoE ops
pggPL Oct 1, 2026
72abef7
Let MoE EP ops allocate their data buffers
pggPL Oct 5, 2026
b14c09f
Fix DeepSeek sanity coverage and grouped FP8 requirements
pggPL Oct 6, 2026
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
42 changes: 38 additions & 4 deletions docs/api/pytorch.rst
Original file line number Diff line number Diff line change
Expand Up @@ -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

Expand Down Expand Up @@ -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
Expand All @@ -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
----------

Expand Down
223 changes: 223 additions & 0 deletions examples/pytorch/deepseek_v3/README.md
Original file line number Diff line number Diff line change
@@ -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).

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

P2 Benchmark provenance removed

The benchmark results no longer identify when they were measured or which commit was used. That makes the reported timings harder to reproduce or compare with later changes. Please keep the measurement date and commit alongside the results.

Suggested change
Every number is the median of three independent runs (spread within 0.3 ms).
Every number is the median of three independent runs (spread within 0.3 ms).
Measured on September 30, 2026, at commit
[`939c9db3`](https://github.com/NVIDIA/TransformerEngine/commit/939c9db36afcdb2617392bab57b4249c7cd4bcb0),
before the subsequent merge of `main`.

Note: If this suggestion doesn't match your team's coding style, reply to this and let me know. I'll remember it for next time!

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=<head>: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_<hostname>.nsys-rep
nsys stats --report cuda_gpu_kern_sum results/deepseek_v3_layer_ep_<hostname>.nsys-rep
```

The launcher wraps `torchrun`, so all local ranks land in one report. For multi-node runs put
the same `nsys profile ... -o <path>_%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.
Loading
Loading