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)
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 computesgroup @ 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: