IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /max/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. /max/get-started.md).

Mojo function

grouped_matmul_sm90

def grouped_matmul_sm90[c_type: DType, a_type: DType, b_type: DType, //, *, transpose_b: Bool = True, wgmma_shape: IndexList[Int(3)] = Index[Int, Int, Int](Int(64), Int(256), Int(16)), config: MatmulConfig[a_type, b_type, c_type, transpose_b] = default_config_sm90[a_type, b_type, c_type, transpose_b, wgmma_shape](), elementwise_lambda_fn: Optional[def[dtype: DType, width: SIMDLength, *, alignment: Int = Int(1)](IndexList[Int(2)], SIMD[dtype, width]) capturing thin -> None] = None](c: TileTensor[c_type, Storage=c.Storage, linear_idx_type=c.linear_idx_type], a: TileTensor[a_type, Storage=a.Storage, linear_idx_type=a.linear_idx_type], a_offsets: TileTensor[DType.uint32, Storage=a_offsets.Storage, linear_idx_type=a_offsets.linear_idx_type], max_num_tokens_per_expert: Int, b: TileTensor[b_type, Storage=b.Storage, linear_idx_type=b.linear_idx_type], expert_ids: TileTensor[DType.int32, Storage=expert_ids.Storage, linear_idx_type=expert_ids.linear_idx_type], num_active_experts: Int, ctx: DeviceContext)

Performs grouped GEMM for MoE routing on SM90 (Hopper) GPUs.

Dispatches a batched expert matmul where each expert has a variable number of tokens, stored contiguously in a at offsets given by a_offsets. Expert weight matrices are stacked along axis 0 in b. Uses TMA-based warp-specialized pipelining via HopperMatmulSM90Kernel.

Parameters:

Args:

Was this page helpful?