IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /get-started.md). For the complete documentation index, see llms.txt.
Skip to main content
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. WM is locked to BM (single warp along M).
  • ​BN (Int): CTA tile size along N in elements, split across num_warps_n = BN // WN warps.
  • ​BK_ELEMS (Int): K tile size in MXFP4 elements per outer-K iteration; must be a multiple of 256 so num_k_mmas is 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 to False).
  • ​b_cache_policy (CacheOperation): CacheOperation hint applied to preshuffled B DRAM loads (defaults to CacheOperation.ALWAYS).
  • ​dram_to_lds (Bool): Routes A loads through the shared swizzled TileLoaderLDS DRAM-to-LDS path instead of a register bounce (defaults to False).
  • ​cluster_drain_sched (Bool): Switches the prefetch steady loop to an interleaved B-issue schedule with per-cluster s_setprio and a partial-vmcnt staircase (defaults to False).
  • ​mfma_cluster (Int): Number of MFMAs per cluster in the scheduled MFMA chain used by cluster_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 > 2 it also switches the steady loop's end-of-iter sync to the non-draining s_waitcnt[lgkmcnt=0] + bare s_barrier and 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-lane voffset plus a wave-uniform whole-tile soffset (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): f8f6f4 operand 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​

AnyType, Deinitable, Movable

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)

Was this page helpful?