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

PreShuffledBGroupedGEMM

struct PreShuffledBGroupedGEMM[cu_count: Int, wg_per_cu: Int = Int(2), matrix_format: CDNA4F8F6F4MatrixFormat = CDNA4F8F6F4MatrixFormat.FLOAT4_E2M1, xcd_stripe: Int = Int(8)]

Grouped GEMM for MXFP4 on AMD CDNA4 with pre-shuffled weights.

This grouped GEMM operates on weights B that have been pre-shuffled into a layout enabling coalesced reads from shared memory and direct MFMA usage. It offers a persistent kernel (grid-stride over work tiles with XCD-aware work-group swizzling) and a direct kernel (one block per output tile, expert dispatched via block_idx.z), selected at launch time by the persistent comptime flag.

Parameters​

  • ​cu_count (Int): Number of compute units on the target device.
  • ​wg_per_cu (Int): Work groups per compute unit (default 2).
  • ​matrix_format (CDNA4F8F6F4MatrixFormat): f8f6f4 operand encoding for A and B (FP4 E2M1 by default). Both kernels below derive their fragment widths from it, so it must reach every BlockScaledMatmulAMD_PreB instantiation -- including the ones inside MAX_THREADS_PER_BLOCK_METADATA, or the launch bounds disagree with the body about num_threads.
  • ​xcd_stripe (Int): Size of the contiguous per-XCD run in to_swizzled_idx's logical-index space (default 8). 0 is the escape hatch: it resolves to wg_per_xcd, recovering the fully-contiguous pre-fix mapping. See to_swizzled_idx for why this bounds the persistent grid's remainder-tile imbalance across XCDs.

Implemented traits​

AnyType, Deinitable, Movable

comptime members​

a_bits​

comptime a_bits = matrix_format.bits_per_element()

b_bits​

comptime b_bits = matrix_format.bits_per_element()

bits_per_element​

comptime bits_per_element = PreShuffledBGroupedGEMM[cu_count, wg_per_cu, matrix_format, xcd_stripe].a_bits

effective_xcd_stripe​

comptime effective_xcd_stripe = (PreShuffledBGroupedGEMM[cu_count, wg_per_cu, matrix_format, xcd_stripe].total_wg // Int(8)) if (xcd_stripe == Int(0)) else xcd_stripe

fmt_suffix​

comptime fmt_suffix = StringSpan("e2m3") if (matrix_format == CDNA4F8F6F4MatrixFormat.FLOAT6_E2M3) else StringSpan(String("e3m2") if (matrix_format == CDNA4F8F6F4MatrixFormat.FLOAT6_E3M2) else String(""))

lane_bytes​

comptime lane_bytes = (Int((mul matrix_format.bits_per_element(), 32)) // Int(8))

num_xcd​

comptime num_xcd = 8

total_wg​

comptime total_wg = (cu_count * wg_per_cu)

wg_per_xcd​

comptime wg_per_xcd = (PreShuffledBGroupedGEMM[cu_count, wg_per_cu, matrix_format, xcd_stripe].total_wg // Int(8))

Methods​

to_swizzled_idx​

static def to_swizzled_idx(linear_idx: Int) -> Int

Returns:

Int

persistent_kernel​

static def persistent_kernel[BM: Int, BN: Int, BK_ELEMS: Int, WN: Int, out_dtype: DType, LayoutC: TensorLayout, LayoutA: TensorLayout, LayoutBPre: TensorLayout, LayoutSFA: TensorLayout, LayoutSFB: TensorLayout, AOffsetsLayout: TensorLayout, ExpertIdsLayout: TensorLayout, CEngine: TensorEngine, AEngine: TensorEngine, BPreEngine: TensorEngine, SFAEngine: TensorEngine, SFBEngine: TensorEngine, AOffsetsEngine: TensorEngine, ExpertIdsEngine: TensorEngine, N: Int, K_BYTES: Int, 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, waves_per_eu: Int = Int(0)](c_tensor: TileTensor[out_dtype, LayoutC, MutAnyOrigin, Engine=CEngine], a_tensor: TileTensor[.uint8, LayoutA, ImmutAnyOrigin, Engine=AEngine], b_pre_tensor: TileTensor[.uint8, LayoutBPre, ImmutAnyOrigin, Engine=BPreEngine], sfa_tensor: TileTensor[.float8_e8m0fnu, LayoutSFA, ImmutAnyOrigin, Engine=SFAEngine], sfb_tensor: TileTensor[.float8_e8m0fnu, LayoutSFB, ImmutAnyOrigin, Engine=SFBEngine], a_offsets: TileTensor[.uint32, AOffsetsLayout, ImmutAnyOrigin, Engine=AOffsetsEngine], expert_ids: TileTensor[.int32, ExpertIdsLayout, ImmutAnyOrigin, Engine=ExpertIdsEngine], num_active_experts: Int32, max_padded_M: Int32)

kernel​

static def kernel[BM: Int, BN: Int, BK_ELEMS: Int, WN: Int, out_dtype: DType, LayoutC: TensorLayout, LayoutA: TensorLayout, LayoutBPre: TensorLayout, LayoutSFA: TensorLayout, LayoutSFB: TensorLayout, AOffsetsLayout: TensorLayout, ExpertIdsLayout: TensorLayout, CEngine: TensorEngine, AEngine: TensorEngine, BPreEngine: TensorEngine, SFAEngine: TensorEngine, SFBEngine: TensorEngine, AOffsetsEngine: TensorEngine, ExpertIdsEngine: TensorEngine, N: Int, K_BYTES: Int, 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, waves_per_eu: Int = Int(0)](c_tensor: TileTensor[out_dtype, LayoutC, MutAnyOrigin, Engine=CEngine], a_tensor: TileTensor[.uint8, LayoutA, ImmutAnyOrigin, Engine=AEngine], b_pre_tensor: TileTensor[.uint8, LayoutBPre, ImmutAnyOrigin, Engine=BPreEngine], sfa_tensor: TileTensor[.float8_e8m0fnu, LayoutSFA, ImmutAnyOrigin, Engine=SFAEngine], sfb_tensor: TileTensor[.float8_e8m0fnu, LayoutSFB, ImmutAnyOrigin, Engine=SFBEngine], a_offsets: TileTensor[.uint32, AOffsetsLayout, ImmutAnyOrigin, Engine=AOffsetsEngine], expert_ids: TileTensor[.int32, ExpertIdsLayout, ImmutAnyOrigin, Engine=ExpertIdsEngine], num_active_experts: Int32, max_padded_M: Int32)

launch​

static def launch[BM: Int, BN: Int, BK_ELEMS: Int, WN: Int, persistent: Bool, 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, waves_per_eu: Int = Int(0), static_grid_z: Bool = False](c: TileTensor[Engine=c.Engine, address_space=c.address_space, linear_idx_type=c.linear_idx_type], a: TileTensor[.uint8, Engine=a.Engine, address_space=a.address_space, linear_idx_type=a.linear_idx_type], b_pre: TileTensor[.uint8, Engine=b_pre.Engine, address_space=b_pre.address_space, linear_idx_type=b_pre.linear_idx_type], a_scales: TileTensor[.float8_e8m0fnu, Engine=a_scales.Engine, address_space=a_scales.address_space, linear_idx_type=a_scales.linear_idx_type], b_scales: TileTensor[.float8_e8m0fnu, Engine=b_scales.Engine, address_space=b_scales.address_space, linear_idx_type=b_scales.linear_idx_type], a_offsets: TileTensor[.uint32, Engine=a_offsets.Engine, linear_idx_type=a_offsets.linear_idx_type], expert_ids: TileTensor[.int32, Engine=expert_ids.Engine, linear_idx_type=expert_ids.linear_idx_type], max_num_tokens_per_expert: Int, num_active_experts: Int, ctx: DeviceContext, grid_m_cap: Int = Int(-1))

Was this page helpful?