IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /get-started.md). For the complete documentation index, see llms.txt.
Skip to main content
For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /get-started.md).

Python function

grouped_matmul_ragged

grouped_matmul_ragged()​

max.experimental.nn.common_layers.functional_kernels.grouped_matmul_ragged(hidden_states, weight, expert_start_indices, expert_ids, expert_usage_stats)

source

Performs the grouped matmul used in the MoE layer.

hidden_states and expert_start_indices are used together to implement the ragged tensor. expert_start_indices indicates where each group starts and ends in hidden_states.

expert_ids is the id of the expert for each group in hidden_states.

expert_usage_stats is a rank-1 uint32 tensor laid out as [max_tokens_per_expert, num_active_experts] (the output of moe_create_indices).

Parameters:

  • hidden_states (TensorValue) – The ragged input activations.
  • weight (TensorValue) – The expert weights, [num_experts, N, K] (Linear convention); each group computes group @ weight.T.
  • expert_start_indices (TensorValue) – The start index of each group in hidden_states.
  • expert_ids (TensorValue) – The id of the expert for each group.
  • expert_usage_stats (TensorValue) – The per-expert usage stats, the output of moe_create_indices.

Returns:

The ragged matmul output.

Return type:

Tensor