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
matmul_gpu_qint4_impl
def matmul_gpu_qint4_impl[c_type: DType, a_type: DType, //, group_size: Int, target: StringSlice[ImmStaticOrigin], 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[DType.uint8, 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], ctx: Optional[DeviceContext])
Dispatches a GPU int4 quantized matrix multiplication to the tuned kernel configuration for the runtime M dimension.
Selects a MatmulConfig specialized for the static K and N dimensions and
the runtime M value, then delegates to multistage_gemm_q.
Parameters:
- βc_type (
DType): The dtype of the output matrix. - βa_type (
DType): The dtype of the A matrix elements. - βgroup_size (
Int): The number of K elements sharing a single scale. - βtarget (
StringSlice[ImmStaticOrigin]): The target platform string, which must identify a GPU. - β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[DType.uint8, 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. - βctx (
Optional[DeviceContext]): The device context used to enqueue the kernel.
Raises:
An error if the input tensors are not rank-2.
Was this page helpful?
Thank you! We'll create more content like this.
Thank you for helping us improve!