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
BlockScaledMatmulAMD_PreB
struct BlockScaledMatmulAMD_PreB[BM: Int = Int(64), BN: Int = Int(128), BK_ELEMS: Int = Int(512), WN: Int = Int(64), b_prefetch: Bool = False, b_cache_policy: CacheOperation = CacheOperation.ALWAYS, dram_to_lds: Bool = False, cluster_drain_sched: Bool = False, mfma_cluster: Int = Int(4), pipeline_depth: Int = Int(2), scale_group: Int = Int(1), b_addr_split: Bool = False, matrix_format: CDNA4F8F6F4MatrixFormat = CDNA4F8F6F4MatrixFormat.FLOAT4_E2M1, elementwise_lambda_fn: Optional[def[dtype: DType, width: SIMDLength, *, alignment: Int = Int(1)](IndexList[Int(2)], SIMD[dtype, width]) capturing thin -> None] = None]
Preshuffled-B variant of BlockScaledMatmulAMD.
The preb path requires num_warps_m == 1 (no LDS staging for B = no
cross-warp M-direction B reuse), so WM is structurally fixed to BM.
When b_prefetch=True, runs a depth-2 outer-K software pipeline: while
the current iter's MFMAs execute, the next iter's B fragments stream
from DRAM into the alternate b_reg slot. Doubles _b_reg size (extra
VGPRs) but hides DRAM B latency across the inner MFMA chain. Targets
K-heavy shapes (e.g. gate/up, K=7168) where outer-iter serialization
dominates.
cluster_drain_sched (b_prefetch only) switches the 1-deep steady loop to
an interleaved B-issue schedule: the next tile's B fragments are issued
per-k between the current tile's MFMA phases (not front-loaded), each phase
pinned by sched_barrier(0) + bracketed by s_setprio, and the
end-of-tile sync is a bare s_barrier + lgkmcnt-only drain so in-flight
B DMAs cross it. (The epilogue fully drains vmcnt -- it has no newer ops
to stage a staircase against -- and keeps the per-cluster s_setprio
bracketing.) Default off: callers bit-identical unless opted in.
b_addr_split splits the B-fragment address into a loop-invariant
per-lane part (voffset) and a wave-uniform whole-tile part (soffset),
hoisting the per-tile v_add chain out of the unrolled K loop. Worth
3-4% on the wide MXFP8 gate+up tiles (BM 32..128) and costs 3-4% on the
BM=16 MXFP4 decode tiles, where there is almost no K loop to amortise the
scalar setup against -- so it is opted into per band, not global.
pipeline_depth sizes the B-fragment register ring (num_b_slots) and,
when > 2, switches the b_prefetch steady loop's end-of-iter sync to the
non-draining s_waitcnt[lgkmcnt=0] + bare s_barrier so in-flight B DMAs
are NOT drained every iteration. At the default (2) the ring is the same
2 slots and the draining barrier() is kept — bit-identical to before.
The deeper prefetch schedule that consumes slots >= 2 is a follow-up; this
param only plumbs the depth + the non-draining-barrier seam.
MFMA consumption order is B-major (n-outer / m-inner): the B fragment is
held resident across the m-loop for better MFMA ILP. See mma.
Parameters
- BM (
Int): CTA tile size along M in elements, either 16 or a multiple of 32.WMis locked toBM(single warp along M). - BN (
Int): CTA tile size along N in elements, split acrossnum_warps_n = BN // WNwarps. - BK_ELEMS (
Int): K tile size in MXFP4 elements per outer-K iteration; must be a multiple of 256 sonum_k_mmasis even.BK_BYTES = BK_ELEMS // 2. - WN (
Int): Per-warp tile size along N in elements, either 16 or a multiple of 32. - b_prefetch (
Bool): Enables a depth-2 outer-K software pipeline that double-buffers B fragments across iterations (defaults toFalse). - b_cache_policy (
CacheOperation):CacheOperationhint applied to preshuffled B DRAM loads (defaults toCacheOperation.ALWAYS). - dram_to_lds (
Bool): Routes A loads through the shared swizzledTileLoaderLDSDRAM-to-LDS path instead of a register bounce (defaults toFalse). - cluster_drain_sched (
Bool): Switches the prefetch steady loop to an interleaved B-issue schedule with per-clusters_setprioand a partial-vmcntstaircase (defaults toFalse). - mfma_cluster (
Int): Number of MFMAs per cluster in the scheduled MFMA chain used bycluster_drain_sched(defaults to 4). - pipeline_depth (
Int): Depth of the B-fragment register ring (num_b_slots) on the prefetch path (defaults to 2). When> 2it also switches the steady loop's end-of-iter sync to the non-drainings_waitcnt[lgkmcnt=0]+ bares_barrierand co-deepens the A/scale rings. At the default (2) the behavior is bit-identical to before. - scale_group (
Int): Number of consecutive outer-K tiles whose A+B scale atoms are fetched by one wide VMEM op per pack and bounced through LDS (1 = off). Trades VMEM instructions for LDS ops; the preconditions it needs are asserted below. - b_addr_split (
Bool): Whether to carry the B-fragment address as a loop-invariant per-lanevoffsetplus a wave-uniform whole-tilesoffset(defaults to False). Hoists the per-tile address chain out of the unrolled K loop; pays only where the warp tile is wide enough to amortise the scalar setup. - matrix_format (
CDNA4F8F6F4MatrixFormat):f8f6f4operand encoding for A and B. A lane covers 32 K-elements in every format; the bytes that occupies -- 16 (FP4), 24 (FP6), 32 (FP8) -- is derived from it. - elementwise_lambda_fn (
Optional[def[dtype: DType, width: SIMDLength, *, alignment: Int = Int(1)](IndexList[Int(2)], SIMD[dtype, width]) capturing thin -> None]): Optional epilogue applied to each output element in registers instead of storing it. The fused QKV ops use it to scatter K/V into the paged cache.
Implemented traits
comptime members
a_bits
comptime a_bits = matrix_format.bits_per_element()
A_SMEM_ROW_BYTES
comptime A_SMEM_ROW_BYTES = BlockScaledMatmulAMD_PreB[BM, BN, BK_ELEMS, WN, b_prefetch, b_cache_policy, dram_to_lds, cluster_drain_sched, mfma_cluster, pipeline_depth, scale_group, b_addr_split, matrix_format, elementwise_lambda_fn].MmaOpType.A_SMEM_ROW_BYTES
b_bits
comptime b_bits = matrix_format.bits_per_element()
B_BK_BYTES
comptime B_BK_BYTES = (Int((mul matrix_format.bits_per_element(), BK_ELEMS)) // Int(8))
bits_per_element
comptime bits_per_element = BlockScaledMatmulAMD_PreB[BM, BN, BK_ELEMS, WN, b_prefetch, b_cache_policy, dram_to_lds, cluster_drain_sched, mfma_cluster, pipeline_depth, scale_group, b_addr_split, matrix_format, elementwise_lambda_fn].a_bits
BK_BYTES
comptime BK_BYTES = (Int((mul matrix_format.bits_per_element(), BK_ELEMS)) // Int(8))
c_frag_size
comptime c_frag_size = BlockScaledMatmulAMD_PreB[BM, BN, BK_ELEMS, WN, b_prefetch, b_cache_policy, dram_to_lds, cluster_drain_sched, mfma_cluster, pipeline_depth, scale_group, b_addr_split, matrix_format, elementwise_lambda_fn].MmaOpType.c_frag_size
MMA_K
comptime MMA_K = 128
MMA_K_BYTES
comptime MMA_K_BYTES = BlockScaledMatmulAMD_PreB[BM, BN, BK_ELEMS, WN, b_prefetch, b_cache_policy, dram_to_lds, cluster_drain_sched, mfma_cluster, pipeline_depth, scale_group, b_addr_split, matrix_format, elementwise_lambda_fn].MmaOpType.MMA_K_BYTES
MMA_M
comptime MMA_M = 16
MMA_N
comptime MMA_N = 16
MmaOpType
comptime MmaOpType = BlockScaledMmaOp_PreB[IndexList(Int(16), Int(16), Int(128), __list_literal__=NoneType(None)), IndexList(BlockScaledMatmulAMD_PreB[BM, BN, BK_ELEMS, WN, b_prefetch, b_cache_policy, dram_to_lds, cluster_drain_sched, mfma_cluster, pipeline_depth, scale_group, b_addr_split, matrix_format, elementwise_lambda_fn].WM, WN, BK_ELEMS, __list_literal__=NoneType(None)), BlockScaledMatmulAMD_PreB[BM, BN, BK_ELEMS, WN, b_prefetch, b_cache_policy, dram_to_lds, cluster_drain_sched, mfma_cluster, pipeline_depth, scale_group, b_addr_split, matrix_format, elementwise_lambda_fn].num_b_slots, BlockScaledMatmulAMD_PreB[BM, BN, BK_ELEMS, WN, b_prefetch, b_cache_policy, dram_to_lds, cluster_drain_sched, mfma_cluster, pipeline_depth, scale_group, b_addr_split, matrix_format, elementwise_lambda_fn].num_scale_slots, scale_group, b_addr_split, matrix_format, not dram_to_lds]
num_a_load_slots
comptime num_a_load_slots = pipeline_depth if (pipeline_depth > Int(2)) else Int(1)
num_a_slots
comptime num_a_slots = Int(2) if b_prefetch else Int(1)
num_b_slots
comptime num_b_slots = pipeline_depth if b_prefetch else Int(1)
num_k_mmas
comptime num_k_mmas = BlockScaledMatmulAMD_PreB[BM, BN, BK_ELEMS, WN, b_prefetch, b_cache_policy, dram_to_lds, cluster_drain_sched, mfma_cluster, pipeline_depth, scale_group, b_addr_split, matrix_format, elementwise_lambda_fn].MmaOpType.num_k_mmas
num_m_mmas
comptime num_m_mmas = BlockScaledMatmulAMD_PreB[BM, BN, BK_ELEMS, WN, b_prefetch, b_cache_policy, dram_to_lds, cluster_drain_sched, mfma_cluster, pipeline_depth, scale_group, b_addr_split, matrix_format, elementwise_lambda_fn].MmaOpType.num_m_mmas
num_n_mmas
comptime num_n_mmas = BlockScaledMatmulAMD_PreB[BM, BN, BK_ELEMS, WN, b_prefetch, b_cache_policy, dram_to_lds, cluster_drain_sched, mfma_cluster, pipeline_depth, scale_group, b_addr_split, matrix_format, elementwise_lambda_fn].MmaOpType.num_n_mmas
num_scale_slots
comptime num_scale_slots = pipeline_depth if (pipeline_depth > Int(2)) else Int(1)
num_threads
comptime num_threads = (BlockScaledMatmulAMD_PreB[BM, BN, BK_ELEMS, WN, b_prefetch, b_cache_policy, dram_to_lds, cluster_drain_sched, mfma_cluster, pipeline_depth, scale_group, b_addr_split, matrix_format, elementwise_lambda_fn].num_warps * _resolve_warp_size())
num_warps
comptime num_warps = BlockScaledMatmulAMD_PreB[BM, BN, BK_ELEMS, WN, b_prefetch, b_cache_policy, dram_to_lds, cluster_drain_sched, mfma_cluster, pipeline_depth, scale_group, b_addr_split, matrix_format, elementwise_lambda_fn].num_warps_n
num_warps_m
comptime num_warps_m = 1
num_warps_n
comptime num_warps_n = (BN // WN)
simd_width
comptime simd_width = simd_width_of[DType.uint8]()
WM
comptime WM = BM
Methods
run
static def run[out_dtype: DType, c_layout: TensorLayout, a_layout: TensorLayout, b_pre_layout: TensorLayout, sfa_layout: TensorLayout, sfb_layout: TensorLayout, c_engine: TensorEngine, a_engine: TensorEngine, b_pre_engine: TensorEngine, sfa_engine: TensorEngine, sfb_engine: TensorEngine, N: Int, K_BYTES: Int](c: TileTensor[out_dtype, c_layout, MutAnyOrigin, Engine=c_engine], a: TileTensor[.uint8, a_layout, ImmutAnyOrigin, Engine=a_engine], b_pre: TileTensor[.uint8, b_pre_layout, ImmutAnyOrigin, Engine=b_pre_engine], sfa: TileTensor[.float8_e8m0fnu, sfa_layout, ImmutAnyOrigin, Engine=sfa_engine], sfb: TileTensor[.float8_e8m0fnu, sfb_layout, ImmutAnyOrigin, Engine=sfb_engine], n_tile_idx: Int, m_tile_idx: Int)