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

BlockScaledMmaOp_PreB

struct BlockScaledMmaOp_PreB[mma_shape: IndexList[Int(3)], warp_tile: IndexList[Int(3)], num_b_slots: Int = Int(1), num_scale_slots: Int = Int(1), scale_group: Int = Int(1), b_addr_split: Bool = False, matrix_format: CDNA4F8F6F4MatrixFormat = CDNA4F8F6F4MatrixFormat.FLOAT4_E2M1, pad_a_row: Bool = True]

Per-warp register state + MFMA dispatch for the preb (preshuffled-B, preshuffled-scale) kernel.

warp_tile is the (M, N, K) region this warp computes per outer-K iteration, in the same element units as mma_shape. Per-warp MFMA counts are derived as warp_tile[i] // mma_shape[i].

Asserted in __init__: warp_tile[i] % mma_shape[i] == 0 per axis, and num_k_mmas % 2 == 0 (k_pack=2 cell halves). num_m_mmas / num_n_mmas may be odd; the constructor rotates the scale i32 per CTA so OPSEL keeps the same comptime formula. See module-level comment for the scale-cell byte ordering.

Parameters​

  • ​mma_shape (IndexList[Int(3)]): The MFMA instruction shape as (M, N, K) in MFMA-native element units, for example (16, 16, 128) for MXFP4.
  • ​warp_tile (IndexList[Int(3)]): The (M, N, K) region this warp computes per outer-K iteration, in the same element units as mma_shape. Per-warp MFMA counts are derived as warp_tile[i] // mma_shape[i].
  • ​num_b_slots (Int): Number of _b_reg slots for software pipelining (defaults to 1). Set to 2 to double-buffer B fragments across outer-K iterations.
  • ​num_scale_slots (Int): Number of A/B scale-register ring slots for software pipelining (defaults to 1). At the depth-2 default this is 1 (1-deep, unchanged behavior); the depth-3 co-deepened prefetch sets it to 2 so the scale ring double-buffers alongside the B fragments.
  • ​scale_group (Int): Outer-K tiles whose scale atoms are fetched by one wide VMEM op (1 = off, per-tile buffer_load_dword).
  • ​b_addr_split (Bool): Whether to split the B-fragment address into a loop-invariant per-lane part and a wave-uniform whole-tile part (defaults to False, one address per fragment).
  • ​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.
  • ​pad_a_row (Bool): Whether to round A's LDS row stride up to a power of two (defaults to True). A no-op at MXFP4/MXFP8; the caller sets it False when a loader writes the tile at the payload stride.

Implemented traits​

AnyType, Deinitable, Movable

comptime members​

a_bits​

comptime a_bits = matrix_format.bits_per_element()

a_frag_width_bytes​

comptime a_frag_width_bytes = ((Int((mul mma_shape[Int(2)], matrix_format.bits_per_element())) // Int(8)) // (_resolve_warp_size() // mma_shape[Int(0)]))

A_K_HALF_STRIDE​

comptime A_K_HALF_STRIDE = ((Int((mul mma_shape[Int(2)], matrix_format.bits_per_element())) // Int(8)) // (((Int((mul mma_shape[Int(2)], matrix_format.bits_per_element())) // Int(8)) // (_resolve_warp_size() // mma_shape[Int(0)])) // Int(16)))

A_MMA_K_BYTES​

comptime A_MMA_K_BYTES = (Int((mul mma_shape[Int(2)], matrix_format.bits_per_element())) // Int(8))

a_num_frag_halves​

comptime a_num_frag_halves = (((Int((mul mma_shape[Int(2)], matrix_format.bits_per_element())) // Int(8)) // (_resolve_warp_size() // mma_shape[Int(0)])) // Int(16))

a_reg_frag_bytes​

comptime a_reg_frag_bytes = SIMD(matrix_format.simd_width())

a_scale_packs​

comptime a_scale_packs = ceildiv((warp_tile[Int(0)] // mma_shape[Int(0)]), Int(2))

A_SMEM_ROW_BYTES​

comptime A_SMEM_ROW_BYTES = _smem_row_bytes[(Int((mul warp_tile[Int(2)], matrix_format.bits_per_element())) // Int(8))]() if pad_a_row else BlockScaledMmaOp_PreB[mma_shape, warp_tile, num_b_slots, num_scale_slots, scale_group, b_addr_split, matrix_format, pad_a_row].BK_BYTES

b_bits​

comptime b_bits = matrix_format.bits_per_element()

b_frag_width_bytes​

comptime b_frag_width_bytes = ((Int((mul mma_shape[Int(2)], matrix_format.bits_per_element())) // Int(8)) // (_resolve_warp_size() // mma_shape[Int(0)]))

B_K_HALF_STRIDE​

comptime B_K_HALF_STRIDE = ((Int((mul mma_shape[Int(2)], matrix_format.bits_per_element())) // Int(8)) // (((Int((mul mma_shape[Int(2)], matrix_format.bits_per_element())) // Int(8)) // (_resolve_warp_size() // mma_shape[Int(0)])) // Int(16)))

B_MMA_K_BYTES​

comptime B_MMA_K_BYTES = (Int((mul mma_shape[Int(2)], matrix_format.bits_per_element())) // Int(8))

b_num_frag_halves​

comptime b_num_frag_halves = (((Int((mul mma_shape[Int(2)], matrix_format.bits_per_element())) // Int(8)) // (_resolve_warp_size() // mma_shape[Int(0)])) // Int(16))

b_reg_frag_bytes​

comptime b_reg_frag_bytes = SIMD(matrix_format.simd_width())

b_scale_packs​

comptime b_scale_packs = ceildiv((warp_tile[Int(1)] // mma_shape[Int(1)]), Int(2))

bits_per_element​

comptime bits_per_element = BlockScaledMmaOp_PreB[mma_shape, warp_tile, num_b_slots, num_scale_slots, scale_group, b_addr_split, matrix_format, pad_a_row].a_bits

BK_BYTES​

comptime BK_BYTES = (Int((mul warp_tile[Int(2)], matrix_format.bits_per_element())) // Int(8))

c_frag_size​

comptime c_frag_size = (Int((mul mma_shape[Int(0)], mma_shape[Int(1)])) // _resolve_warp_size())

FRAG_HALF_BYTES​

comptime FRAG_HALF_BYTES = 16

K_HALF_STRIDE​

comptime K_HALF_STRIDE = BlockScaledMmaOp_PreB[mma_shape, warp_tile, num_b_slots, num_scale_slots, scale_group, b_addr_split, matrix_format, pad_a_row].A_K_HALF_STRIDE

lane_bytes​

comptime lane_bytes = BlockScaledMmaOp_PreB[mma_shape, warp_tile, num_b_slots, num_scale_slots, scale_group, b_addr_split, matrix_format, pad_a_row].mma_frag_width_bytes

lanes_per_row​

comptime lanes_per_row = (_resolve_warp_size() // mma_shape[Int(0)])

mma_frag_width_bytes​

comptime mma_frag_width_bytes = BlockScaledMmaOp_PreB[mma_shape, warp_tile, num_b_slots, num_scale_slots, scale_group, b_addr_split, matrix_format, pad_a_row].a_frag_width_bytes

MMA_K​

comptime MMA_K = mma_shape[Int(2)]

MMA_K_BYTES​

comptime MMA_K_BYTES = BlockScaledMmaOp_PreB[mma_shape, warp_tile, num_b_slots, num_scale_slots, scale_group, b_addr_split, matrix_format, pad_a_row].A_MMA_K_BYTES

MMA_M​

comptime MMA_M = mma_shape[Int(0)]

MMA_N​

comptime MMA_N = mma_shape[Int(1)]

mx_format​

comptime mx_format = matrix_format

num_frag_halves​

comptime num_frag_halves = BlockScaledMmaOp_PreB[mma_shape, warp_tile, num_b_slots, num_scale_slots, scale_group, b_addr_split, matrix_format, pad_a_row].a_num_frag_halves

num_k_mmas​

comptime num_k_mmas = (warp_tile[Int(2)] // mma_shape[Int(2)])

num_m_mmas​

comptime num_m_mmas = (warp_tile[Int(0)] // mma_shape[Int(0)])

num_n_mmas​

comptime num_n_mmas = (warp_tile[Int(1)] // mma_shape[Int(1)])

reg_frag_bytes​

comptime reg_frag_bytes = BlockScaledMmaOp_PreB[mma_shape, warp_tile, num_b_slots, num_scale_slots, scale_group, b_addr_split, matrix_format, pad_a_row].a_reg_frag_bytes

SCALE_ATOM_BYTES​

comptime SCALE_ATOM_BYTES = Int(256)

scale_packs​

comptime scale_packs = (ceildiv((warp_tile[Int(0)] // mma_shape[Int(0)]), Int(2)) + ceildiv((warp_tile[Int(1)] // mma_shape[Int(1)]), Int(2)))

Methods​

__init__​

def __init__(out self, warp_m_off: Int, warp_n_off: Int)

accum_tile​

def accum_tile(self) -> ref[self._c_reg] TileTensor[.float32, Layout[TypeList[ComptimeInt[(warp_tile[Int(0)] // mma_shape[Int(0)])], ComptimeInt[((warp_tile[Int(1)] // mma_shape[Int(1)]) * (Int((mul mma_shape[Int(0)], mma_shape[Int(1)])) // _resolve_warp_size()))]](), TypeList[ComptimeInt[Int((mul (warp_tile[Int(1)] // mma_shape[Int(1)]), (Int((mul mma_shape[Int(0)], mma_shape[Int(1)])) // _resolve_warp_size())))], ComptimeInt[Int(1)]]()], MutUntrackedOrigin, address_space=AddressSpace.LOCAL]

Returns:

ref[self._c_reg] TileTensor[.float32, Layout[TypeList[ComptimeInt[(warp_tile[Int(0)] // mma_shape[Int(0)])], ComptimeInt[((warp_tile[Int(1)] // mma_shape[Int(1)]) * (Int((mul mma_shape[Int(0)], mma_shape[Int(1)])) // _resolve_warp_size()))]](), TypeList[ComptimeInt[Int((mul (warp_tile[Int(1)] // mma_shape[Int(1)]), (Int((mul mma_shape[Int(0)], mma_shape[Int(1)])) // _resolve_warp_size())))], ComptimeInt[Int(1)]]()], MutUntrackedOrigin, address_space=AddressSpace.LOCAL]

load_a_frag_from_smem​

def load_a_frag_from_smem[mma_k_idx: Int](self, a_smem_warp: TileTensor[.uint8, Engine=a_smem_warp.Engine, address_space=AddressSpace.SHARED, linear_idx_type=a_smem_warp.linear_idx_type])

Loads the A fragment for MFMA-K position mma_k_idx from row-major SMEM.

XOR-16 swizzled read (matches the write in copy_a_tile_to_smem): each lane reads the 16B vec at slot-tile (row, col_byte), then swizzles the flat in-tile byte offset before the raw_load. WM==BM so a_smem_warp IS the contiguous [BM, A_SMEM_ROW_BYTES] slot tile and raw_load indexes it directly.

Parameters:

  • ​mma_k_idx (Int): Index of the MFMA step along K within the warp tile, in [0, num_k_mmas).

Args:

load_b_frag_preshuffled​

def load_b_frag_preshuffled[B_N: Int, B_K_BYTES: Int, b_cache: CacheOperation, b_lane_bytes: Int, //, mma_k_idx: Int, slot: Int = Int(0)](self, b_loader: PreshuffledBLoader[B_N, B_K_BYTES, b_cache, b_lane_bytes], warp_n_off: Int, k_byte_base: Int)

Loads B fragments direct from preshuffled DRAM into b_reg slot slot.

Parameters:

  • ​mma_k_idx (Int): Index of the MFMA step along K within the warp tile, in [0, num_k_mmas).
  • ​slot (Int): The _b_reg slot to load into (defaults to 0).

Args:

load_a_scales_preshuffled​

def load_a_scales_preshuffled[k_pair: Int, slot: Int = Int(0)](mut self, a_scale_loader: PreshuffledScaleLoader, warp_m_off: Int, k_pair_idx: Int)

Issues per-lane i32 scale loads for A at one k_pair slot.

Caller provides the absolute k_pair_idx (= k_iter * (num_k_mmas / 2) + k_pair); each step advances by 8 K-scales (= 2 MFMAs along K). One i32 per (mi_pair, k_pair) per lane.

Parameters:

  • ​k_pair (Int): Index of the k_pair slot within the current outer-K tile, in [0, num_k_mmas // 2).
  • ​slot (Int): Scale-ring slot (comptime; 0 at the 1-deep default).

Args:

  • ​a_scale_loader (PreshuffledScaleLoader): The PreshuffledScaleLoader for the A scale tensor.
  • ​warp_m_off (Int): Global M offset of this warp's tile.
  • ​k_pair_idx (Int): Absolute k_pair index across all outer-K iterations, equal to k_iter * (num_k_mmas / 2) + k_pair; each step advances by 8 K-scales.

load_b_scales_preshuffled​

def load_b_scales_preshuffled[k_pair: Int, slot: Int = Int(0)](mut self, b_scale_loader: PreshuffledScaleLoader, warp_n_off: Int, k_pair_idx: Int)

Mirror of load_a_scales_preshuffled along N.

Parameters:

  • ​k_pair (Int): Index of the k_pair slot within the current outer-K tile, in [0, num_k_mmas // 2).
  • ​slot (Int): Scale-ring slot (comptime; 0 at the 1-deep default).

Args:

  • ​b_scale_loader (PreshuffledScaleLoader): The PreshuffledScaleLoader for the B scale tensor.
  • ​warp_n_off (Int): Global N offset of this warp's tile.
  • ​k_pair_idx (Int): Absolute k_pair index across all outer-K iterations, equal to k_iter * (num_k_mmas / 2) + k_pair; each step advances by 8 K-scales.

stage_scale_group​

def stage_scale_group(mut self, a_scale_loader: PreshuffledScaleLoader, b_scale_loader: PreshuffledScaleLoader, warp_m_off: Int, warp_n_off: Int, k_pair_base: Int)

Fetches scale_group outer-K tiles of A+B scales in one op per pack.

A wave64 VMEM op costs the same whether it presents 4 or 16 bytes per lane, so this cuts scale VMEM instructions by scale_group.

Issue it at the top of a tile and publish_scale_group at the bottom, so the A landing registers' existing vmcnt covers the fetch.

Args:

  • ​a_scale_loader (PreshuffledScaleLoader): The PreshuffledScaleLoader for A scales.
  • ​b_scale_loader (PreshuffledScaleLoader): The PreshuffledScaleLoader for B scales.
  • ​warp_m_off (Int): Global M offset of this warp's tile.
  • ​warp_n_off (Int): Global N offset of this warp's tile.
  • ​k_pair_base (Int): First k_pair of the window.

publish_scale_group​

def publish_scale_group(self, scale_smem: TileTensor[.uint8, Engine=scale_smem.Engine, address_space=AddressSpace.SHARED, linear_idx_type=scale_smem.linear_idx_type], warp_id: Int)

Writes the staged window to this warp's LDS strip, one op per pack.

load_group delivers the window lane-transposed; bouncing it through LDS undoes that far more cheaply than the VMEM ops it replaces.

Args:

read_scale_group​

def read_scale_group[phase: Int, slot: Int = Int(0)](mut self, scale_smem: TileTensor[.uint8, Engine=scale_smem.Engine, address_space=AddressSpace.SHARED, linear_idx_type=scale_smem.linear_idx_type], warp_id: Int)

Reads tile phase of the published window into the scale registers.

Parameters:

  • ​phase (Int): Position of this outer-K tile inside the window, in [0, scale_group).
  • ​slot (Int): Scale-ring slot to fill (0 at the 1-deep default).

Args:

mma​

def mma[mma_k_idx: Int, slot: Int = Int(0), scale_slot: Int = Int(0)](self)

Block-scaled MFMA at MFMA-K position mma_k_idx using B from slot.

B-major / n-outer / m-inner: hoist the B fragment + b_byte + b_scale (VMEM-loaded _b_reg) once per n, then cycle A (m-inner, LDS-loaded _a_reg). Keeping B resident across the m-loop improves MFMA ILP.

OPSEL byte selection from the 2x2 cell: a_byte = (mma_k_idx % 2) * 2 + (m % 2) b_byte = (mma_k_idx % 2) * 2 + (n % 2) Scale dword lives at _*_scale_packed[mn // 2, mma_k_idx // 2]. WM/WN=16 CTAs see only m=0 / n=0, so the constructor shrui (_a_scale_shift / _b_scale_shift) rotates the i32 to the right OPSEL byte.

Parameters:

  • ​mma_k_idx (Int): Index of the MFMA step along K within the warp tile, in [0, num_k_mmas).
  • ​slot (Int): The _b_reg slot to read B fragments from (defaults to 0).
  • ​scale_slot (Int): The scale-ring slot to read A/B scales from (defaults to 0; non-zero only under the depth-3 co-deepened prefetch).

Was this page helpful?