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
MXFP4MatmulAMD
struct MXFP4MatmulAMD[BM: Int = Int(128), BN: Int = Int(128), BK_ELEMS: Int = Int(128), WM: Int = Int(64), WN: Int = Int(64), MMA_M: Int = Int(16), MMA_N: Int = Int(16), MMA_K: Int = Int(128), lane_bytes: Int = Int(16)]
Native MXFP4 block-scaled matmul for AMD CDNA4.
Uses cdna4_block_scaled_mfma with FLOAT4_E2M1 format directly. Single-buffer pipeline with schedule-driven prologue/kernel/epilogue. SMEM is plain row-major (no blocked-product, no swizzle).
Parameters
- BM (
Int): Block tile rows (output M per block). Default 128. - BN (
Int): Block tile cols (output N per block). Default 128. - BK_ELEMS (
Int): Block tile K in logical FP4 elements. Default 128. - WM (
Int): Warp tile rows. BM must be divisible by WM. Default 64. - WN (
Int): Warp tile cols. BN must be divisible by WN. Default 64. - MMA_M (
Int): MFMA tile rows. WM must be divisible by MMA_M. Default 16. - MMA_N (
Int): MFMA tile cols. WN must be divisible by MMA_N. Default 16. - MMA_K (
Int): MFMA K-depth in logical FP4 elements. Default 128. - lane_bytes (
Int): Operand bytes per lane per MFMA: 16 (MXFP4, default) or 32 (MXFP8).BK_ELEMScounts ELEMENTS, so a givenBK_ELEMScosts twice the registers and LDS at MXFP8.
Implemented traits
comptime members
BK_BYTES
comptime BK_BYTES = (BK_ELEMS // MXFP4MatmulAMD[BM, BN, BK_ELEMS, WM, WN, MMA_M, MMA_N, MMA_K, lane_bytes].elems_per_byte)
c_frag_size
comptime c_frag_size = ((MMA_M * MMA_N) // _resolve_warp_size())
elems_per_byte
comptime elems_per_byte = (Int(32) // lane_bytes)
k_tile_size
comptime k_tile_size = MXFP4MatmulAMD[BM, BN, BK_ELEMS, WM, WN, MMA_M, MMA_N, MMA_K, lane_bytes].BK_BYTES
num_k_tiles
comptime num_k_tiles = (MXFP4MatmulAMD[BM, BN, BK_ELEMS, WM, WN, MMA_M, MMA_N, MMA_K, lane_bytes].BK_BYTES // MXFP4MatmulAMD[BM, BN, BK_ELEMS, WM, WN, MMA_M, MMA_N, MMA_K, lane_bytes].packed_k_per_mma)
num_m_mmas
comptime num_m_mmas = (WM // MMA_M)
num_n_mmas
comptime num_n_mmas = (WN // MMA_N)
num_threads
comptime num_threads = (MXFP4MatmulAMD[BM, BN, BK_ELEMS, WM, WN, MMA_M, MMA_N, MMA_K, lane_bytes].num_warps * _resolve_warp_size())
num_warps
comptime num_warps = (MXFP4MatmulAMD[BM, BN, BK_ELEMS, WM, WN, MMA_M, MMA_N, MMA_K, lane_bytes].num_warps_m * MXFP4MatmulAMD[BM, BN, BK_ELEMS, WM, WN, MMA_M, MMA_N, MMA_K, lane_bytes].num_warps_n)
num_warps_m
comptime num_warps_m = (BM // WM)
num_warps_n
comptime num_warps_n = (BN // WN)
packed_k_per_mma
comptime packed_k_per_mma = (MMA_K // MXFP4MatmulAMD[BM, BN, BK_ELEMS, WM, WN, MMA_M, MMA_N, MMA_K, lane_bytes].elems_per_byte)
scales_per_mma
comptime scales_per_mma = (MMA_K // Int(32))
simd_width
comptime simd_width = simd_width_of[DType.uint8]()
Methods
run
static def run[out_dtype: DType, c_layout: TensorLayout, a_layout: TensorLayout, b_layout: TensorLayout, sfa_layout: TensorLayout, sfb_layout: TensorLayout, num_splits: Int = Int(1)](c: TileTensor[out_dtype, c_layout, MutAnyOrigin], a: TileTensor[DType.uint8, a_layout, ImmutAnyOrigin], b: TileTensor[DType.uint8, b_layout, ImmutAnyOrigin], sfa: TileTensor[DType.float8_e8m0fnu, sfa_layout, ImmutAnyOrigin], sfb: TileTensor[DType.float8_e8m0fnu, sfb_layout, ImmutAnyOrigin])
MXFP4 block-scaled GEMM kernel with SMEM pipeline.
With num_splits > 1 this is the inter-block split-K body: each
block_idx.z slice accumulates one disjoint K-band into its own
[M, N] region of a stacked (num_splits * M, N) float32
workspace (out_dtype is float32 in that mode). A separate
reduce kernel sums the num_splits partials and casts to the
real output dtype. num_splits == 1 is byte-identical to the
no-split path (split_id == 0, full K range, zero offset).
Parameters:
- out_dtype (
DType): Element type of the output tensorc; must befloat32whennum_splits > 1. - c_layout (
TensorLayout): Compile-time layout of the output tensorc. - a_layout (
TensorLayout): Compile-time layout of the A operand. - b_layout (
TensorLayout): Compile-time layout of the B operand. - sfa_layout (
TensorLayout): Compile-time layout of the A scales tensorsfa. - sfb_layout (
TensorLayout): Compile-time layout of the B scales tensorsfb. - num_splits (
Int): Number of disjoint K-bands the K dimension is partitioned into (defaults to 1, no split).
Args:
- c (
TileTensor[out_dtype, c_layout, MutAnyOrigin]): Output matrix[M, N]of dtypeout_dtype; in split-K mode a stacked(num_splits * M, N)float32 workspace. - a (
TileTensor[DType.uint8, a_layout, ImmutAnyOrigin]): Packed A operand[M, K//2]uint8, two MXFP4 nibbles per byte. - b (
TileTensor[DType.uint8, b_layout, ImmutAnyOrigin]): Packed B operand[N, K//2]uint8, transposed with two MXFP4 nibbles per byte. - sfa (
TileTensor[DType.float8_e8m0fnu, sfa_layout, ImmutAnyOrigin]): A block scales[M, K//32]asfloat8_e8m0fnu, one scale per 32 MXFP4 elements. - sfb (
TileTensor[DType.float8_e8m0fnu, sfb_layout, ImmutAnyOrigin]): B block scales[N, K//32]asfloat8_e8m0fnu, one scale per 32 MXFP4 elements.