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).

Mojo module

grouped_matmul

Apple M5 grouped (MoE) matmul for bf16 or fp16 operands: a grouped GEMV for decode and a grouped simdgroup-MMA GEMM for prefill.

Target: Apple M5 (compute_capability == 5, Metal 4).

For each expert group z, C[a_offsets[z]:a_offsets[z+1]] = A[a_offsets[z]:a_offsets[z+1]] @ b[expert_ids[z]]^T, with b the weight stack [num_experts, N, K] and fp32 accumulation. A and b share one 16-bit float type (bf16 or fp16). One dispatch covers every group (block_idx.z), as in matmul2d_fp8.Matmul2dFp8.run_grouped.

Two kernels, picked on the host by max_num_tokens_per_expert:

  • Decode (few tokens per expert): grouped_gemv_kernel. At 1-2 tokens per expert there is no matrix to feed the 16x16 MMA, and a 64-row MMA tile wastes most of its rows, so this is a register-resident GEMV that streams each active expert's weight slab once per pass. One simdgroup owns rows_per_sg consecutive weight rows; its 32 lanes stride down K with 16-byte loads, so a simdgroup reads rows_per_sg contiguous 512-byte runs per step. The tokens_per_pass tokens of a pass reuse the same weight registers, so a pass reads the weight once. The ceiling is weight-read bandwidth.
  • Prefill: grouped_matmul_mma_kernel drives the dense AppleM5MatMul._run_gemm_body (the NT 16-bit use_x2 configuration enqueue_apple_matmul picks) on per-group views, so the grouped path shares the dense kernel's MMA, Morton tile order, bounded edges and epilogue. The host dispatches a few groups at a time (see enqueue_apple_grouped_mma).

comptime values​

APPLE_GROUPED_GEMV_MAX_TOKENS​

comptime APPLE_GROUPED_GEMV_MAX_TOKENS = 16

Functions​

Was this page helpful?