Skip to content

[Docs] Add Mixture of Experts guide - #3494

Open
pggPL wants to merge 61 commits into
NVIDIA:mainfrom
pggPL:docs_moe
Open

pggPL wants to merge 61 commits into
NVIDIA:mainfrom
pggPL:docs_moe

Conversation

@pggPL

@pggPL pggPL commented Sep 7, 2026

Copy link
Copy Markdown
Collaborator

Adds a Mixture of Experts guide covering routing, token permutation, grouped expert computation, and expert parallelism.

Includes concise PyTorch and JAX examples.

pggPL and others added 30 commits May 4, 2026 15:00
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…PI entries

- Add code snippets and SVG figures referenced by mixture_of_experts.rst
  (moe_permute / moe_unpermute / grouped_linear tabbed examples for both
  PyTorch and JAX)
- Add JAX API reference entries for token_dispatch, token_combine and
  grouped_dense so the cross-references from the MoE page resolve
- Make wording framework-neutral where it was PyTorch-only (Grouped GEMM
  instead of GroupedLinear/grouped linear in shared sections, both
  m_splits and group_sizes mentioned, figure labels generalized)
- Tighten routing-kernel intro: consolidate the redundant "multiple
  variants exist / see API ref" notes into one paragraph next to the
  example, and explicitly state that the kernels are differentiable
- Sharpen merging_probs explanation (top-1 vs top-k) and explicitly
  describe what token_dispatch / token_combine return
- Snippet cleanups: define previously undefined symbols, drop the JAX
  probs= argument from the basic example and explain its purpose in a
  comment, document the ignored permuted_probs / pad_offsets outputs
- Reorder MoE entry in the docs/index.rst toctree

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Co-authored-by: Cursor <cursoragent@cursor.com>
…ecision

Add Router (score function + top-k + load-balancing loss) and Putting-it-together
sections, plus token-probabilities / padding-and-alignment / chunk-sort subsections
and a fused-expert-MLP note. New SVG figures and PyTorch/JAX snippets. Add the
router and moe_permute_and_pad_with_probs API reference entries (PyTorch and JAX)
and sort_chunks_by_index (JAX).

Co-Authored-By: Claude Opus 4.8 <noreply@anthropic.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…ument EP APIs

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…iagram

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…MM, grouped MLP, EP

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
… in introduction

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…e introduction

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…section

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…mework specifics to snippets

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…P conditions

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…LP figure

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…w layer figure

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…ections

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…-framework API)

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
…, MXFP8 dispatch, shared experts

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
pggPL and others added 5 commits September 7, 2026 19:54
…viour

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
@pggPL pggPL added the documentation Improvements or additions to documentation label Sep 8, 2026
…th -W

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>

# Conflicts:
#	docs/_static/css/diagram-colors.css
@pggPL
pggPL marked this pull request as ready for review September 8, 2026 10:48
@greptile-apps

greptile-apps Bot commented Sep 8, 2026 •

Copy link
Copy Markdown
Contributor

RetriggerConfidence Score: 5/5

[Low risk] Adds Mixture of Experts documentation and examples.

The PR appears safe to merge; no new actionable issue was identified.

Summary

The PR adds a Mixture of Experts guide, PyTorch and JAX examples, API references, and diagrams. The change since the previous review clarifies the required dtype, device, and layout of PyTorch’s compact top-k index buffer.

Diagram

%%{init: {'theme': 'neutral'}}%%
flowchart LR
  Router --> Dispatch[Token dispatch]
  Dispatch --> Experts[Grouped expert computation]
  Experts --> Combine[Token combine]
  Dispatch -. Expert parallelism .-> EPDispatch[All-to-all dispatch]
  EPDispatch --> Experts
  Experts --> EPCombine[All-to-all combine]
Loading

Reviews (8) · Last reviewed commit: "Clarify top-k index buffer requirements"

Comment thread docs/features/mixture_of_experts/moe_expert_parallel_pytorch.py Outdated
…P snippet

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
@pggPL

pggPL commented Sep 10, 2026

Copy link
Copy Markdown
Collaborator Author

Hi @vthumbe1503 @phu0ngng this is PR with docs for MoE operations and EP. Can you have a look?

)

# Transformer Engine: one grouped dense call. group_sizes is a device array.
# On Blackwell, BF16 and MXFP8 inputs without bias run as a single grouped GEMM

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Small update, Hopper BF16 is now supported since this PR merged yesterday: #3083

# On Blackwell, BF16 and MXFP8 inputs without bias run as a single grouped GEMM
# with the group sizes kept on the device; other cases launch one GEMM per
# expert and copy group_sizes to the host first.
grouped_out = te_dense.grouped_dense(

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

@pggPL How do we want to handle group alignment in the docs? This does support MXFP8 but we require group sizes to be aligned to a multiple of 128 for MXFP8, so if a user passes in un-aligned groups we can get IMA, illegal instruction, or incorrect results

For now, I've been guiding users to our more monolithic MoE block to avoid this complexity. But the same constraints apply to TE/PyTorch, so if you've found a better way to explain the nuances of this alignment, let me know and I'm open to adding it.

The on-device group size alignment is difficult since we can't assert it on Host without introducing runtime overhead

)

# 3. Experts: one grouped call over all expert token blocks.
expert_out = te_dense.grouped_dense(permuted, expert_weights, group_sizes=group_sizes)

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

At the moment, I'd prefer guiding users to our monolithic MoE layer here:

It's simpler and handles things like group size alignment automatically and hides it from the user. We do also want to highlight the lower-level APIs like this at some point in the future, but when we do so we need to significantly mark all the caveats like group size alignment and other constraints, which may be changing as we add new fused kernel support. Additionally, te_permutation.token_dispatch/combine are only recommended for non-EP, with EP enabled the TE EP APIs should be used. This is handled automatically by the TE MoEBlock, so my preference is to guide users to this monolithic MoE block as a starting point rather than these lower-level APIs

If Phuong wants to include a dedicated doc on TE EP APIs specifically, I'm okay with that lower-level API being highlighted. For other things, like grouped GEMM, I'd prefer to hide the alignment complexity in the TE MoEBlock and limit usage of lower-level APIs like grouped_dense to only users who are not doing a standard MoE architecture and are okay with handling the alignment constraints.

x_by_expert = jnp.split(x, split_indices, axis=0)

# Baseline: one matmul per expert.
loop_out = jnp.concatenate(

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Thanks for adding JAX documentation as well @pggPL! 🙌

I've reviewed and left a few comments. Let me know what you think. Thanks!

@vthumbe1503 vthumbe1503 left a comment •

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

I have reviewed the GroupedMLP side of documentation for now. And it looks pretty solid @pggPL !

Comment on lines +429 to +430
supported block-scaled recipe (MXFP8 or NVFP4). When the configuration is not
supported, the three operations run separately.

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

Suggested change
supported block-scaled recipe (MXFP8 or NVFP4). When the configuration is not
supported, the three operations run separately.
supported block-scaled recipe (MXFP8 or NVFP4). Fusion is only supported for a specific configuration [link](https://github.com/NVIDIA/TransformerEngine/blob/ea1a165decdd09c1ef98b273d965796c89c47473/transformer_engine/pytorch/ops/fused/grouped_mlp.py#L816-L817). When the configuration is not
supported, the three operations run separately.

We should specify what configuration means here. Feel free to modify the suggestion.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

It was intentionally left unspecified, since this is going to rapidly change in the future I guess and we usually forget to update the docs. I added link and specified that this is for some release and may change in the future.

Comment thread docs/features/mixture_of_experts/grouped_mlp_pytorch.py
@ptrendx ptrendx added the 2.21 label Sep 30, 2026
Document MoeDispatch and MoeCombine composition with internal channels. Preserve actual routing decisions for auxiliary loss, use the fusion-compatible GLU layout, and distinguish PyTorch and JAX EP behavior.

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Comment thread docs/features/mixture_of_experts/moe_expert_parallel_pytorch.py
pggPL added 2 commits October 1, 2026 19:47
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Comment thread docs/api/jax.rst
Comment on lines +73 to +75
.. autoapifunction:: transformer_engine.jax.permutation.token_dispatch

.. autoapifunction:: transformer_engine.jax.permutation.token_combine

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

I don't think we should expose these API now.
I would rather deprecate them and expose the permutation ops alone instead, in case there are users interested in the permutation alone.

Copy link
Copy Markdown
Collaborator Author

Choose a reason for hiding this comment

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

I'm not sure what you mean by "permutation ops". Can you please elaborate on that?

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

I mean permute alone, but not token_dispatch/token_combine.
It's confusing for users seeing ep.ep_dispatch vs permutation.token_dispatch

Comment thread docs/features/mixture_of_experts/mixture_of_experts.rst Outdated
Explain the PyTorch topk_indices output buffer, the unchanged probability tensor shape, and frontend support in the MoE router guide.

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>
Comment thread docs/features/mixture_of_experts/mixture_of_experts.rst Outdated
Document the CUDA device, contiguous layout and supported integer dtypes, show an allocation example, and clarify the routing map format restriction.

Signed-off-by: Pawel Gadzinski <pgadzinski@nvidia.com>

This branch has not been deployed

No deployments
Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

2.21 documentation Improvements or additions to documentation

Projects

None yet

Development

Successfully merging this pull request may close these issues.

5 participants