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
multistage_gemm_q
def multistage_gemm_q[c_type: DType, a_type: DType, b_type: DType, //, *, group_size: Int, pack_factor: Int, config: MatmulConfig[a_type, b_type, c_type, True], 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, 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, 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, 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], runtime_config: MatmulConfig[a_type, b_type, c_type, True], ctx: DeviceContext)
Enqueues the multi-stage quantized GEMM kernel, reducing pipeline stages or warp partitions when the shared memory budget is exceeded.
Parameters:
- c_type (
DType): The dtype of the output matrix. - a_type (
DType): The dtype of the A matrix elements. - b_type (
DType): The dtype of the packed B weight buffer. - group_size (
Int): The number of K elements sharing a single scale. - pack_factor (
Int): The number of 4-bit elements packed into oneuint32. - config (
MatmulConfig[a_type, b_type, c_type, True]): The matmul configuration describing tile and warp shapes. - elementwise_lambda_fn (
Optional[def[dtype: DType, width: SIMDLength, *, alignment: Int = Int(1)](IndexList[Int(2)], SIMD[dtype, width]) capturing thin -> None]): An optional elementwise epilogue applied per output element.
Args:
- c (
LayoutTensor[c_type, 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]): The output matrix in global memory. - a (
LayoutTensor[a_type, 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]): The left-hand (activation) matrix in global memory. - b (
LayoutTensor[b_type, 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]): The packed quantized weight buffer in global memory. - runtime_config (
MatmulConfig[a_type, b_type, c_type, True]): The runtime matmul configuration used for grid and block dimensions. - ctx (
DeviceContext): The device context used to enqueue the kernel.
Raises:
An error if the input tensors are not rank-2.