MoE Grouped Matmul
Overview
The MoE Grouped Matmul operation computes a grouped matrix multiplication across experts, as used in Mixture-of-Experts (MoE) layers. Each expert has its own weight matrix, and tokens are routed to experts via first_token_offset.
Three routing modes are supported:
None mode (tokens already routed per expert):
Gather mode (gather tokens from unrouted layout before matmul):
Scatter mode (scatter output back to token order after matmul):
where = number of experts, = number of tokens, = hidden size, = output (weight) size.
Tensor Roles by Mode
Support Matrix
The support matrix is based on the latest cuDNN backend.
Important Notes
FirstTokenOffsetcontainsB * Evalues with the total token count implicit from the token tensor dimension.- In Scatter mode, both
TokenIndexandTokenKsare required, andtop_kmust be explicitly provided. - In Gather mode,
TokenIndexis required.
MoE Grouped Matmul Forward
C++ API
Moe_grouped_matmul_attributes is a lightweight structure with setters:
Python API
Low-level graph API (cudnn.pygraph)
Args:
token(cudnn_tensor): Token data.- None/Scatter mode: shape
(1, S*topK, K) - Gather mode: shape
(1, S, K)
- None/Scatter mode: shape
weight(cudnn_tensor): Expert weight data with shape(E, K, N).first_token_offset(cudnn_tensor): INT32 tensor of shape(B*E, 1, 1). The -th entry is the index of the first token assigned to expert .token_index(Optional[cudnn_tensor]): INT32 tensor of shape(1, S*topK, 1). Maps each routed slot to a source token index. Required for Gather and Scatter modes.token_ks(Optional[cudnn_tensor]): INT32 tensor of shape(1, S*topK, 1). The expert index for each routed token. Required for Scatter mode.mode(cudnn.moe_grouped_matmul_mode): Routing mode —NONE,GATHER, orSCATTER.top_k(int): Top-k routing value. Must be provided for Scatter mode.compute_data_type(Optional[cudnn.data_type]): Data type for internal computation. Defaults to FLOAT.name(Optional[str]): Name for the operation.
Returns:
output(cudnn_tensor): Output tensor of shape(1, M_out, N), whereM_out = token_index.shape[1]for Gather mode, otherwisetoken.shape[1].
High-level experimental API
The high-level API handles cuDNN handle management and graph caching automatically. cuDNN graphs are built once per unique (shape, dtype, mode, top_k) configuration and reused across subsequent calls.
Configurable Options
-
Mode (
mode): Controls how tokens are routed to and from expert weight matrices.NONE: Tokens are already ordered by expert (pre-routed). Direct grouped matmul with no reordering.GATHER: Tokens are in original (un-routed) order.TokenIndexspecifies which source token each expert slot reads from.SCATTER: Tokens are pre-routed, but the output is scattered back to the original token order. Requires bothTokenIndexandTokenKs.
-
TopK (
top_k): The number of experts each token is routed to. Required in Scatter mode for the scatter-back computation. -
Compute data type (
compute_data_type): Sets the precision for internal accumulation. Defaults to FLOAT for fp16/bf16 I/O.
Example (Python)
MoE Grouped Matmul Backward
The backward operation computes the weight gradient given the upstream gradient and the forward token activations:
per expert, where the per-expert token slices are determined by FirstTokenOffset.
C++ API
Moe_grouped_matmul_bwd_attributes is a lightweight structure with setters:
Python API
Low-level graph API (cudnn.pygraph)
Args:
doutput(cudnn_tensor): Upstream gradient with shape(1, S, N), same layout as the forward output.token(cudnn_tensor): Forward token activations with shape(1, S, K).first_token_offset(cudnn_tensor): INT32 tensor of shape(B*E, 1, 1), same as used in the forward pass.compute_data_type(Optional[cudnn.data_type]): Data type for internal accumulation.name(Optional[str]): Name for the operation.
Returns:
dweight(cudnn_tensor): Weight gradient with shape(E, K, N).