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 ownsrows_per_sgconsecutive weight rows; its 32 lanes stride down K with 16-byte loads, so a simdgroup readsrows_per_sgcontiguous 512-byte runs per step. Thetokens_per_passtokens 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_kerneldrives the denseAppleM5MatMul._run_gemm_body(the NT 16-bituse_x2configurationenqueue_apple_matmulpicks) 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 (seeenqueue_apple_grouped_mma).
comptime values
APPLE_GROUPED_GEMV_MAX_TOKENS
comptime APPLE_GROUPED_GEMV_MAX_TOKENS = 16
Functions
-
enqueue_apple_grouped_gemv: Enqueues the grouped GEMV. -
enqueue_apple_grouped_matmul: Enqueues the Apple M5 grouped matmulC = A @ b[expert]^T. -
enqueue_apple_grouped_mma: Enqueues the grouped simdgroup-MMA GEMM (any token count). -
grouped_gemv_kernel: Grouped GEMV: one simdgroup perrows_per_sgoutput columns of a group. -
grouped_matmul_mma_kernel: Grouped GEMM: the denseAppleM5MatMulbody on groupfirst_group + block_idx.z.