diff --git a/docs/_static/css/diagram-colors.css b/docs/_static/css/diagram-colors.css index 9ee5827bd1..f1e90a09b8 100644 --- a/docs/_static/css/diagram-colors.css +++ b/docs/_static/css/diagram-colors.css @@ -279,6 +279,38 @@ html[data-theme="dark"] .subtitle, html[data-theme="dark"] .memory-label { fill: #e0e0e0; } html[data-theme="dark"] .connector { stroke: #bdbdbd; } +/* mixture_of_experts diagrams */ +html[data-theme="dark"] .small-label, +html[data-theme="dark"] .tiny-label, +html[data-theme="dark"] .row-label, +html[data-theme="dark"] .rank-label { fill: #bdbdbd; } +html[data-theme="dark"] .grid-title, +html[data-theme="dark"] .brace-label, +html[data-theme="dark"] .row-label-l, +html[data-theme="dark"] .row-label-r, +html[data-theme="dark"] .mono { fill: #e0e0e0; } +html[data-theme="dark"] .expert-e0 { fill: #10375c; stroke: #64b5f6; } +html[data-theme="dark"] .expert-e1 { fill: #1e4620; stroke: #81c784; } +html[data-theme="dark"] .expert-e2 { fill: #5c3a10; stroke: #ffb74d; } +html[data-theme="dark"] .expert-e3 { fill: #5c1f38; stroke: #f48fb1; } +html[data-theme="dark"] .stage, +html[data-theme="dark"] .grouped, +html[data-theme="dark"] .a2a, +html[data-theme="dark"] .fused-mlp { fill: #3a2f5c; stroke: #b39ddb; } +html[data-theme="dark"] .moe-op { fill: #10375c; stroke: #64b5f6; } +html[data-theme="dark"] .act, +html[data-theme="dark"] .sel { fill: #1e4620; stroke: #81c784; } +html[data-theme="dark"] .linear, +html[data-theme="dark"] .param, +html[data-theme="dark"] .hbm, +html[data-theme="dark"] .cell, +html[data-theme="dark"] .box { fill: #2b2b2b; stroke: #9e9e9e; } +html[data-theme="dark"] .unsel, +html[data-theme="dark"] .pad { fill: #1f1f1f; stroke: #616161; } +html[data-theme="dark"] .soft-arrow, +html[data-theme="dark"] .brace { stroke: #bdbdbd; } +html[data-theme="dark"] .skip, +html[data-theme="dark"] .farrow { stroke: #b39ddb; } /* fine_grained_quantization diagrams */ html[data-theme="dark"] .fmt-mxfp8 { fill: #10375c; stroke: #64b5f6; } diff --git a/docs/api/jax.rst b/docs/api/jax.rst index 7a31c9d379..1a17db6ebf 100644 --- a/docs/api/jax.rst +++ b/docs/api/jax.rst @@ -59,3 +59,31 @@ Modules :members: __call__ .. autoapifunction:: transformer_engine.jax.flax.extend_logical_axis_rules + + +Mixture of Experts +------------------ +Router, routing kernels, and grouped dense for MoE layers. See +:doc:`Mixture of Experts <../features/mixture_of_experts/mixture_of_experts>` for an overview. + +.. autoapifunction:: transformer_engine.jax.router.fused_topk_with_score_function + +.. autoapifunction:: transformer_engine.jax.router.fused_moe_aux_loss + +.. autoapifunction:: transformer_engine.jax.permutation.token_dispatch + +.. autoapifunction:: transformer_engine.jax.permutation.token_combine + +.. autoapifunction:: transformer_engine.jax.permutation.sort_chunks_by_index + +.. autoapifunction:: transformer_engine.jax.dense.grouped_dense + +.. autoapifunction:: transformer_engine.jax.moe.moe + +.. autoapifunction:: transformer_engine.jax.ep.ep_bootstrap + +.. autoapifunction:: transformer_engine.jax.ep.ep_finalize + +.. autoapifunction:: transformer_engine.jax.ep.ep_dispatch + +.. autoapifunction:: transformer_engine.jax.ep.ep_combine diff --git a/docs/api/pytorch.rst b/docs/api/pytorch.rst index 73a08974c8..502446ee1a 100644 --- a/docs/api/pytorch.rst +++ b/docs/api/pytorch.rst @@ -96,16 +96,55 @@ Recipe availability Mixture of Experts (MoE) functions ---------------------------------- +See :doc:`Mixture of Experts <../features/mixture_of_experts/mixture_of_experts>` for an overview of +how these functions fit together. ``GroupedLinear`` (the grouped GEMM used for the +expert layers) is documented above with the other modules. + .. autoapifunction:: transformer_engine.pytorch.moe_permute .. autoapifunction:: transformer_engine.pytorch.moe_permute_with_probs +.. autoapifunction:: transformer_engine.pytorch.moe_permute_and_pad_with_probs + .. autoapifunction:: transformer_engine.pytorch.moe_unpermute .. autoapifunction:: transformer_engine.pytorch.moe_sort_chunks_by_index .. autoapifunction:: transformer_engine.pytorch.moe_sort_chunks_by_index_with_probs +.. autoapifunction:: transformer_engine.pytorch.router.fused_topk_with_score_function + +.. autoapifunction:: transformer_engine.pytorch.router.fused_compute_score_for_moe_aux_loss + +.. autoapifunction:: transformer_engine.pytorch.router.fused_moe_aux_loss + +Expert parallelism +~~~~~~~~~~~~~~~~~~ + +NCCL-based dispatch and combine for experts sharded across ranks. See +:doc:`Expert parallelism <../features/mixture_of_experts/mixture_of_experts>`. + +``MoeDispatch`` and ``MoeCombine`` can wrap the local expert operations in a +``transformer_engine.pytorch.ops.Sequential``. They share an ``EpConfig`` and +an ``EpBuffer``; the guide shows how to connect the per-expert counts and +routing weights through the sequence's extra-tensor channels. + +.. autoapiclass:: transformer_engine.pytorch.ops.MoeDispatch + +.. autoapiclass:: transformer_engine.pytorch.ops.MoeCombine + +.. autoapiclass:: transformer_engine.pytorch.ep.EpConfig + +.. autoapifunction:: transformer_engine.pytorch.ep.ep_bootstrap + +.. autoapifunction:: transformer_engine.pytorch.ep.ep_finalize + +.. autoapiclass:: transformer_engine.pytorch.ep.EpBuffer + +.. autoapifunction:: transformer_engine.pytorch.ep.ep_dispatch + +.. autoapifunction:: transformer_engine.pytorch.ep.ep_combine + Communication-computation overlap --------------------------------- diff --git a/docs/features/mixture_of_experts/grouped_linear_jax.py b/docs/features/mixture_of_experts/grouped_linear_jax.py new file mode 100644 index 0000000000..cc66e13596 --- /dev/null +++ b/docs/features/mixture_of_experts/grouped_linear_jax.py @@ -0,0 +1,30 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +# START_GROUPED_LINEAR_JAX +import jax.numpy as jnp +from transformer_engine.jax import dense as te_dense + +# x: [sum(group_sizes), hidden_size], expert-contiguous tokens +# kernel: [num_experts, hidden_size, ffn_hidden_size], stacked per-expert weights +# bias: [num_experts, ffn_hidden_size], stacked per-expert biases +# group_sizes: [num_experts] int array; group_sizes[i] is the number of routed +# tokens for expert i +split_indices = jnp.cumsum(group_sizes)[:-1] +x_by_expert = jnp.split(x, split_indices, axis=0) + +# Baseline: one matmul per expert. +loop_out = jnp.concatenate( + [x_i @ kernel_i + bias_i for x_i, kernel_i, bias_i in zip(x_by_expert, kernel, bias)], + axis=0, +) + +# Transformer Engine: one grouped dense call. group_sizes is a device array. +grouped_out = te_dense.grouped_dense( + x, + kernel, + group_sizes=group_sizes, + bias=bias, +) +# END_GROUPED_LINEAR_JAX diff --git a/docs/features/mixture_of_experts/grouped_linear_pytorch.py b/docs/features/mixture_of_experts/grouped_linear_pytorch.py new file mode 100644 index 0000000000..313098ac9f --- /dev/null +++ b/docs/features/mixture_of_experts/grouped_linear_pytorch.py @@ -0,0 +1,43 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +# START_GROUPED_LINEAR_PYTORCH +import torch +import transformer_engine.pytorch as te + +# x: [sum(m_splits), hidden_size], expert-contiguous tokens +# m_splits: list[int] of length num_experts; m_splits[i] is the number +# of routed tokens for expert i +# torch_experts: list[torch.nn.Linear] of length num_experts, one per expert +# (used only by the baseline loop below) +x_by_expert = torch.split(x, m_splits, dim=0) + +# Baseline: one Linear call per expert. +loop_out = torch.cat( + [expert(x_i) for expert, x_i in zip(torch_experts, x_by_expert)], + dim=0, +) + +# Transformer Engine: one grouped linear call. +grouped_linear = te.GroupedLinear( + num_experts, + hidden_size, + ffn_hidden_size, + bias=True, + params_dtype=torch.bfloat16, +).cuda() +grouped_out = grouped_linear(x, m_splits) + +# Keep token counts on the device with use_grouped_tensor=True. +grouped_linear = te.GroupedLinear( + num_experts, + hidden_size, + ffn_hidden_size, + bias=True, + params_dtype=torch.bfloat16, + use_grouped_tensor=True, +).cuda() +m_splits_dev = torch.tensor(m_splits, dtype=torch.int64, device="cuda") +grouped_out = grouped_linear(x, m_splits_dev) +# END_GROUPED_LINEAR_PYTORCH diff --git a/docs/features/mixture_of_experts/grouped_mlp_pytorch.py b/docs/features/mixture_of_experts/grouped_mlp_pytorch.py new file mode 100644 index 0000000000..9db0634647 --- /dev/null +++ b/docs/features/mixture_of_experts/grouped_mlp_pytorch.py @@ -0,0 +1,23 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +# START_GROUPED_MLP_PYTORCH +import transformer_engine.pytorch as te + +# FC1 produces gate and value features interleaved in blocks of 32. +expert_mlp = te.ops.Sequential( + te.ops.GroupedLinear(num_experts, hidden_size, 2 * ffn_hidden_size), + te.ops.ScaledSwiGLU(glu_interleave_size=32), + te.ops.GroupedLinear(num_experts, ffn_hidden_size, hidden_size, scale_bias=True), +) + +# Dispatch outputs: expert-contiguous tokens, aligned counts, and routing weights. +expert_out = expert_mlp( + permuted, + tokens_per_expert, + permuted_probs, + tokens_per_expert, + permuted_probs, +) +# END_GROUPED_MLP_PYTORCH diff --git a/docs/features/mixture_of_experts/img/grouped_linear.svg b/docs/features/mixture_of_experts/img/grouped_linear.svg new file mode 100644 index 0000000000..686593a9b4 --- /dev/null +++ b/docs/features/mixture_of_experts/img/grouped_linear.svg @@ -0,0 +1,148 @@ + + + + + + + + + + + + + + Loop over experts + + + + X0 + + W0 + + b0 + + + + Expert 0 (Linear) + + + Y0 + + + + + X1 + + W1 + + b1 + + + + Expert 1 (Linear) + + + Y1 + + + + + X2 + + W2 + + b2 + + + + Expert 2 (Linear) + + + Y2 + + + + + + + Grouped GEMM + + + + + W0 + + b0 + + W1 + + b1 + + W2 + + b2 + + + + + + + + + + + X0 + X1 + X2 + + + tokens for expert 0 + + tokens for expert 1 + + tokens for expert 2 + + + + + + + + Grouped GEMM + multiple experts + + + + tokens per expert + + + + + + + + + + + Y0 + Y1 + Y2 + + + diff --git a/docs/features/mixture_of_experts/img/moe_ep_one_step.svg b/docs/features/mixture_of_experts/img/moe_ep_one_step.svg new file mode 100644 index 0000000000..6f5c6f8f84 --- /dev/null +++ b/docs/features/mixture_of_experts/img/moe_ep_one_step.svg @@ -0,0 +1,44 @@ + + + + + + + + + generic all-to-all + + tokens + + + permute + sort by rank + + + all-to-all + + + permute + group by expert + + + experts + + NCCL EP dispatch + + tokens + + + dispatch + token → expert slot + + + experts + diff --git a/docs/features/mixture_of_experts/img/moe_expert_parallel.svg b/docs/features/mixture_of_experts/img/moe_expert_parallel.svg new file mode 100644 index 0000000000..d82d4e3071 --- /dev/null +++ b/docs/features/mixture_of_experts/img/moe_expert_parallel.svg @@ -0,0 +1,76 @@ + + + + + + + Expert parallelism: experts E0, E1 on rank 0 and E2, E3 on rank 1 + + tokens + dispatch all-to-all + receive buffer + local experts + combine all-to-all + output + + + Rank 0 + + Rank 1 + + + t0 → E0 + t1 → E2 + t2 → E0 + t3 → E1 + t4 → E3 + t5 → E2 + + + + + + + + + E0: t0 + E0: t2 + E1: t3 + E2: t1 + E2: t5 + E3: t4 + + + + Grouped MLPE0, E1 + + Grouped MLPE2, E3 + + + + + + + + + y0 + y1 + y2 + y3 + y4 + y5 + diff --git a/docs/features/mixture_of_experts/img/moe_expert_placement.svg b/docs/features/mixture_of_experts/img/moe_expert_placement.svg new file mode 100644 index 0000000000..6080f12085 --- /dev/null +++ b/docs/features/mixture_of_experts/img/moe_expert_placement.svg @@ -0,0 +1,43 @@ + + + + + + Rank 0 + tokens + + t0 → E0 + + t1 → E2 + + t2 → E0 + local experts + + E0 + + E1 + + Rank 1 + tokens + + t3 → E1 + + t4 → E3 + + t5 → E2 + local experts + + E2 + + E3 + diff --git a/docs/features/mixture_of_experts/img/moe_grouped_mlp.svg b/docs/features/mixture_of_experts/img/moe_grouped_mlp.svg new file mode 100644 index 0000000000..0f3840d654 --- /dev/null +++ b/docs/features/mixture_of_experts/img/moe_grouped_mlp.svg @@ -0,0 +1,40 @@ + + + + + + + + + FC1 grouped GEMM + + + + activation + + + + FC2 grouped GEMM + + + operation fuser + + + fused grouped-MLP operation + + FC1 grouped GEMM + + activation + + intermediate + device memory + + FC2 grouped GEMM + diff --git a/docs/features/mixture_of_experts/img/moe_layer_ep.svg b/docs/features/mixture_of_experts/img/moe_layer_ep.svg new file mode 100644 index 0000000000..ca4cda5546 --- /dev/null +++ b/docs/features/mixture_of_experts/img/moe_layer_ep.svg @@ -0,0 +1,45 @@ + + + + + + + + + + single device + + Router + + + Token dispatch + + + Grouped MLP + + + Token combine + + + expert parallelism (experts sharded across ranks) + + Router + + + All-to-all dispatch + grouped by local expert + + + Grouped MLP + local experts + + + All-to-all combine + back in token order + diff --git a/docs/features/mixture_of_experts/img/moe_padding.svg b/docs/features/mixture_of_experts/img/moe_padding.svg new file mode 100644 index 0000000000..d4ddddebb3 --- /dev/null +++ b/docs/features/mixture_of_experts/img/moe_padding.svg @@ -0,0 +1,56 @@ + + + + + + + + + Fused padding to expert-aligned blocks + + + permuted + expert 0 + + + + expert 1 + + expert 2 + + + tokens_per_expert = [3, 1, 2] + + + + token dispatch + align_size = 2 + + + padded + expert 0 + + + + pad + expert 1 + + pad + expert 2 + + + padded = [4, 2, 2] + diff --git a/docs/features/mixture_of_experts/img/moe_permute.svg b/docs/features/mixture_of_experts/img/moe_permute.svg new file mode 100644 index 0000000000..99b60b884f --- /dev/null +++ b/docs/features/mixture_of_experts/img/moe_permute.svg @@ -0,0 +1,86 @@ + + + + + + + + + Token Dispatch + + + + tokens + + t0 + + t1 + + t2 + + t3 + + t4 + + t5 + + + + + + + + Token Dispatch + group tokens by + destination expert + + + + + + + routing_map + + + + + row_id_map + + + + + permuted tokens + + t2 + + t3 + + t0 + + t4 + + t1 + + t5 + + + expert 0 + + expert 1 + + expert 2 + + diff --git a/docs/features/mixture_of_experts/img/moe_router.svg b/docs/features/mixture_of_experts/img/moe_router.svg new file mode 100644 index 0000000000..952082db10 --- /dev/null +++ b/docs/features/mixture_of_experts/img/moe_router.svg @@ -0,0 +1,103 @@ + + + + + + + + + Router: scoring and top-k selection + + + logits + e0 + e1 + e2 + e3 + + t0 + 1.2 + 3.1 + 0.4 + 2.7 + + t1 + 2.9 + 0.8 + 2.2 + 1.0 + + t2 + 0.5 + 1.1 + 3.3 + 2.5 + + + + + + + score function + + top-k + + + + + + + routing_map + e0 + e1 + e2 + e3 + + 0 + 1 + 0 + 1 + + 1 + 0 + 1 + 0 + + 0 + 0 + 1 + 1 + + + probs + e0 + e1 + e2 + e3 + + 0 + .6 + 0 + .4 + + .7 + 0 + .3 + 0 + + 0 + 0 + .6 + .4 + diff --git a/docs/features/mixture_of_experts/img/moe_unpermute.svg b/docs/features/mixture_of_experts/img/moe_unpermute.svg new file mode 100644 index 0000000000..aa27a7642e --- /dev/null +++ b/docs/features/mixture_of_experts/img/moe_unpermute.svg @@ -0,0 +1,81 @@ + + + + + + + + + Token Combine + + + + expert outputs + + y2 + + y3 + + y0 + + y4 + + y1 + + y5 + + + expert 0 + + expert 1 + + expert 2 + + + + + + + + Token Combine + restore original + token order + + + + + + + row_id_map + + + + + tokens (original order) + + y0 + + y1 + + y2 + + y3 + + y4 + + y5 + + diff --git a/docs/features/mixture_of_experts/mixture_of_experts.rst b/docs/features/mixture_of_experts/mixture_of_experts.rst new file mode 100644 index 0000000000..05f8cbd509 --- /dev/null +++ b/docs/features/mixture_of_experts/mixture_of_experts.rst @@ -0,0 +1,726 @@ +.. + Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + + See LICENSE for license information. + +.. _moe-overview: + +Mixture of Experts +=================================== + +.. warning:: + + MoE feature availability depends on the GPU, framework, Transformer Engine + and dependency versions, and precision recipe. Not every operation or + optimization works in every configuration. This guide explains the building + blocks; consult the :doc:`PyTorch API reference ` and + :doc:`JAX API reference ` for compatibility and input requirements. + +Introduction +------------ + +A Mixture of Experts (MoE) layer replaces a dense feed-forward network with a set +of expert networks and a router that sends each token to one or more experts. +A token passes through the layer in the following stages: + +#. The **router** scores the experts for each token and selects the top-k of + them. +#. **Token dispatch** gathers the tokens into expert-contiguous order. +#. The **grouped MLP** (the experts) runs a single batched computation over all + expert blocks. +#. **Token combine** scatters the expert outputs back into the original token + order, merging the contributions when a token was sent to more than one + expert. + +With expert parallelism the experts are sharded across ranks, and an +**all-to-all dispatch** and **all-to-all combine** take the place of token +dispatch and token combine: the dispatch takes the router output directly and +delivers each rank's tokens already grouped by local expert, and the combine +returns the outputs to the source rank in the original token order. + +.. raw:: html + :file: img/moe_layer_ep.svg + +*Figure 1. The stages of an MoE layer on a single device and with expert +parallelism.* + +Transformer Engine provides a building block for each stage. They are exposed as +standalone functions, so they can be assembled into a complete MoE layer or +dropped into an existing implementation one piece at a time: + +* :ref:`Router `: fused score function and top-k selection, and a + fused :ref:`load-balancing loss `. +* :ref:`Token permutation `: token dispatch and combine + kernels that move tokens between their original order and the + expert-contiguous layout. +* :ref:`Grouped GEMM `: the expert linear layers as one call + over expert-contiguous blocks; the :ref:`grouped MLP ` fuses + activation and supported quantization steps into the expert GEMMs. +* :ref:`Expert parallelism `: all-to-all dispatch and + combine for experts sharded across devices. + +The :ref:`single-device example ` wires the blocks +into a complete MoE layer; :ref:`Expert parallelism ` +covers the sharded case. + +.. _moe-router: + +Router +------ + +The router decides which experts each token is sent to. It applies a score +function to the gating logits, selects the top-k experts per token, and produces +the two tensors that drive the rest of the layer: + +* ``routing_map`` - by default, a ``[num_tokens, num_experts]`` mask marking + the selected experts. Token dispatch uses it to lay the tokens out by expert. +* ``probs`` - the routing weight of each selected expert. Token combine uses + these as merging weights when a token was routed to more than one expert. + +In PyTorch, passing a preallocated ``topk_indices`` buffer of shape +``[num_tokens, topk]`` to ``fused_topk_with_score_function`` selects a compact +list of expert indices instead. The buffer must be a contiguous CUDA tensor on +the same device as ``logits``, with dtype ``torch.int16``, ``torch.int32`` or +``torch.int64``. For example: + +.. code-block:: python + + topk_indices = torch.empty( + (num_tokens, topk), dtype=torch.int32, device=logits.device + ) + +Pass this buffer as ``topk_indices=topk_indices`` and leave +``routing_map_format`` at its default value. The function fills the buffer and +returns it as its second output without materializing the routing mask. The +``probs`` output still has shape ``[num_tokens, num_experts]``. The JAX router +does not currently expose this index-output option. Both frontends also support +a bit-packed mask through ``routing_map_format="bitmap_u8"``. The examples and +figure below use the default mask format. + +``fused_topk_with_score_function`` runs the score function and the top-k +selection in a single differentiable kernel. All internal math runs in FP32, +regardless of the logits dtype. + +.. raw:: html + :file: img/moe_router.svg + +*Figure 2. The router scores the experts for each token and keeps the top-k. +The selected entries populate* ``routing_map`` *(a 0/1 mask) and* ``probs`` *(the +routing weights); all other entries are zero.* + +Options: + +* **Score function:** softmax, sigmoid or sqrtsoftplus. With softmax, ``use_pre_softmax`` + selects whether the softmax is applied before or after the top-k. +* **Grouped routing:** the experts are split into ``num_groups`` equal groups. + Each group is scored by the sum of its best expert scores, the top + ``group_topk`` groups are kept, and the top-k experts are chosen only from + those groups. +* **Expert bias:** ``expert_bias`` is added to the scores before the top-k + selection (see :ref:`Load balancing `). +* **Scaling:** ``scaling_factor`` rescales the returned probabilities. + +.. tabs:: + + .. tab:: PyTorch + + .. literalinclude:: router_pytorch.py + :language: python + :start-after: # START_ROUTER_PYTORCH + :end-before: # END_ROUTER_PYTORCH + + .. tab:: JAX + + .. literalinclude:: router_jax.py + :language: python + :start-after: # START_ROUTER_JAX + :end-before: # END_ROUTER_JAX + +.. _moe-load-balancing: + +Load balancing +~~~~~~~~~~~~~~ + +``fused_moe_aux_loss`` computes the auxiliary load-balancing loss that penalizes +uneven token counts across experts. It takes the per-expert token counts and the +*dense* routing scores (one value per expert, not only the selected top-k), so +the loss has a gradient with respect to every expert's logit. The dense scores +are returned by the router functions shown below; add the scaled loss to the +training loss. + +Count tokens from the original router's ``routing_map``. The auxiliary-score +helpers ignore expert bias and grouped routing, so their returned maps can +select different experts. Use the same score function for routing and the +auxiliary scores. + +``expert_bias`` balances the load without an extra loss term: + +* with sigmoid or sqrtsoftplus it is added to the scores only for the top-k + selection, so it changes which experts are picked but not the returned + routing weights; +* update it between steps: lower it for overloaded experts, raise it for + under-used ones. + +.. tabs:: + + .. tab:: PyTorch + + .. literalinclude:: router_pytorch.py + :language: python + :start-after: # START_ROUTER_AUX_PYTORCH + :end-before: # END_ROUTER_AUX_PYTORCH + + .. tab:: JAX + + .. literalinclude:: router_jax.py + :language: python + :start-after: # START_ROUTER_AUX_JAX + :end-before: # END_ROUTER_AUX_JAX + +.. _moe-token-permutation: + +Token permutation +----------------- + +Token dispatch moves the tokens into the expert-contiguous layout expected by +the grouped GEMM, and token combine moves the expert outputs back. + +* All of these kernels are differentiable. +* The snippets below use the mask-map routing variant. Other variants (for + example index-map routing) follow the same pattern, see the :doc:`PyTorch API + reference ` and :doc:`JAX API reference `. + +Token dispatch +~~~~~~~~~~~~~~ + +Token dispatch takes the token tensor and a routing map describing each +token's destination experts, and returns a permuted token buffer in which all +rows assigned to the same expert are stored contiguously. This is the layout the +grouped GEMM consumes, together with the per-expert token counts. + +.. raw:: html + :file: img/moe_permute.svg + +*Figure 3. Token dispatch consumes the input token tensor together with the +routing map and produces an expert-contiguous token tensor; rows assigned to the +same expert are stored back-to-back.* + +A typical call looks like: + +.. tabs:: + + .. tab:: PyTorch + + .. literalinclude:: moe_permute_pytorch.py + :language: python + :start-after: # START_MOE_PERMUTE_PYTORCH + :end-before: # END_MOE_PERMUTE_PYTORCH + + .. tab:: JAX + + .. literalinclude:: moe_permute_jax.py + :language: python + :start-after: # START_MOE_PERMUTE_JAX + :end-before: # END_MOE_PERMUTE_JAX + +The call returns the permuted token buffer of shape +``[num_out_tokens, hidden_size]`` together with a ``row_id_map`` that token +combine uses to restore the original token order after the experts have run. + +Token combine +~~~~~~~~~~~~~ + +Token combine is the inverse operation: it takes the expert-contiguous output +of the grouped GEMM and the ``row_id_map`` returned by token dispatch, and +returns a tensor of shape ``[num_tokens, hidden_size]`` with the rows written +back into the original token order. + +Pass the routing weights as ``merging_probs``; the kernel then computes the +weighted sum of the per-expert contributions in the same fused pass. Without +them the contributions are summed unweighted. They may be omitted for top-1 +only when the router guarantees that every selected weight is one, for example +post-top-k softmax with ``scaling_factor=1``. Pre-softmax, sigmoid, +sqrtsoftplus, or an additional scaling factor can produce non-unit top-1 +weights that must still be applied. + +.. raw:: html + :file: img/moe_unpermute.svg + +*Figure 4. Token combine reads the expert-contiguous output tensor and the* +``row_id_map``\ *, and writes each row back to its original token slot. With* +``merging_probs``\ *, contributions from multiple experts to the same token are +combined in the same fused kernel.* + +A typical call looks like: + +.. tabs:: + + .. tab:: PyTorch + + .. literalinclude:: moe_unpermute_pytorch.py + :language: python + :start-after: # START_MOE_UNPERMUTE_PYTORCH + :end-before: # END_MOE_UNPERMUTE_PYTORCH + + .. tab:: JAX + + .. literalinclude:: moe_unpermute_jax.py + :language: python + :start-after: # START_MOE_UNPERMUTE_JAX + :end-before: # END_MOE_UNPERMUTE_JAX + +Token probabilities +~~~~~~~~~~~~~~~~~~~ + +The routing weights must be applied exactly once: + +* **In token combine.** Pass the original routing weights as ``merging_probs``, + as in the examples above. +* **Before token combine.** Pass the weights to token dispatch as ``probs`` to + receive them in expert-contiguous order, multiply the completed expert + outputs by the permuted weights, and call token combine without + ``merging_probs``. An implementation may instead fold the weights into a + mathematically equivalent point, such as immediately before the expert's + final bias-free linear layer. + +Token dispatch only permutes ``probs``; it does not apply them. In particular, +scaling the input to a nonlinear expert is not equivalent to weighting that +expert's output. + +Padding and alignment +~~~~~~~~~~~~~~~~~~~~~ + +Grouped GEMM backends require or prefer each expert's token block to start at +an aligned offset (for example, a multiple of 128 rows). Token dispatch can pad +each block up to a multiple of ``align_size`` in the same kernel. + +.. raw:: html + :file: img/moe_padding.svg + +*Figure 5. Each expert's block is rounded up to a multiple of* ``align_size``\ *. +The per-expert padding offsets are returned so that token combine can drop the +padding again.* + +The padded dispatch returns the padded token buffer, the aligned per-expert +token counts to pass to the grouped GEMM, and the per-expert ``pad_offsets`` +that token combine needs to remove the padding. + +.. tabs:: + + .. tab:: PyTorch + + .. literalinclude:: moe_permute_pad_pytorch.py + :language: python + :start-after: # START_MOE_PERMUTE_PAD_PYTORCH + :end-before: # END_MOE_PERMUTE_PAD_PYTORCH + + .. tab:: JAX + + .. literalinclude:: moe_permute_pad_jax.py + :language: python + :start-after: # START_MOE_PERMUTE_PAD_JAX + :end-before: # END_MOE_PERMUTE_PAD_JAX + +Reordering expert chunks +~~~~~~~~~~~~~~~~~~~~~~~~ + +``sort_chunks_by_index`` reorders whole blocks of rows: + +* the input ``[num_tokens, hidden_size]`` is split along the first dimension + into chunks of the given ``split_sizes``; +* the chunks are concatenated again in the order given by ``sorted_indices``: + output chunk ``i`` is input chunk ``sorted_indices[i]``, rows inside a chunk + keep their order; +* the operation is differentiable, and a ``_with_probs`` variant moves a + per-row probability tensor along with the rows. + +The typical use is expert parallelism over a generic all-to-all, where the +received buffer is ordered by source rank and then by expert, while the grouped +GEMM needs all rows of one expert together. With two source ranks and two local +experts: + +* received chunks: ``(rank 0, E4)``, ``(rank 0, E5)``, ``(rank 1, E4)``, + ``(rank 1, E5)``; +* ``sorted_indices = [0, 2, 1, 3]`` regroups them into ``E4, E4, E5, E5`` for + the grouped GEMM; +* after the experts have run, the inverse permutation restores the rank-major + order for the combine all-to-all. + +.. _moe-grouped-gemm: + +Grouped GEMM +------------ + +The grouped GEMM applies the per-expert linear layers in one call, replacing a +loop of one ``Linear`` call per expert and producing the same outputs. + +Let ``G`` be the number of experts. For expert ``i``, ``X_i`` is the routed +token block, ``W_i`` is the expert weight, and ``b_i`` is the optional bias: + +.. math:: + + Y_i = X_i W_i^T + b_i,\quad i = 0, \ldots, G - 1 + +The full layer output is the concatenation of all expert outputs: + +.. math:: + + Y = \mathrm{concat}(Y_0, Y_1, \ldots, Y_{G-1}) + +The number of token rows belonging to each expert is passed as a per-expert +token-count argument. + +.. raw:: html + :file: img/grouped_linear.svg + +*Figure 6. Both paths produce the same outputs from the same inputs. The +baseline launches one* ``Linear`` *per expert; the grouped GEMM replaces the +loop with one call.* + +Grouped GEMMs can use :doc:`low-precision training recipes +`: the inputs are quantized per expert +and the expert GEMMs run in the recipe's precision. + +The computation can be implemented in two ways: + +* **Per-expert GEMMs.** The input is split by expert and each expert's matrix + multiplication is launched separately. +* **Single grouped GEMM.** The experts run in one grouped operation. Keeping + token counts on the device avoids a host synchronization and helps with + CUDA graph capture. + +The snippets assume the tokens have already been permuted into +expert-contiguous order. + +.. tabs:: + + .. tab:: PyTorch + + .. literalinclude:: grouped_linear_pytorch.py + :language: python + :start-after: # START_GROUPED_LINEAR_PYTORCH + :end-before: # END_GROUPED_LINEAR_PYTORCH + + .. tab:: JAX + + .. literalinclude:: grouped_linear_jax.py + :language: python + :start-after: # START_GROUPED_LINEAR_JAX + :end-before: # END_GROUPED_LINEAR_JAX + + With MXFP8, pad each expert's token block to a multiple of 128 rows and + pass the padded counts in ``group_sizes``. Aligning only the total token + buffer is insufficient. The high-level JAX ``moe`` API handles this + per-expert padding during dispatch. + +.. _moe-grouped-mlp: + +Grouped MLP +----------- + +An expert MLP is two grouped GEMMs with an activation between them. Transformer +Engine can fuse the activation and quantization into the expert GEMMs. In the +fused path illustrated below, FC1 includes the activation and FC2 runs as a +separate grouped GEMM, reading the intermediate tensor from device memory. + +.. raw:: html + :file: img/moe_grouped_mlp.svg + +*Figure 7. The operation fuser replaces the two grouped GEMMs and the activation +between them with one grouped-MLP operation that launches separate FC1 and FC2 +kernels.* + +The fusion is applied by the :doc:`operation fuser `: +a grouped linear, a scaled GLU (or SReLU) activation and another grouped linear +in sequence are replaced with one fused grouped-MLP operation. +Fusion requires compatible layer dimensions, activation layout, and precision +recipe. + +Fusion eligibility may change between releases; this guide summarizes the +requirements, while the implementation contains the current checks. For example, +see the `fusion eligibility checks in Transformer Engine v2.19 +`_. + +For GLU fusion, use ``glu_interleave_size=32``: FC1 must produce alternating +blocks of 32 gate features and 32 value features. When loading a checkpoint +whose FC1 outputs store all gate features before all value features, convert +each expert's FC1 weight and bias with ``te.interleave_glu_tensor(tensor, 32)``. +``ScaledClampedQGeGLU`` uses the same layout. For ``ScaledSReLU``, FC1 instead +produces ``ffn_hidden_size`` features, without the GLU's doubled width. + +``ScaledSwiGLU`` multiplies the activation output by the routing weights before +FC2. To weight the entire expert output, FC2 must also scale its bias: + +.. math:: + + p\,(h W_2 + b_2) = (p\,h) W_2 + p\,b_2, + +where :math:`h` is the activation output and :math:`p` is the routing weight. +Set ``scale_bias=True`` on FC2 and pass the same routing weights to both the +activation and FC2. With ``bias=False`` on FC2, bias scaling is unnecessary. + +.. tabs:: + + .. tab:: PyTorch + + .. literalinclude:: grouped_mlp_pytorch.py + :language: python + :start-after: # START_GROUPED_MLP_PYTORCH + :end-before: # END_GROUPED_MLP_PYTORCH + +The arguments after ``permuted`` are consumed in operation order: FC1 takes the +per-expert token counts, the activation takes the routing weights, and FC2 takes +the counts and the same weights for its bias. ``expert_out`` is already weighted; +the subsequent combine must not apply the routing weights again. + +Set ``NVTE_CUTEDSL_FUSED_GROUPED_MLP=1`` before importing +``transformer_engine.pytorch`` to enable fusion for eligible configurations. + +.. _moe-putting-it-together: + +Example: MoE layer on a single device +------------------------------------- + +The runnable examples below wire the blocks together for top-k routing on a +single supported NVIDIA GPU: route, dispatch, run the experts, combine. To keep +them short, each expert is represented by one grouped linear layer; a full +expert MLP uses the same routing around two grouped linear layers and an +activation. + +.. tabs:: + + .. tab:: PyTorch + + .. literalinclude:: moe_layer_pytorch.py + :language: python + :start-after: # START_MOE_LAYER_PYTORCH + :end-before: # END_MOE_LAYER_PYTORCH + + .. tab:: JAX + + .. literalinclude:: moe_layer_jax.py + :language: python + :start-after: # START_MOE_LAYER_JAX + :end-before: # END_MOE_LAYER_JAX + +The example uses dropless routing (``num_out_tokens = num_tokens * top_k``), so +the dispatch buffer is sized statically rather than from a device-to-host sync. +Every stage is differentiable, so the assembled layer trains end to end. + +.. _moe-expert-parallelism: + +Expert parallelism +------------------ + +With expert parallelism (EP) the experts are sharded across ranks: every rank +keeps its own shard of the tokens and holds only a slice of the experts. + +.. raw:: html + :file: img/moe_expert_placement.svg + +*Figure 8. Expert placement: each rank holds its token shard and a subset of the +experts. Tokens t1, t3 and t5 are routed to experts on the other rank.* + +A token routed to an expert on another rank has to travel there and back. Two +all-to-all collectives wrap the local expert computation: + +* **Dispatch** sends each token to the rank that owns its expert. It takes the + router output (expert indices and weights) directly and delivers a receive + buffer grouped by local expert. +* The local grouped GEMM runs on the receive buffer. +* **Combine** returns the results to the source rank and writes them back in the + original token order. + +No separate token dispatch or token combine is needed. Shared experts, which +every token passes through, are not part of the dispatch: they run as a regular +dense MLP on the local tokens on every rank. + +.. raw:: html + :file: img/moe_expert_parallel.svg + +*Figure 9. Dispatch routes each token to the rank owning its expert, the local +experts run on the receive buffer, and combine returns the outputs to the source +rank.* + +Transformer Engine provides dispatch and combine as ready, differentiable +operations built on the NCCL EP library (``libnccl_ep``, loaded at runtime), so +an MoE layer with expert parallelism is just router, dispatch, local experts and +combine. + +**Communication and permutation in one step** + +Over a generic all-to-all the tokens have to be permuted twice: sorted by +destination rank before the collective and regrouped by expert after it. The +NCCL EP kernels fold both permutations into the transfer: + +* dispatch computes the destination slot of every token from the top-k indices + and writes it there directly, so no permutation kernel runs before or after + the communication; +* combine does the inverse in one pass: it returns each expert output to its + source rank and sums it into the original token position. + +.. raw:: html + :file: img/moe_ep_one_step.svg + +*Figure 10. Dispatch over a generic all-to-all needs a permutation on each side; +the NCCL EP dispatch writes every token straight into its expert slot.* + +**Receive buffer** + +The number of tokens a rank receives depends on the routing. Local expert +blocks are packed consecutively, with +each block optionally padded to the configured alignment. Their offsets are +the cumulative per-expert counts; any remaining capacity forms an unused tail. + +* **Fixed capacity**, ``recv_capacity_per_rank`` (an integer): the total buffer + size is fixed, while the expert block sizes and offsets depend on the routing. + + * Sizing needs no host synchronization, allowing CUDA graph capture when + the expert operations also support it. The low-level zero-copy EP path + additionally requires persistent caller-owned receive and gradient buffers. + * Without per-expert padding, ``ep_size * max_tokens_per_rank * top_k`` is a + dropless upper bound. With alignment enabled, capacity must also include + padding for every local expert. For example, counts of 129 and 127 need + 384 rows at alignment 128, even though only 256 tokens arrive. JAX provides + ``get_moe_recv_capacity_per_rank`` to compute an aligned bound. + * With balanced routing, a rank receives about + ``max_tokens_per_rank * top_k`` assignments. A smaller capacity based on + this estimate saves memory but can overflow; include alignment padding + when sizing it. + * On overflow the dispatch fails with a device-side error. With + ``drop_on_overflow=True`` it instead drops the tokens that do not fit; they + are not sent, so their experts contribute nothing to the output, like + token dropping in capacity-limited MoE. + * ``total_recv_tokens`` reports the required receive rows before dropping, + including alignment padding; compare it with capacity after the step to + detect an overflow. + +* **Eager (PyTorch)**, no capacity given: the buffer is sized from the actual + receive count each step. + + * This costs a host synchronization per step and is not CUDA-graph + capturable; ``drop_on_overflow`` does not apply. + +JAX requires a static receive capacity. In the high-level ``moe`` API, +``recv_capacity_per_rank=None`` selects an aligned dropless worst-case bound; +it does not size the buffer from each step's actual receive count. Pass the +same static capacity to ``ep_bootstrap``. + +**Framework APIs** + +.. tabs:: + + .. tab:: PyTorch + + ``te.ops.MoeDispatch`` and ``te.ops.MoeCombine`` compose with the local + expert MLP in one ``te.ops.Sequential``. The sequence runs dispatch, + grouped linear, scaled activation, grouped linear, and combine, with + autograd through both the tokens and routing weights. + + Initialize NCCL EP once per process with ``ep_bootstrap(ep_group, ...)``. + Both operations take the same immutable ``EpConfig`` and ``EpBuffer``. + Their settings must match the bootstrap configuration, including the + process group, token limit, receive capacity, and overflow policy. + ``EpBuffer`` stores the routing handle and per-expert counts until backward + completes; use a separate buffer for each concurrently in-flight layer + call, including overlapping pipeline microbatches. + + Dispatch produces two extra tensors: the per-expert row counts (including + alignment padding) and the received routing weights. Named channels pass + the counts to both grouped linear operations and the weights to the + activation. Bind these channels before constructing or calling the + sequence. ``output_to_caller=False`` keeps the metadata internal, so the + sequence returns only the combined output. + + This assembly example assumes an initialized NCCL process group, the + current CUDA device set for each rank, and router outputs for the local + BF16 tokens. ``num_experts`` must be divisible by the EP group size, and + ``max_tokens_per_rank`` must bound each rank's input token count. The + expert MLP uses BF16 without bias or per-expert padding. + + .. literalinclude:: moe_expert_parallel_pytorch.py + :language: python + :start-after: # START_MOE_EXPERT_PARALLEL_PYTORCH + :end-before: # END_MOE_EXPERT_PARALLEL_PYTORCH + + The public extra arguments follow operation order: dispatch consumes + ``topk_idx`` and ``topk_w``, and combine consumes ``topk_idx`` again. The + grouped linear and activation inputs come from the internal channels. + NCCL EP uses the shared buffer's routing state for combine, but the + operation still requires the routing-index argument. + + ``ScaledSwiGLU`` applies the routing weights before the bias-free FC2, so + combine only sums the already weighted expert outputs. For an FC2 with + bias, set ``scale_bias=True`` and bind its second extra input to the same + ``routing_weights`` channel, as described in :ref:`moe-grouped-mlp`. + + PyTorch dispatch can quantize tokens before sending them, reducing + communication volume; expert GEMMs can consume the quantized groups. + Communication precision is configured separately from the expert GEMMs: + set ``dispatch_fwd_quant_recipe`` and ``combine_bwd_quant_recipe`` in both + ``EpConfig`` and ``EpBuffer`` for the corresponding communication paths. + Use ``te.autocast`` to select the expert-compute recipe, and choose an + alignment and receive capacity compatible with both the communication and + expert kernels. See the configurations in + `tests/pytorch/distributed/run_ep.py `_. + + The lower-level ``ep_dispatch`` and ``ep_combine`` functions remain + available for custom expert code or explicit receive-buffer management. + ``ep_dispatch`` returns ``(recv_tokens, recv_topk_weights, + tokens_per_expert)``; unlike the operation's extra outputs, the weights + precede the counts. Weight the expert outputs before ``ep_combine``, which + returns the summed results in the original token order. When operating on + a full fixed-capacity buffer, clear unused tail rows and weights before + multiplying so uninitialized values cannot introduce NaNs. NCCL EP zeroes + alignment padding, and combine ignores slots without a routed token. + + .. tab:: JAX + + JAX EP sends BF16 tokens. Low-precision expert computation quantizes the + received tokens after dispatch; it does not change communication precision. + + JAX offers two levels of API, both experimental: + + * ``transformer_engine.jax.moe.moe`` runs the whole MoE block (router, + dispatch, expert MLPs, combine) as a single differentiable call. It is + executed inside a ``Mesh``; ``ep_axis`` names the mesh axis the experts + are sharded over and the dispatch and combine become all-to-all + collectives over that axis. It also returns the load-balancing loss when + ``aux_loss_coeff`` is non-zero. Before the first call, initialize EP + eagerly with ``ep_bootstrap`` and record the same configuration with + ``record_ep_bootstrap_signature_for_moe``. The snippet below shows the + required TE EP bootstrap sequence. + * ``transformer_engine.jax.ep`` exposes the primitives separately. Unlike + the PyTorch ``EpBuffer``, the routing state is not kept in an object: + dispatch returns it as arrays and the caller passes them on to combine. + + * ``ep_bootstrap(world_size, rank, num_experts, max_tokens_per_rank, + recv_capacity_per_rank, hidden_dim, ...)`` initializes the EP group + once per process. It runs inside the active ``Mesh`` and reads the EP + axis (and the data-parallel axes) from ``MeshResource``. Multiple local + devices per process require TE's support for borrowing the XLA NCCL + communicator; the example below uses one process per device. + * ``EpLayerConfig(top_k, ...)`` is a small per-layer configuration that + every per-step call takes as its first argument. + * ``ep_dispatch(cfg, topk_idx, tokens, topk_weights, + recv_capacity_per_rank)`` scatters the tokens to the expert ranks and + returns ``(recv_tokens, recv_topk_weights, handle_mem, token_counts, + total_recv_tokens)``: the receive buffer grouped by local expert, the + weights of the received tokens, the routing handle and per-expert + token counts needed by combine, and the pre-drop receive total that can + be used to detect overflow. + * ``ep_combine(cfg, handle_mem, token_counts, expert_out, + num_local_tokens)`` sums the expert outputs back on the source ranks in + the original token order. It is unweighted: multiply ``expert_out`` by + ``recv_topk_weights`` before calling it. NCCL EP zeroes alignment + padding; combine ignores slots without a routed token. + ``num_local_tokens`` must be static because it fixes the output shape. + + .. literalinclude:: moe_expert_parallel_jax.py + :language: python + :start-after: # START_MOE_EXPERT_PARALLEL_JAX + :end-before: # END_MOE_EXPERT_PARALLEL_JAX + +Complete runnable examples: +`examples/pytorch/ep `_ +and `examples/jax/ep `_. diff --git a/docs/features/mixture_of_experts/moe_expert_parallel_jax.py b/docs/features/mixture_of_experts/moe_expert_parallel_jax.py new file mode 100644 index 0000000000..2dd5367c4c --- /dev/null +++ b/docs/features/mixture_of_experts/moe_expert_parallel_jax.py @@ -0,0 +1,71 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +# START_MOE_EXPERT_PARALLEL_JAX +import jax + +from transformer_engine.jax.ep import ep_bootstrap +from transformer_engine.jax.moe import ( + get_moe_recv_capacity_per_rank, + moe, + record_ep_bootstrap_signature_for_moe, +) +from transformer_engine.jax.sharding import MeshResource, global_shard_guard + +num_experts = 8 +top_k = 2 +ep_axis = "ep" +ep_size = mesh.shape[ep_axis] +max_tokens_per_rank = x.shape[0] * x.shape[1] // jax.process_count() +recv_capacity_per_rank = get_moe_recv_capacity_per_rank( + num_experts=num_experts, + num_experts_per_tok=top_k, + max_tokens_per_rank=max_tokens_per_rank, + ep_size=ep_size, +) + +# Initialize EP eagerly once per process. The mesh has one device per process. +mesh_resource = MeshResource(ep_resource=ep_axis) +with mesh, global_shard_guard(mesh_resource): + ep_bootstrap( + world_size=jax.process_count(), + rank=jax.process_index(), + num_experts=num_experts, + max_tokens_per_rank=max_tokens_per_rank, + recv_capacity_per_rank=recv_capacity_per_rank, + hidden_dim=x.shape[-1], + max_token_dtype=x.dtype, + ) +record_ep_bootstrap_signature_for_moe( + num_experts=num_experts, + max_tokens_per_rank=max_tokens_per_rank, + recv_capacity_per_rank=recv_capacity_per_rank, + hidden_dim=x.shape[-1], + ep_size=ep_size, +) + +# mesh: jax.sharding.Mesh with an "ep" axis and one device per process +# x: [batch, seq, hidden_size], BF16 and sharded over the mesh +# gate_kernel: [hidden_size, num_experts] router projection +# wi: [num_experts, hidden_size, 2 * ffn] gated FC1 (gate and value) +# wo: [num_experts, ffn, hidden_size] FC2 +with mesh, global_shard_guard(mesh_resource): + output, aux_loss, total_recv_tokens = moe( + x, + gate_kernel, + wi, + wo, + num_experts=num_experts, + num_experts_per_tok=top_k, + activation_type="silu", + score_function="softmax", + aux_loss_coeff=1e-2, # load-balancing loss; 0 disables it + ep_axis=ep_axis, + dtype=x.dtype, + recv_capacity_per_rank=recv_capacity_per_rank, + ) +# output: [batch, seq, hidden_size] +# aux_loss: scalar load-balancing loss (None when aux_loss_coeff == 0) +# total_recv_tokens: receive count before any capacity-based token dropping +# END_MOE_EXPERT_PARALLEL_JAX diff --git a/docs/features/mixture_of_experts/moe_expert_parallel_pytorch.py b/docs/features/mixture_of_experts/moe_expert_parallel_pytorch.py new file mode 100644 index 0000000000..d433427029 --- /dev/null +++ b/docs/features/mixture_of_experts/moe_expert_parallel_pytorch.py @@ -0,0 +1,78 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +# START_MOE_EXPERT_PARALLEL_PYTORCH +import torch +import torch.distributed as dist +import transformer_engine.pytorch as te +from transformer_engine.pytorch.ep import EpBuffer, EpConfig, ep_bootstrap + +# ep_group: process group the experts are sharded over +# tokens: [num_tokens, hidden_size] bf16 tokens local to this rank +# topk_idx: [num_tokens, top_k] global expert index per selected expert +# topk_w: [num_tokens, top_k] fp32 routing weights from the router +ep_size = dist.get_world_size(ep_group) +num_local_experts = num_experts // ep_size +recv_capacity = ep_size * max_tokens_per_rank * top_k # alignment=0 +config = EpConfig( + top_k=top_k, + hidden_dim=hidden_size, + num_local_experts=num_local_experts, + max_tokens_per_rank=max_tokens_per_rank, + recv_capacity_per_rank=recv_capacity, + ep_group=ep_group, + alignment=0, +) + +# Once per process: sets up NCCL EP on ep_group's communicator. +ep_bootstrap( + ep_group, + num_experts=num_experts, + max_tokens_per_rank=config.max_tokens_per_rank, + hidden_dim=hidden_size, + num_topk=top_k, + recv_capacity_per_rank=recv_capacity, +) +# One buffer per in-flight layer call (e.g. per pipeline microbatch). +buffer = EpBuffer( + top_k=config.top_k, + max_tokens_per_rank=config.max_tokens_per_rank, + recv_capacity_per_rank=config.recv_capacity_per_rank, + hidden_dim=config.hidden_dim, + num_local_experts=config.num_local_experts, + alignment=config.alignment, + device=tokens.device, +) + +dispatch = te.ops.MoeDispatch(config, buffer) +fc1 = te.ops.GroupedLinear( + num_local_experts, + hidden_size, + 2 * ffn_hidden_size, + bias=False, + device=tokens.device, + dtype=torch.bfloat16, +) +activation = te.ops.ScaledSwiGLU() +fc2 = te.ops.GroupedLinear( + num_local_experts, + ffn_hidden_size, + hidden_size, + bias=False, + device=tokens.device, + dtype=torch.bfloat16, +) +combine = te.ops.MoeCombine(config, buffer) + +# Bind dispatch metadata before constructing or calling the sequence. +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") +moe = te.ops.Sequential(dispatch, fc1, activation, fc2, combine) + +# Dispatch takes indices and weights; combine takes the same indices. +output = moe(tokens, topk_idx, topk_w, topk_idx) # [num_tokens, hidden_size] +# END_MOE_EXPERT_PARALLEL_PYTORCH diff --git a/docs/features/mixture_of_experts/moe_layer_jax.py b/docs/features/mixture_of_experts/moe_layer_jax.py new file mode 100644 index 0000000000..f500287d89 --- /dev/null +++ b/docs/features/mixture_of_experts/moe_layer_jax.py @@ -0,0 +1,49 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +# START_MOE_LAYER_JAX +import jax +import jax.numpy as jnp +from transformer_engine.jax import permutation as te_permutation +from transformer_engine.jax import dense as te_dense +from transformer_engine.jax.router import fused_topk_with_score_function + +num_tokens, hidden_size, num_experts, top_k = 16, 64, 4, 2 +keys = jax.random.split(jax.random.key(0), 3) +hidden_states = jax.random.normal(keys[0], (num_tokens, hidden_size), dtype=jnp.bfloat16) +gate_kernel = jax.random.normal(keys[1], (hidden_size, num_experts), dtype=jnp.bfloat16) +kernel = jax.random.normal( + keys[2], + (num_experts, hidden_size, hidden_size), + dtype=jnp.bfloat16, +) + + +@jax.jit +def moe_layer(tokens, gate_weight, expert_weights): + # 1. Router: score the experts and pick the top-k for each token. + logits = tokens @ gate_weight + probs, routing_map = fused_topk_with_score_function( + logits, + topk=top_k, + score_function="softmax", + ) + + # 2. Dispatch: gather tokens into expert-contiguous order. + permuted, _, row_id_map, _, group_sizes = te_permutation.token_dispatch( + tokens, + routing_map.astype(jnp.int32), + num_out_tokens=num_tokens * top_k, + ) + + # 3. Experts: one grouped call over all expert token blocks. + expert_out = te_dense.grouped_dense(permuted, expert_weights, group_sizes=group_sizes) + + # 4. Combine: restore token order and merge the top-k contributions. + return te_permutation.token_combine(expert_out, row_id_map, merging_probs=probs) + + +output = moe_layer(hidden_states, gate_kernel, kernel) +output.block_until_ready() +# END_MOE_LAYER_JAX diff --git a/docs/features/mixture_of_experts/moe_layer_pytorch.py b/docs/features/mixture_of_experts/moe_layer_pytorch.py new file mode 100644 index 0000000000..211bd417f5 --- /dev/null +++ b/docs/features/mixture_of_experts/moe_layer_pytorch.py @@ -0,0 +1,67 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +# START_MOE_LAYER_PYTORCH +import torch +import transformer_engine.pytorch as te +from transformer_engine.pytorch.router import fused_topk_with_score_function + +num_tokens, hidden_size, num_experts, top_k = 16, 64, 4, 2 +hidden_states = torch.randn( + num_tokens, + hidden_size, + device="cuda", + dtype=torch.bfloat16, + requires_grad=True, +) +gate = torch.nn.Linear( + hidden_size, + num_experts, + bias=False, + device="cuda", + dtype=torch.bfloat16, +) +experts = te.GroupedLinear( + num_experts, + hidden_size, + hidden_size, + bias=False, + params_dtype=torch.bfloat16, + device="cuda", +) + +# 1. Router: score the experts and pick the top-k for each token. +logits = gate(hidden_states) +probs, routing_map = fused_topk_with_score_function( + logits, + topk=top_k, + use_pre_softmax=False, + num_groups=None, + group_topk=None, + scaling_factor=None, + score_function="softmax", + expert_bias=None, +) + +# 2. Dispatch: gather tokens into expert-contiguous order. +routing_map = routing_map.to(torch.int32) +permuted, row_id_map = te.moe_permute( + hidden_states, + routing_map, + num_out_tokens=num_tokens * top_k, +) + +# 3. Experts: one grouped call over all expert token blocks. +m_splits = routing_map.sum(dim=0).tolist() # tokens routed to each expert +expert_out = experts(permuted, m_splits) + +# 4. Combine: scatter the outputs back and merge the top-k contributions. +output = te.moe_unpermute( + expert_out, + row_id_map, + merging_probs=probs, + restore_shape=(num_tokens, hidden_size), +) +output.square().mean().backward() +# END_MOE_LAYER_PYTORCH diff --git a/docs/features/mixture_of_experts/moe_permute_jax.py b/docs/features/mixture_of_experts/moe_permute_jax.py new file mode 100644 index 0000000000..3740c2a126 --- /dev/null +++ b/docs/features/mixture_of_experts/moe_permute_jax.py @@ -0,0 +1,31 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +# START_MOE_PERMUTE_JAX +import jax.numpy as jnp +from transformer_engine.jax import permutation as te_permutation + +# tokens: [num_tokens, hidden_size] +# routing_map: [num_tokens, num_experts] mask, 1 if token routed to expert +# top_k is the statically configured number of experts selected per token. +# tokens.shape[0] is static while tracing, so this stays valid under ``jit``. +permuted, _, row_id_map, _, group_sizes = te_permutation.token_dispatch( + tokens, + routing_map.astype(jnp.int32), + num_out_tokens=tokens.shape[0] * top_k, +) + +# permuted: [num_out_tokens, hidden_size], expert-contiguous +# group_sizes: [num_experts], per-expert token counts; can be passed directly +# to ``grouped_dense`` as ``group_sizes``. +# row_id_map: opaque tensor used by ``token_combine`` to reverse the permutation. +# +# The two ignored outputs are ``permuted_probs`` and ``pad_offsets``: +# - ``permuted_probs`` (returned only when ``probs=`` is supplied) holds the +# routing probabilities in expert-contiguous order. Multiply the completed +# expert outputs by these weights before ``token_combine``; dispatch does +# not apply them itself. Do not also pass ``merging_probs`` in that case. +# - ``pad_offsets`` is only used together with ``align_size`` for fused +# padding to expert-aligned blocks. +# END_MOE_PERMUTE_JAX diff --git a/docs/features/mixture_of_experts/moe_permute_pad_jax.py b/docs/features/mixture_of_experts/moe_permute_pad_jax.py new file mode 100644 index 0000000000..2c4ce61ac2 --- /dev/null +++ b/docs/features/mixture_of_experts/moe_permute_pad_jax.py @@ -0,0 +1,37 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +# START_MOE_PERMUTE_PAD_JAX +from transformer_engine.jax import permutation as te_permutation + +# tokens: [num_tokens, hidden_size] +# probs: [num_tokens, num_experts] routing probabilities +# routing_map: [num_tokens, num_experts] int32 mask +# +# Passing align_size enables the same fused padding. token_dispatch allocates a +# fixed worst-case buffer (so it stays jit-compatible) and reports the aligned +# per-expert counts together with the padding offsets. +padded, permuted_probs, row_id_map, pad_offsets, tokens_per_expert = te_permutation.token_dispatch( + tokens, + routing_map, + num_out_tokens=num_tokens * top_k, + probs=probs, + align_size=128, +) + +# tokens_per_expert: aligned per-expert counts -> group_sizes for grouped_dense + +# ... run the grouped GEMM on `padded`, producing expert_out ... + +# Apply the permuted routing weights to the completed expert outputs. Since the +# weights are applied here, do not pass them to token_combine as well. +expert_out = expert_out * permuted_probs[:, None] + +# Pass pad_offsets so token combine removes the padding it added. +output = te_permutation.token_combine( + expert_out, + row_id_map, + pad_offsets=pad_offsets, +) +# END_MOE_PERMUTE_PAD_JAX diff --git a/docs/features/mixture_of_experts/moe_permute_pad_pytorch.py b/docs/features/mixture_of_experts/moe_permute_pad_pytorch.py new file mode 100644 index 0000000000..c7b900fe61 --- /dev/null +++ b/docs/features/mixture_of_experts/moe_permute_pad_pytorch.py @@ -0,0 +1,43 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +# START_MOE_PERMUTE_PAD_PYTORCH +from transformer_engine.pytorch import moe_permute_and_pad_with_probs, moe_unpermute + +# tokens: [num_tokens, hidden_size] +# probs: [num_tokens, num_experts] routing probabilities +# routing_map: [num_tokens, num_experts] int32 mask +# +# Pad each expert's token block up to a multiple of align_size (here 128) so the +# grouped GEMM sees aligned blocks. Permutation and padding happen in one kernel. +tokens_per_expert = routing_map.sum(dim=0) # [num_experts] +padded, permuted_probs, row_id_map, pad_offsets, padded_tokens_per_expert = ( + moe_permute_and_pad_with_probs( + tokens, + probs, + routing_map, + tokens_per_expert, + align_size=128, + ) +) + +# padded: [sum(padded_tokens_per_expert), hidden_size] +# pad_offsets: per-expert cumulative padding (None if already aligned) +# padded_tokens_per_expert: aligned per-expert counts -> m_splits for GroupedLinear + +# ... run the grouped MLP on `padded`, producing expert_out ... + +# Apply the permuted routing weights to the completed expert outputs. Since the +# weights are applied here, do not pass them to moe_unpermute as well. +expert_out = expert_out * permuted_probs[:, None] + +# Pass pad_offsets so token combine removes the padding it added, and +# restore_shape so the result has the original [num_tokens, hidden_size] shape. +output = moe_unpermute( + expert_out, + row_id_map, + restore_shape=tokens.shape, + pad_offsets=pad_offsets, +) +# END_MOE_PERMUTE_PAD_PYTORCH diff --git a/docs/features/mixture_of_experts/moe_permute_pytorch.py b/docs/features/mixture_of_experts/moe_permute_pytorch.py new file mode 100644 index 0000000000..cb5a4e263d --- /dev/null +++ b/docs/features/mixture_of_experts/moe_permute_pytorch.py @@ -0,0 +1,25 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +# START_MOE_PERMUTE_PYTORCH +import torch +from transformer_engine.pytorch import moe_permute + +# tokens: [num_tokens, hidden_size] +# routing_map: [num_tokens, num_experts] mask, 1 if token routed to expert +# +# num_out_tokens is the number of rows in the permuted buffer. Reading it from +# the routing map (``int(routing_map.sum())``) triggers a device-to-host sync; +# when the value is known statically (e.g. ``num_tokens * top_k`` for dropless +# routing), prefer passing that constant directly. +permuted, row_id_map = moe_permute( + tokens, + routing_map, + num_out_tokens=int(routing_map.sum()), +) + +# permuted: [num_out_tokens, hidden_size], expert-contiguous +# row_id_map: opaque tensor used by ``moe_unpermute`` to reverse the +# permutation after the experts have run. +# END_MOE_PERMUTE_PYTORCH diff --git a/docs/features/mixture_of_experts/moe_unpermute_jax.py b/docs/features/mixture_of_experts/moe_unpermute_jax.py new file mode 100644 index 0000000000..960bae99f5 --- /dev/null +++ b/docs/features/mixture_of_experts/moe_unpermute_jax.py @@ -0,0 +1,21 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +# START_MOE_UNPERMUTE_JAX +from transformer_engine.jax import permutation as te_permutation + +# expert_out: [num_out_tokens, hidden_size], expert-contiguous, +# produced by grouped_dense (or a grouped MLP). +# row_id_map: returned by token_dispatch. +# router_probs: [num_tokens, num_experts]; the original (un-permuted) routing +# probabilities. Provide whenever the selected weights are not +# all one, including top-1 configurations with non-unit weights. +tokens_out = te_permutation.token_combine( + expert_out, + row_id_map, + merging_probs=router_probs, +) + +# tokens_out: [num_tokens, hidden_size], in the original token order +# END_MOE_UNPERMUTE_JAX diff --git a/docs/features/mixture_of_experts/moe_unpermute_pytorch.py b/docs/features/mixture_of_experts/moe_unpermute_pytorch.py new file mode 100644 index 0000000000..b66d0186a3 --- /dev/null +++ b/docs/features/mixture_of_experts/moe_unpermute_pytorch.py @@ -0,0 +1,25 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +# START_MOE_UNPERMUTE_PYTORCH +from transformer_engine.pytorch import moe_unpermute + +# expert_out: [num_out_tokens, hidden_size], expert-contiguous, +# produced by GroupedLinear (or a grouped MLP). +# row_id_map: returned by moe_permute. +# merging_probs: [num_tokens, num_experts]; routing probabilities used to +# weight the per-expert contributions to each token. They may +# be omitted only when every selected routing weight is one. +# restore_shape: the original [num_tokens, hidden_size]. Required when the +# permuted buffer has more rows than the input (top-k routing); +# for top-1 it can be omitted and is inferred from expert_out. +tokens_out = moe_unpermute( + expert_out, + row_id_map, + merging_probs=merging_probs, + restore_shape=(num_tokens, hidden_size), +) + +# tokens_out: [num_tokens, hidden_size], in the original token order +# END_MOE_UNPERMUTE_PYTORCH diff --git a/docs/features/mixture_of_experts/router_jax.py b/docs/features/mixture_of_experts/router_jax.py new file mode 100644 index 0000000000..bbe0ba8bb4 --- /dev/null +++ b/docs/features/mixture_of_experts/router_jax.py @@ -0,0 +1,45 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +# START_ROUTER_JAX +from transformer_engine.jax.router import fused_topk_with_score_function + +# logits: [num_tokens, num_experts], produced by the gating (router) projection. +# +# Select the top-k experts for each token. The score function and the top-k +# selection run in a single fused kernel. Most arguments have defaults, so a +# basic call only needs the logits, topk and score_function. +score_function = "softmax" # "softmax", "sigmoid" or "sqrtsoftplus" +probs, routing_map = fused_topk_with_score_function( + logits, + topk=2, + score_function=score_function, +) + +# probs: [num_tokens, num_experts], non-zero only at the selected experts. +# Pass to token_combine as merging_probs. +# routing_map: [num_tokens, num_experts] bool mask. Cast to int32 for token_dispatch. +# END_ROUTER_JAX + + +# START_ROUTER_AUX_JAX +from transformer_engine.jax.router import fused_moe_aux_loss + +# Use dense scores, but keep the actual router's map for counts: auxiliary mode +# ignores expert bias, grouping and scaling. +scores, _ = fused_topk_with_score_function( + logits, + topk=2, + score_function=score_function, + compute_aux_scores=True, +) +tokens_per_expert = routing_map.sum(axis=0) # [num_experts] + +aux_loss = fused_moe_aux_loss( + scores, + tokens_per_expert, + topk=2, + coeff=1e-2, # loss weight; add aux_loss to the training loss +) +# END_ROUTER_AUX_JAX diff --git a/docs/features/mixture_of_experts/router_pytorch.py b/docs/features/mixture_of_experts/router_pytorch.py new file mode 100644 index 0000000000..c2188ecabc --- /dev/null +++ b/docs/features/mixture_of_experts/router_pytorch.py @@ -0,0 +1,55 @@ +# Copyright (c) 2022-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. +# +# See LICENSE for license information. + +# START_ROUTER_PYTORCH +from transformer_engine.pytorch.router import fused_topk_with_score_function + +# logits: [num_tokens, num_experts], produced by the gating (router) projection. +# +# Select the top-k experts for each token and return their routing weights. The +# score function and the top-k selection run in a single fused kernel (all math +# is done in fp32 internally for numerical stability). +score_function = "softmax" # "softmax", "sigmoid" or "sqrtsoftplus" +probs, routing_map = fused_topk_with_score_function( + logits, + topk=2, + use_pre_softmax=False, # softmax after top-k; True selects softmax-then-top-k + num_groups=None, # set with group_topk to enable grouped (device-limited) routing + group_topk=None, + scaling_factor=None, # optional scalar multiplied into the returned probs + score_function=score_function, + expert_bias=None, # [num_experts] selection bias for sigmoid or sqrtsoftplus +) + +# probs: [num_tokens, num_experts], non-zero only at the selected experts. +# Pass directly to moe_unpermute as merging_probs. +# routing_map: [num_tokens, num_experts] bool mask, True at the selected experts. +# Cast to int32 and pass to moe_permute. +# END_ROUTER_PYTORCH + + +# START_ROUTER_AUX_PYTORCH +from transformer_engine.pytorch.router import ( + fused_compute_score_for_moe_aux_loss, + fused_moe_aux_loss, +) + +# Use dense scores, but keep the actual router's map for counts: this helper +# does not apply expert bias or grouped routing. +_, scores = fused_compute_score_for_moe_aux_loss( + logits, + topk=2, + score_function=score_function, +) +tokens_per_expert = routing_map.sum(dim=0) # [num_experts] + +aux_loss = fused_moe_aux_loss( + scores, + tokens_per_expert, + total_num_tokens=logits.shape[0], + num_experts=logits.shape[1], + topk=2, + coeff=1e-2, # loss weight; add aux_loss to the training loss +) +# END_ROUTER_AUX_PYTORCH diff --git a/docs/index.rst b/docs/index.rst index 68060974f4..0d3746fcc2 100644 --- a/docs/index.rst +++ b/docs/index.rst @@ -55,6 +55,7 @@ Transformer Engine documentation :caption: Features features/low_precision_training/index.rst + features/mixture_of_experts/mixture_of_experts.rst features/other_optimizations/index.rst diff --git a/transformer_engine/jax/moe.py b/transformer_engine/jax/moe.py index c672db2fec..e697ab3b57 100644 --- a/transformer_engine/jax/moe.py +++ b/transformer_engine/jax/moe.py @@ -1272,21 +1272,23 @@ def moe( the capacity used by ``ep_bootstrap``. Overflow is reported through ``total_recv_tokens`` when bootstrap used ``drop_on_overflow=True``. - Note that the per-expert dispatch-slot alignment is fixed internally + Notes + ----- + The per-expert dispatch-slot alignment is fixed internally at 128 tokens (``_ALIGN_SIZE``); see that constant's docstring for rationale and how to extend if a future recipe needs >128. Axis-name parameters: - * ``ep_axis`` and ``data_parallelism_axes`` are *physical mesh - axis names* -- they index ``jax.sharding.Mesh.shape`` directly + * ``ep_axis`` and ``data_parallelism_axes`` are physical mesh + axis names: they index ``jax.sharding.Mesh.shape`` directly (to compute ``num_ep`` / ``dp_size`` and to construct ``P((dp..., ep), None, None)`` for the physical ``jax.lax.with_sharding_constraint`` calls that JAX requires to refer to real mesh axes). * ``input_axes``, ``gate_kernel_axes``, ``wi_kernel_axes``, - ``wo_kernel_axes`` are *logical axis names* (e.g. - ``"batch"``, ``"embed"``, ``"mlp"``, ``"exp"``) -- they get + ``wo_kernel_axes`` are logical axis names (e.g. + ``"batch"``, ``"embed"``, ``"mlp"``, ``"exp"``): they get resolved via the active Flax logical-axis rules and consumed by ``with_sharding_constraint_by_logical_axes``. They are ``Optional[str]`` tuples so a rule of ``None`` means diff --git a/transformer_engine/jax/router.py b/transformer_engine/jax/router.py index cd5c1dec6d..0778155e14 100644 --- a/transformer_engine/jax/router.py +++ b/transformer_engine/jax/router.py @@ -145,14 +145,15 @@ def fused_topk_with_score_function( Returns ------- probs_or_scores : jnp.ndarray - When compute_aux_scores=False: Sparse probability tensor, shape [num_tokens, num_experts]. - Non-zero only at selected expert positions. - When compute_aux_scores=True: Dense score tensor, shape [num_tokens, num_experts]. - All expert positions contain scores. + When ``compute_aux_scores=False``: sparse probability tensor, shape + ``[num_tokens, num_experts]``, non-zero only at selected expert positions. + When ``compute_aux_scores=True``: dense score tensor, shape + ``[num_tokens, num_experts]``, all expert positions contain scores. routing_map : jnp.ndarray - Shape/dtype depend on routing_map_format: - - BYTEMAP: bool[num_tokens, num_experts]; True at selected expert positions. - - BITMAP_U8: uint8[num_tokens, ceil(num_experts/8)]; LSB-first bit-packed + Shape/dtype depend on ``routing_map_format``: + + - BYTEMAP: ``bool[num_tokens, num_experts]``; True at selected expert positions. + - BITMAP_U8: ``uint8[num_tokens, ceil(num_experts/8)]``; LSB-first bit-packed along the expert axis. """ if not isinstance(scaling_factor, (int, float)): diff --git a/transformer_engine/pytorch/router.py b/transformer_engine/pytorch/router.py index 898310b0d6..946ea73da8 100644 --- a/transformer_engine/pytorch/router.py +++ b/transformer_engine/pytorch/router.py @@ -293,6 +293,7 @@ def fused_topk_with_score_function( ): """ Fused topk with score function router. + Parameters ---------- logits : torch.Tensor in fp32/bf16/fp16 @@ -337,9 +338,10 @@ def fused_topk_with_score_function( Same shape as ``logits``. routing_map : torch.Tensor Same leading dims as ``logits``; trailing dim and dtype depend on - routing_map_format, or dense top-k indices when topk_indices is provided: - - BYTEMAP: bool[*logits.shape[:-1], num_experts] - - BITMAP_U8: uint8[*logits.shape[:-1], ceil(num_experts/8)] + ``routing_map_format``, or dense top-k indices when ``topk_indices`` is provided: + + - BYTEMAP: ``bool[*logits.shape[:-1], num_experts]`` + - BITMAP_U8: ``uint8[*logits.shape[:-1], ceil(num_experts/8)]``, LSB-first bit-packed. """ if logits.dtype == torch.float64: @@ -537,6 +539,7 @@ def fused_moe_aux_loss( ) -> torch.Tensor: """ Fused MoE aux loss. + Parameters ---------- probs : torch.Tensor in fp32/bf16/fp16