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 function
naive_blockwise_scaled_fp8_grouped_matmul
def naive_blockwise_scaled_fp8_grouped_matmul[c_type: DType, a_type: DType, b_type: DType, a_scales_type: DType, b_scales_type: DType, a_offsets_type: DType, expert_ids_type: DType, c_layout: Layout, a_layout: Layout, b_layout: Layout, a_scale_layout: Layout, b_scale_layout: Layout, a_offsets_layout: Layout, expert_ids_layout: Layout, //, BLOCK_DIM_N: Int = Int(32), BLOCK_DIM_M: Int = Int(16), transpose_b: Bool = True, scales_granularity_mnk: Optional[IndexList[Int(3)]] = None, elementwise_lambda_fn: Optional[def[dtype: DType, width: SIMDLength, *, alignment: Int = Int(1)](IndexList[Int(2)], SIMD[dtype, width]) capturing thin -> None] = None](c: LayoutTensor[c_type, c_layout, address_space=c.address_space, element_layout=c.element_layout, layout_int_type=c.layout_int_type, linear_idx_type=c.linear_idx_type, masked=c.masked, alignment=c.alignment], a: LayoutTensor[a_type, a_layout, address_space=a.address_space, element_layout=a.element_layout, layout_int_type=a.layout_int_type, linear_idx_type=a.linear_idx_type, masked=a.masked, alignment=a.alignment], b: LayoutTensor[b_type, b_layout, address_space=b.address_space, element_layout=b.element_layout, layout_int_type=b.layout_int_type, linear_idx_type=b.linear_idx_type, masked=b.masked, alignment=b.alignment], a_scales: LayoutTensor[a_scales_type, a_scale_layout, address_space=a_scales.address_space, element_layout=a_scales.element_layout, layout_int_type=a_scales.layout_int_type, linear_idx_type=a_scales.linear_idx_type, masked=a_scales.masked, alignment=a_scales.alignment], b_scales: LayoutTensor[b_scales_type, b_scale_layout, address_space=b_scales.address_space, element_layout=b_scales.element_layout, layout_int_type=b_scales.layout_int_type, linear_idx_type=b_scales.linear_idx_type, masked=b_scales.masked, alignment=b_scales.alignment], a_offsets: LayoutTensor[a_offsets_type, a_offsets_layout, address_space=a_offsets.address_space, element_layout=a_offsets.element_layout, layout_int_type=a_offsets.layout_int_type, linear_idx_type=a_offsets.linear_idx_type, masked=a_offsets.masked, alignment=a_offsets.alignment], expert_ids: LayoutTensor[expert_ids_type, expert_ids_layout, address_space=expert_ids.address_space, element_layout=expert_ids.element_layout, layout_int_type=expert_ids.layout_int_type, linear_idx_type=expert_ids.linear_idx_type, masked=expert_ids.masked, alignment=expert_ids.alignment], max_num_tokens_per_expert: Int, num_active_experts: Int, ctx: DeviceContext)
Dispatches the naive blockwise scaled FP8 grouped matmul kernel on the GPU.
Enqueues naive_blockwise_scaled_fp8_grouped_matmul_kernel with one
expert per grid-Z slice, tiling the per-expert M_local x N output
with BLOCK_DIM_M x BLOCK_DIM_N thread blocks.
Args:
- c (
LayoutTensor[c_type, c_layout, address_space=c.address_space, element_layout=c.element_layout, layout_int_type=c.layout_int_type, linear_idx_type=c.linear_idx_type, masked=c.masked, alignment=c.alignment]): Rank-2 output accumulator tensor holding all expert outputs. - a (
LayoutTensor[a_type, a_layout, address_space=a.address_space, element_layout=a.element_layout, layout_int_type=a.layout_int_type, linear_idx_type=a.linear_idx_type, masked=a.masked, alignment=a.alignment]): Rank-2 FP8 input matrix in K-major format. - b (
LayoutTensor[b_type, b_layout, address_space=b.address_space, element_layout=b.element_layout, layout_int_type=b.layout_int_type, linear_idx_type=b.linear_idx_type, masked=b.masked, alignment=b.alignment]): Rank-3 FP8 weight tensor indexed by expert, in K-major format. - a_scales (
LayoutTensor[a_scales_type, a_scale_layout, address_space=a_scales.address_space, element_layout=a_scales.element_layout, layout_int_type=a_scales.layout_int_type, linear_idx_type=a_scales.linear_idx_type, masked=a_scales.masked, alignment=a_scales.alignment]): Per-block scales forain M-major format. - b_scales (
LayoutTensor[b_scales_type, b_scale_layout, address_space=b_scales.address_space, element_layout=b_scales.element_layout, layout_int_type=b_scales.layout_int_type, linear_idx_type=b_scales.linear_idx_type, masked=b_scales.masked, alignment=b_scales.alignment]): Per-block scales forbindexed by expert. - a_offsets (
LayoutTensor[a_offsets_type, a_offsets_layout, address_space=a_offsets.address_space, element_layout=a_offsets.element_layout, layout_int_type=a_offsets.layout_int_type, linear_idx_type=a_offsets.linear_idx_type, masked=a_offsets.masked, alignment=a_offsets.alignment]): Prefix-sum offsets delimiting each expert's rows ina. - expert_ids (
LayoutTensor[expert_ids_type, expert_ids_layout, address_space=expert_ids.address_space, element_layout=expert_ids.element_layout, layout_int_type=expert_ids.layout_int_type, linear_idx_type=expert_ids.linear_idx_type, masked=expert_ids.masked, alignment=expert_ids.alignment]): Expert id (or-1to skip) for each grid-Z slice. - max_num_tokens_per_expert (
Int): Maximum row count assigned to any single expert. - num_active_experts (
Int): Number of active experts to dispatch. - ctx (
DeviceContext): Device context used to enqueue the kernel.
def naive_blockwise_scaled_fp8_grouped_matmul[c_type: DType, a_type: DType, b_type: DType, a_scales_type: DType, b_scales_type: DType, a_offsets_type: DType, expert_ids_type: DType, c_tt_layout: TensorLayout, a_tt_layout: TensorLayout, b_tt_layout: TensorLayout, a_scale_tt_layout: TensorLayout, b_scale_tt_layout: TensorLayout, a_offsets_tt_layout: TensorLayout, expert_ids_tt_layout: TensorLayout, //, BLOCK_DIM_N: Int = Int(32), BLOCK_DIM_M: Int = Int(16), transpose_b: Bool = True, scales_granularity_mnk: Optional[IndexList[Int(3)]] = None, 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, c_tt_layout, Engine=c.Engine, address_space=c.address_space, linear_idx_type=c.linear_idx_type], a: TileTensor[a_type, a_tt_layout, Engine=a.Engine, address_space=a.address_space, linear_idx_type=a.linear_idx_type], b: TileTensor[b_type, b_tt_layout, Engine=b.Engine, address_space=b.address_space, linear_idx_type=b.linear_idx_type], a_scales: TileTensor[a_scales_type, a_scale_tt_layout, Engine=a_scales.Engine, address_space=a_scales.address_space, linear_idx_type=a_scales.linear_idx_type], b_scales: TileTensor[b_scales_type, b_scale_tt_layout, Engine=b_scales.Engine, address_space=b_scales.address_space, linear_idx_type=b_scales.linear_idx_type], a_offsets: TileTensor[a_offsets_type, a_offsets_tt_layout, Engine=a_offsets.Engine, address_space=a_offsets.address_space, linear_idx_type=a_offsets.linear_idx_type], expert_ids: TileTensor[expert_ids_type, expert_ids_tt_layout, Engine=expert_ids.Engine, address_space=expert_ids.address_space, linear_idx_type=expert_ids.linear_idx_type], max_num_tokens_per_expert: Int, num_active_experts: Int, ctx: DeviceContext)
TileTensor overload of the naive blockwise scaled FP8 grouped matmul.
Bridges to the LayoutTensor implementation, which stays the reference until the grouped FP8 path is TileTensor-native.
Parameters:
- c_type (
DType): Element type of the output accumulator. - a_type (
DType): Element type of the FP8 input matrix. - b_type (
DType): Element type of the FP8 weight tensor. - a_scales_type (
DType): Element type of the per-block scales fora. - b_scales_type (
DType): Element type of the per-block scales forb. - a_offsets_type (
DType): Element type of the per-expert row offsets. - expert_ids_type (
DType): Element type of the expert id tensor. - c_tt_layout (
TensorLayout): Compile-timeTensorLayoutofc. - a_tt_layout (
TensorLayout): Compile-timeTensorLayoutofa. - b_tt_layout (
TensorLayout): Compile-timeTensorLayoutofb. - a_scale_tt_layout (
TensorLayout): Compile-timeTensorLayoutofa_scales. - b_scale_tt_layout (
TensorLayout): Compile-timeTensorLayoutofb_scales. - a_offsets_tt_layout (
TensorLayout): Compile-timeTensorLayoutofa_offsets. - expert_ids_tt_layout (
TensorLayout): Compile-timeTensorLayoutofexpert_ids. - BLOCK_DIM_N (
Int): Thread-block width over the N dimension. - BLOCK_DIM_M (
Int): Thread-block height over the M dimension. - transpose_b (
Bool):Truewhenbis stored transposed. - scales_granularity_mnk (
Optional[IndexList[Int(3)]]): Optional per-dimension scale block sizes. - elementwise_lambda_fn (
Optional[def[dtype: DType, width: SIMDLength, *, alignment: Int = Int(1)](IndexList[Int(2)], SIMD[dtype, width]) capturing thin -> None]): Optional elementwise epilogue.
Args:
- c (
TileTensor[c_type, c_tt_layout, Engine=c.Engine, address_space=c.address_space, linear_idx_type=c.linear_idx_type]): Rank-2 output accumulator tensor holding all expert outputs. - a (
TileTensor[a_type, a_tt_layout, Engine=a.Engine, address_space=a.address_space, linear_idx_type=a.linear_idx_type]): Rank-2 FP8 input matrix in K-major format. - b (
TileTensor[b_type, b_tt_layout, Engine=b.Engine, address_space=b.address_space, linear_idx_type=b.linear_idx_type]): Rank-3 FP8 weight tensor indexed by expert, in K-major format. - a_scales (
TileTensor[a_scales_type, a_scale_tt_layout, Engine=a_scales.Engine, address_space=a_scales.address_space, linear_idx_type=a_scales.linear_idx_type]): Per-block scales forain M-major format. - b_scales (
TileTensor[b_scales_type, b_scale_tt_layout, Engine=b_scales.Engine, address_space=b_scales.address_space, linear_idx_type=b_scales.linear_idx_type]): Per-block scales forbindexed by expert. - a_offsets (
TileTensor[a_offsets_type, a_offsets_tt_layout, Engine=a_offsets.Engine, address_space=a_offsets.address_space, linear_idx_type=a_offsets.linear_idx_type]): Prefix-sum offsets delimiting each expert's rows ina. - expert_ids (
TileTensor[expert_ids_type, expert_ids_tt_layout, Engine=expert_ids.Engine, address_space=expert_ids.address_space, linear_idx_type=expert_ids.linear_idx_type]): Expert id (or-1to skip) for each grid-Z slice. - max_num_tokens_per_expert (
Int): Maximum row count assigned to any single expert. - num_active_experts (
Int): Number of active experts to dispatch. - ctx (
DeviceContext): Device context used to enqueue the kernel.