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 struct
BlockScaledMatmulConfig
struct BlockScaledMatmulConfig[a_type: DType, b_type: DType, c_type: DType, sfa_dtype: DType, sfb_dtype: DType, transpose_b: Bool = True]
Static configuration of GPU matmul.
Parameters
- a_type (
DType):DTypeof the A (left) operand elements;uint8indicates packed FP4; must equalb_type. - b_type (
DType):DTypeof the B (right) operand elements. - c_type (
DType):DTypeof the output matrix elements. - sfa_dtype (
DType):DTypeof the A operand block scaling factors; selects the block-scaled MMA kind. - sfb_dtype (
DType):DTypeof the B operand block scaling factors; must equalsfa_dtype. - transpose_b (
Bool): Whether the B operand is stored transposed (defaults toTrue).
Fields
- cta_group (
Int): - mma_shape (
IndexList[Int(3)]): - cluster_shape (
IndexList[Int(3)]): - AB_swapped (
Bool): - block_swizzle_size (
Int): - raster_order (
RasterOrder): - register_based_epilogue (
Bool): - block_tile_shape (
IndexList[Int(3)]): - num_split_k (
Int): - num_pipeline_stages (
Int): - num_clc_pipeline_stages (
Int): - num_accum_pipeline_stages (
Int): - num_output_stages (
Int): - output_tile_shape (
IndexList[Int(2)]): - a_swizzle (
TensorMapSwizzle): - b_swizzle (
TensorMapSwizzle): - c_swizzle (
TensorMapSwizzle): - k_group_size (
Int): - scaling_kind (
UMMAKind): - vec_sf_size (
Int): - num_sf_k_tiles (
Int): - is_small_bn (
Bool): - gemm_kind (
GEMMKind): - prefetch_tiles_n (
Int):
Implemented traits
AnyType,
Copyable,
Deinitable,
Equatable,
Hashable,
ImplicitlyCopyable,
Movable,
RegisterPassable,
TrivialRegisterPassable,
Writable
comptime members
accum_type
comptime accum_type = get_accum_type[a_type]()
sf_block_atom_size
comptime sf_block_atom_size = (((load_from_mem SF_ATOM_M.__getitem_param__[Int(0)]()) * (load_from_mem SF_ATOM_M.__getitem_param__[Int(1)]())) * Int(4))
Methods
__init__
def __init__(*, scaling_kind: UMMAKind, cta_group: Int = Int(2), mma_shape: IndexList[Int(3)] = get_mma_shape[a_type, BlockScaledMatmulConfig[a_type, b_type, c_type, sfa_dtype, sfb_dtype, transpose_b].accum_type](), cluster_shape: IndexList[Int(3)] = Index[Int, Int, Int](Int(2), Int(1), Int(1)), AB_swapped: Bool = False, num_split_k: Int = Int(1), block_swizzle_size: Int = Int(0), raster_order: RasterOrder = RasterOrder.AlongM, k_group_size: Int = Int(1), num_pipeline_stages: Optional[Int] = None, num_accum_pipeline_stages: Int = Int(2), num_clc_pipeline_stages: Int = Int(2), is_gmm: Bool = False, is_small_bn: Bool = False, register_based_epilogue: Bool = True, gemm_kind: GEMMKind = GEMMKind.GEMM, prefetch_tiles_n: Int = Int(0)) -> Self
swap_AB_type
def swap_AB_type(self) -> BlockScaledMatmulConfig[b_type, a_type, c_type, sfb_dtype, sfa_dtype, transpose_b]
Returns:
BlockScaledMatmulConfig[b_type, a_type, c_type, sfb_dtype, sfa_dtype, transpose_b]
write_to
def write_to[W: Writer](self, mut writer: W)
write_repr_to
def write_repr_to(self, mut writer: T)