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

BlackwellBlockScaledMatmulKernel

struct BlackwellBlockScaledMatmulKernel[a_type: DType, b_type: DType, c_type: DType, sfa_dtype: DType, sfb_dtype: DType, transpose_b: Bool, config: BlockScaledMatmulConfig[a_type, b_type, c_type, sfa_dtype, sfb_dtype, transpose_b], cluster_shape: StaticTuple[Int32, Int(3)] = StaticTuple(Int32(1)), elementwise_compute_lambda_fn: Optional[def[dtype: DType, width: SIMDLength, *, alignment: Int = Int(1)](IndexList[Int(2)], SIMD[dtype, width]) capturing thin -> SIMD[dtype, width]] = None, pdl_level: PDLLevel = PDLLevel(), max_profiled_tiles_per_SM: UInt32 = UInt32(0)]

Block-scaled matmul kernel V3 - ported from working legacy kernel.

This struct provides the structured interface while internally using the proven legacy kernel logic.

Parameters

Implemented traits

AnyType, Deinitable, Movable

comptime members

a_expected_bytes

comptime a_expected_bytes = (Int((mul config.block_tile_shape[Int(0)], config.block_tile_shape[Int(2)])) * size_of[a_type]())

a_internal_layout

comptime a_internal_layout = Layout(Coord(Coord(ComptimeInt(), ComptimeInt()), Coord(ComptimeInt(), ComptimeInt())), Coord(Coord(ComptimeInt(), ComptimeInt()), Coord(ComptimeInt(), ComptimeInt())))

a_smem_layout

comptime a_smem_layout = Layout(Coord(Coord(ComptimeInt(), ComptimeInt()), Coord(ComptimeInt(), ComptimeInt())), Coord(Coord(ComptimeInt(), ComptimeInt()), Coord(ComptimeInt(), ComptimeInt()))).to_layout()

a_swizzle_elems

comptime a_swizzle_elems = (config.a_swizzle.bytes() // size_of[a_type]())

a_tile_dim0

comptime a_tile_dim0 = compute_tma_tile_dims[config.block_tile_shape[Int(0)], config.block_tile_shape[Int(1)], config.mma_shape[Int(0)], config.output_tile_shape[Int(0)], config.cluster_shape[Int(0)], config.cluster_shape[Int(1)], config.cta_group, config.AB_swapped]()[Int(0)]

a_tma_load_size

comptime a_tma_load_size = (BlackwellBlockScaledMatmulKernel[a_type, b_type, c_type, sfa_dtype, sfb_dtype, transpose_b, config, cluster_shape, elementwise_compute_lambda_fn, pdl_level, max_profiled_tiles_per_SM].a_tile_dim0 * (config.a_swizzle.bytes() // size_of[a_type]()))

a_tma_rows

comptime a_tma_rows = BlackwellBlockScaledMatmulKernel[a_type, b_type, c_type, sfa_dtype, sfb_dtype, transpose_b, config, cluster_shape, elementwise_compute_lambda_fn, pdl_level, max_profiled_tiles_per_SM].a_tile_dim0

accum_pipeline_consumer_arv_count

comptime accum_pipeline_consumer_arv_count = compute_accum_barrier_counts[Int((mul _resolve_warp_size(), 4)), config.cta_group]()[Int(1)]

accum_pipeline_producer_arv_count

comptime accum_pipeline_producer_arv_count = compute_accum_barrier_counts[Int((mul _resolve_warp_size(), 4)), config.cta_group]()[Int(0)]

accum_type

comptime accum_type = DType.float32

ADescLayout

comptime ADescLayout = Layout[*(), *()]

ATileLayout

comptime ATileLayout = Layout[*(), *()]

ATmaOp

comptime ATmaOp = TMATensorTile[a_type, Int(3), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(3), Layout[*(), *()]]()]

b_expected_bytes

comptime b_expected_bytes = (Int((mul config.block_tile_shape[Int(1)], config.block_tile_shape[Int(2)])) * size_of[b_type]())

b_internal_layout

comptime b_internal_layout = Layout(Coord(Coord(ComptimeInt(), ComptimeInt()), Coord(ComptimeInt(), ComptimeInt())), Coord(Coord(ComptimeInt(), ComptimeInt()), Coord(ComptimeInt(), ComptimeInt())))

b_smem_layout

comptime b_smem_layout = Layout(Coord(Coord(ComptimeInt(), ComptimeInt()), Coord(ComptimeInt(), ComptimeInt())), Coord(Coord(ComptimeInt(), ComptimeInt()), Coord(ComptimeInt(), ComptimeInt()))).to_layout() if transpose_b else Layout(Coord(Coord(ComptimeInt(), ComptimeInt()), Coord(ComptimeInt(), ComptimeInt())), Coord(Coord(ComptimeInt(), ComptimeInt()), Coord(ComptimeInt(), ComptimeInt()))).transpose().to_layout()

b_swizzle_elems

comptime b_swizzle_elems = (config.b_swizzle.bytes() // size_of[b_type]())

b_tile_dim0

comptime b_tile_dim0 = compute_tma_tile_dims[config.block_tile_shape[Int(0)], config.block_tile_shape[Int(1)], config.mma_shape[Int(0)], config.output_tile_shape[Int(0)], config.cluster_shape[Int(0)], config.cluster_shape[Int(1)], config.cta_group, config.AB_swapped]()[Int(1)]

b_tma_load_size

comptime b_tma_load_size = (BlackwellBlockScaledMatmulKernel[a_type, b_type, c_type, sfa_dtype, sfb_dtype, transpose_b, config, cluster_shape, elementwise_compute_lambda_fn, pdl_level, max_profiled_tiles_per_SM].b_tile_dim0 * (config.b_swizzle.bytes() // size_of[b_type]()))

b_tma_rows

comptime b_tma_rows = BlackwellBlockScaledMatmulKernel[a_type, b_type, c_type, sfa_dtype, sfb_dtype, transpose_b, config, cluster_shape, elementwise_compute_lambda_fn, pdl_level, max_profiled_tiles_per_SM].b_tile_dim0

BDescLayout

comptime BDescLayout = Layout[*(), *()]

BK

comptime BK = config.block_tile_shape[Int(2)]

BM

comptime BM = config.block_tile_shape[Int(0)]

BN

comptime BN = config.block_tile_shape[Int(1)]

BTileLayout

comptime BTileLayout = Layout[*(), *()]

BTmaOp

comptime BTmaOp = TMATensorTile[b_type, Int(3), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(3), Layout[*(), *()]]()]

c_smem_layout

comptime c_smem_layout = Layout.row_major(config.output_tile_shape[Int(0)], config.output_tile_shape[Int(1)])

c_swizzle_elems

comptime c_swizzle_elems = (config.c_swizzle.bytes() // size_of[c_type]())

c_tile_dim0

comptime c_tile_dim0 = compute_tma_tile_dims[config.block_tile_shape[Int(0)], config.block_tile_shape[Int(1)], config.mma_shape[Int(0)], config.output_tile_shape[Int(0)], config.cluster_shape[Int(0)], config.cluster_shape[Int(1)], config.cta_group, config.AB_swapped]()[Int(2)]

c_tile_dim1

comptime c_tile_dim1 = BlackwellBlockScaledMatmulKernel[a_type, b_type, c_type, sfa_dtype, sfb_dtype, transpose_b, config, cluster_shape, elementwise_compute_lambda_fn, pdl_level, max_profiled_tiles_per_SM].OutputN if not config.AB_swapped.__bool__() else BlackwellBlockScaledMatmulKernel[a_type, b_type, c_type, sfa_dtype, sfb_dtype, transpose_b, config, cluster_shape, elementwise_compute_lambda_fn, pdl_level, max_profiled_tiles_per_SM].c_swizzle_elems

CDescLayout

comptime CDescLayout = Layout[*(), *()]

clc_consumer_arv_count

comptime clc_consumer_arv_count = compute_clc_barrier_counts[_resolve_warp_size(), _resolve_warp_size(), _resolve_warp_size(), Int((mul _resolve_warp_size(), 4)), Int((mul config.cluster_shape[Int(0)], config.cluster_shape[Int(1)])), config.cta_group]()[Int(1)]

clc_producer_arv_count

comptime clc_producer_arv_count = compute_clc_barrier_counts[_resolve_warp_size(), _resolve_warp_size(), _resolve_warp_size(), Int((mul _resolve_warp_size(), 4)), Int((mul config.cluster_shape[Int(0)], config.cluster_shape[Int(1)])), config.cta_group]()[Int(0)]

clc_throttle_consumer_arv_count

comptime clc_throttle_consumer_arv_count = compute_clc_barrier_counts[_resolve_warp_size(), _resolve_warp_size(), _resolve_warp_size(), Int((mul _resolve_warp_size(), 4)), Int((mul config.cluster_shape[Int(0)], config.cluster_shape[Int(1)])), config.cta_group]()[Int(3)]

clc_throttle_producer_arv_count

comptime clc_throttle_producer_arv_count = compute_clc_barrier_counts[_resolve_warp_size(), _resolve_warp_size(), _resolve_warp_size(), Int((mul _resolve_warp_size(), 4)), Int((mul config.cluster_shape[Int(0)], config.cluster_shape[Int(1)])), config.cta_group]()[Int(2)]

CLUSTER_M

comptime CLUSTER_M = config.cluster_shape[Int(0)]

CLUSTER_N

comptime CLUSTER_N = config.cluster_shape[Int(1)]

CLUSTER_SIZE

comptime CLUSTER_SIZE = (config.cluster_shape[Int(0)] * config.cluster_shape[Int(1)])

Context

comptime Context = KernelContext[config.num_clc_pipeline_stages, config.cta_group, config.cluster_shape[Int(0)], config.cluster_shape[Int(1)]]

cta_group

comptime cta_group = config.cta_group

CTileLayout

comptime CTileLayout = Layout[*(), *()]

CTmaOp

comptime CTmaOp = TMATensorTile[c_type, Int(3), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(3), Layout[*(), *()]]()]

EPILOGUE_THREADS

comptime EPILOGUE_THREADS = (Int(4) * _resolve_warp_size())

EpilogueCtx

comptime EpilogueCtx = EpilogueWarpContext[BlackwellBlockScaledMatmulKernel[a_type, b_type, c_type, sfa_dtype, sfb_dtype, transpose_b, config, cluster_shape, elementwise_compute_lambda_fn, pdl_level, max_profiled_tiles_per_SM].opc, _resolve_warp_size(), Int((mul _resolve_warp_size(), 4))]

input_expected_bytes

comptime input_expected_bytes = (Int((add (mul config.block_tile_shape[Int(0)], config.block_tile_shape[Int(2)], size_of[a_type](), config.cta_group), (mul config.block_tile_shape[Int(1)], config.block_tile_shape[Int(2)], size_of[b_type](), config.cta_group), (mul tile_sf_layout_k_major[config.block_tile_shape[Int(0)], Int((mul config.num_sf_k_tiles, config.vec_sf_size, 4)), config.vec_sf_size]().size(), size_of[sfa_dtype](), config.cta_group), (mul tile_sf_layout_k_major[config.mma_shape[Int(1)], Int((mul config.num_sf_k_tiles, config.vec_sf_size, 4)), config.vec_sf_size]().size(), size_of[sfb_dtype](), config.cta_group))) * config)

InputTilePipeline

comptime InputTilePipeline = InputTilePipeline[BlockScaledTilePayload[a_type, b_type, sfa_dtype, sfb_dtype, IndexList(config.block_tile_shape[Int(0)], config.block_tile_shape[Int(2)], __list_literal__=NoneType(None)), IndexList(config.block_tile_shape[Int(1)], config.block_tile_shape[Int(2)], __list_literal__=NoneType(None)), IndexList(sfa_dim0[config](), sfa_dim1[config](), __list_literal__=NoneType(None)), IndexList(sfb_dim0[config](), sfb_dim1[config](), __list_literal__=NoneType(None)), config.num_pipeline_stages], (config // config), config.k_group_size]

MMA_K

comptime MMA_K = config.mma_shape[Int(2)]

MMA_M

comptime MMA_M = config.mma_shape[Int(0)]

MMA_N

comptime MMA_N = config.mma_shape[Int(1)]

MMA_THREADS

comptime MMA_THREADS = WARP_SIZE

MmaCtx

comptime MmaCtx = MmaWarpContext[BlackwellBlockScaledMatmulKernel[a_type, b_type, c_type, sfa_dtype, sfb_dtype, transpose_b, config, cluster_shape, elementwise_compute_lambda_fn, pdl_level, max_profiled_tiles_per_SM].opc, _resolve_warp_size(), Int((mul _resolve_warp_size(), 4))]

MmaEpilogueSync

comptime MmaEpilogueSync = WarpGroupBarrier[(_resolve_warp_size() + Int((mul _resolve_warp_size(), 4))), Int(1)]

MmaOp

comptime MmaOp = MmaOpSM100_BlockScaled_SS[c_type, a_type, b_type, sfa_dtype, sfb_dtype, config.scaling_kind, config.block_tile_shape, config.mma_shape, cta_group=config.cta_group, cluster_shape=config.cluster_shape, a_swizzle=config.a_swizzle, b_swizzle=config.b_swizzle, transpose_b=transpose_b]

num_accum_pipeline_stages

comptime num_accum_pipeline_stages = config.num_accum_pipeline_stages

num_clc_pipeline_stages

comptime num_clc_pipeline_stages = config.num_clc_pipeline_stages

num_group_pipeline_stages

comptime num_group_pipeline_stages = (config // config)

num_output_stages

comptime num_output_stages = config.num_output_stages

num_output_warps

comptime num_output_warps = 4

num_pipeline_stages

comptime num_pipeline_stages = config.num_pipeline_stages

NUM_THREADS

comptime NUM_THREADS = (Int((mul _resolve_warp_size(), 3)) + Int((mul _resolve_warp_size(), 4)))

NUM_TMEM_COLS

comptime NUM_TMEM_COLS = 512

opc

comptime opc = OutputPipelineConfig(BlackwellBlockScaledMatmulKernel[a_type, b_type, c_type, sfa_dtype, sfb_dtype, transpose_b, config, cluster_shape, elementwise_compute_lambda_fn, pdl_level, max_profiled_tiles_per_SM].num_accum_pipeline_stages, BlackwellBlockScaledMatmulKernel[a_type, b_type, c_type, sfa_dtype, sfb_dtype, transpose_b, config, cluster_shape, elementwise_compute_lambda_fn, pdl_level, max_profiled_tiles_per_SM].stage_stride_cols, BlackwellBlockScaledMatmulKernel[a_type, b_type, c_type, sfa_dtype, sfb_dtype, transpose_b, config, cluster_shape, elementwise_compute_lambda_fn, pdl_level, max_profiled_tiles_per_SM].cta_group)

OutputM

comptime OutputM = config.output_tile_shape[Int(0)]

OutputN

comptime OutputN = config.output_tile_shape[Int(1)]

OutputPipeline

comptime OutputPipeline = OutputTilePipeline[BlackwellBlockScaledMatmulKernel[a_type, b_type, c_type, sfa_dtype, sfb_dtype, transpose_b, config, cluster_shape, elementwise_compute_lambda_fn, pdl_level, max_profiled_tiles_per_SM].opc]

Scheduler

comptime Scheduler = TileScheduler[config.num_clc_pipeline_stages, Index[Int, Int, Int, dtype=DType.uint32](config.cluster_shape[Int(0)], config.cluster_shape[Int(1)], config.cluster_shape[Int(2)]), config.raster_order, config.block_swizzle_size]

SCHEDULER_THREADS

comptime SCHEDULER_THREADS = WARP_SIZE

SF_K_GROUP_SIZE

comptime SF_K_GROUP_SIZE = (Int(4) * config)

sfa_expected_bytes

comptime sfa_expected_bytes = (tile_sf_layout_k_major[config.block_tile_shape[Int(0)], Int((mul config.num_sf_k_tiles, config.vec_sf_size, 4)), config.vec_sf_size]().size() * size_of[sfa_dtype]())

SFA_NUM_COLS

comptime SFA_NUM_COLS = (config * (config.block_tile_shape[Int(0)] // Int(32)))

sfa_smem_layout

comptime sfa_smem_layout = tile_sf_layout_k_major[config.block_tile_shape[Int(0)], (Int((mul config.vec_sf_size, 4)) * config), config.vec_sf_size]()

SFADescLayout

comptime SFADescLayout = Layout[*(), *()]

SFATileLayout

comptime SFATileLayout = Layout[*(), *()]

SFATmaOp

comptime SFATmaOp = TMATensorTile[sfa_dtype, Int(5), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(5), Layout[*(), *()]]()]

sfb_expected_bytes

comptime sfb_expected_bytes = (tile_sf_layout_k_major[config.mma_shape[Int(1)], Int((mul config.num_sf_k_tiles, config.vec_sf_size, 4)), config.vec_sf_size]().size() * size_of[sfb_dtype]())

SFB_NUM_COLS

comptime SFB_NUM_COLS = (config * (config.mma_shape[Int(1)] // Int(32)))

sfb_smem_layout

comptime sfb_smem_layout = tile_sf_layout_k_major[config.mma_shape[Int(1)], (Int((mul config.vec_sf_size, 4)) * config), config.vec_sf_size]()

SFBDescLayout

comptime SFBDescLayout = Layout[*(), *()]

SFBTileLayout

comptime SFBTileLayout = Layout[*(), *()]

SFBTmaOp

comptime SFBTmaOp = TMATensorTile[sfb_dtype, Int(5), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(5), Layout[*(), *()]]()]

SmemType

comptime SmemType = BlockScaledSmem[a_type, b_type, c_type, sfa_dtype, sfb_dtype, transpose_b, config=config]

stage_stride_cols

comptime stage_stride_cols = BlackwellBlockScaledMatmulKernel[a_type, b_type, c_type, sfa_dtype, sfb_dtype, transpose_b, config, cluster_shape, elementwise_compute_lambda_fn, pdl_level, max_profiled_tiles_per_SM].MMA_N

TilePayload

comptime TilePayload = BlockScaledTilePayload[a_type, b_type, sfa_dtype, sfb_dtype, IndexList(config.block_tile_shape[Int(0)], config.block_tile_shape[Int(2)], __list_literal__=NoneType(None)), IndexList(config.block_tile_shape[Int(1)], config.block_tile_shape[Int(2)], __list_literal__=NoneType(None)), IndexList(sfa_dim0[config](), sfa_dim1[config](), __list_literal__=NoneType(None)), IndexList(sfb_dim0[config](), sfb_dim1[config](), __list_literal__=NoneType(None)), config.num_pipeline_stages]

TileWriterType

comptime TileWriterType = TileWriter[a_type, DType.float32, config.block_tile_shape, config.mma_shape, BlackwellBlockScaledMatmulKernel[a_type, b_type, c_type, sfa_dtype, sfb_dtype, transpose_b, config, cluster_shape, elementwise_compute_lambda_fn, pdl_level, max_profiled_tiles_per_SM].opc, config.c_swizzle, config.AB_swapped, config.output_tile_shape[Int(0)], config.output_tile_shape[Int(1)], config.num_output_stages, Int(4), batched=True]

TMA_LOAD_THREADS

comptime TMA_LOAD_THREADS = WARP_SIZE

Tmem

comptime Tmem = TmemAllocation[OutputPipelineConfig(config.num_accum_pipeline_stages, config.mma_shape[Int(1)], config.cta_group).cta_group]

TmemDealloc

comptime TmemDealloc = TmemDeallocBarrier[OutputPipelineConfig(config.num_accum_pipeline_stages, config.mma_shape[Int(1)], config.cta_group).cta_group]

Methods

load_input_tiles

static def load_input_tiles[tiles_origin: MutOrigin, //](a_tma_op: TMATensorTile[a_type, Int(3), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(3), Layout[*(), *()]]()], b_tma_op: TMATensorTile[b_type, Int(3), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(3), Layout[*(), *()]]()], sfa_tma_op: TMATensorTile[sfa_dtype, Int(5), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(5), Layout[*(), *()]]()], sfb_tma_op: TMATensorTile[sfb_dtype, Int(5), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(5), Layout[*(), *()]]()], tiles: ProducerTiles[tiles_origin, BlockScaledTilePayload[a_type, b_type, sfa_dtype, sfb_dtype, IndexList(config.block_tile_shape[Int(0)], config.block_tile_shape[Int(2)], __list_literal__=NoneType(None)), IndexList(config.block_tile_shape[Int(1)], config.block_tile_shape[Int(2)], __list_literal__=NoneType(None)), IndexList(sfa_dim0[config](), sfa_dim1[config](), __list_literal__=NoneType(None)), IndexList(sfb_dim0[config](), sfb_dim1[config](), __list_literal__=NoneType(None)), config.num_pipeline_stages], (config // config), config.k_group_size], peer_cta_coord: Tuple[Int, Int, Int], work_tile_coord: Tuple[Int, Int, Int], a_multicast_mask: UInt16, b_multicast_mask: UInt16, iter_idx: UInt32, elect_one_cta: Bool)

Load A, B, SFA, SFB tiles using TMA with ProducerTiles.

This method uses the structured ProducerStage pattern from matmul_kernels.mojo, with tiles and barrier encapsulated in the stage.

Parameters:

  • tiles_origin (MutOrigin): Memory origin for the producer tiles (inferred).

Args:

mma

static def mma[tiles_origin: MutOrigin, //](tiles: ConsumerTiles[tiles_origin, BlockScaledTilePayload[a_type, b_type, sfa_dtype, sfb_dtype, IndexList(config.block_tile_shape[Int(0)], config.block_tile_shape[Int(2)], __list_literal__=NoneType(None)), IndexList(config.block_tile_shape[Int(1)], config.block_tile_shape[Int(2)], __list_literal__=NoneType(None)), IndexList(sfa_dim0[config](), sfa_dim1[config](), __list_literal__=NoneType(None)), IndexList(sfb_dim0[config](), sfb_dim1[config](), __list_literal__=NoneType(None)), config.num_pipeline_stages], (config // config), config.k_group_size], mma_op: MmaOpSM100_BlockScaled_SS[c_type, a_type, b_type, sfa_dtype, sfb_dtype, config.scaling_kind, config.block_tile_shape, config.mma_shape, cta_group=config.cta_group, cluster_shape=config.cluster_shape, a_swizzle=config.a_swizzle, b_swizzle=config.b_swizzle, transpose_b=transpose_b], tmem_addr: UInt32, sfa_tmem: UInt32, sfb_tmem: UInt32, iter_idx: UInt32, k_start: UInt32)

Execute MMA operations using ConsumerTiles.

This method uses the structured ConsumerStage pattern from matmul_kernels.mojo, with tiles and barrier encapsulated in the stage.

Parameters:

  • tiles_origin (MutOrigin): Memory origin for the consumer tiles (inferred).

Args:

epilogue

static def epilogue(c_tiles: SMemTileArray2DRowMajor[c_type, config.output_tile_shape[Int(0)], config.output_tile_shape[Int(1)], config.num_output_stages], c_tma_op: TMATensorTile[c_type, Int(3), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(3), Layout[*(), *()]]()], stage: OutputStage[Self.opc], work_tile_coord: Tuple[UInt32, UInt32, UInt32], M: UInt32, N: UInt32, alpha: Float32)

Execute epilogue to store accumulated results to global memory.

Uses TileWriter which encapsulates:

  • TmemArrayType.load_fragments() for TMEM load
  • AccumBarrier.arrive() for barrier signaling
  • TMEMToSMemWriter.write_fragments() for SMEM write
  • 3D TMA store (M, N, Batch coordinates)
  • tma_wait_pipelined() for TMA wait

Barrier synchronization (wait/step) is handled by caller via consumer() context.

Args:

validate_config

static def validate_config()

Validate configuration constraints at compile time.

init_barriers

static def init_barriers(ctx: KernelContext[config.num_clc_pipeline_stages, config.cta_group, config.cluster_shape[Int(0)], config.cluster_shape[Int(1)]], a_tma_op: TMATensorTile[a_type, Int(3), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(3), Layout[*(), *()]]()], b_tma_op: TMATensorTile[b_type, Int(3), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(3), Layout[*(), *()]]()], c_tma_op: TMATensorTile[c_type, Int(3), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(3), Layout[*(), *()]]()], sfa_tma_op: TMATensorTile[sfa_dtype, Int(5), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(5), Layout[*(), *()]]()], sfb_tma_op: TMATensorTile[sfb_dtype, Int(5), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(5), Layout[*(), *()]]()], input_barriers: SMemArray[SharedMemBarrier, ((config // config) * Int(2))], accum_barriers: SMemArray[SharedMemBarrier, (config * Int(2))], clc_throttle: SMemArray[SharedMemBarrier, (config * Int(2))], clc_full: SMemArray[SharedMemBarrier, config.num_clc_pipeline_stages], clc_empty: SMemArray[SharedMemBarrier, config.num_clc_pipeline_stages], tmem_dealloc: SMemArray[SharedMemBarrier, Int(1)])

Initialize barriers and prefetch TMA descriptors.

Args:

run

static def run(a_tma_op: TMATensorTile[a_type, Int(3), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(3), Layout[*(), *()]]()], b_tma_op: TMATensorTile[b_type, Int(3), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(3), Layout[*(), *()]]()], c_tma_op: TMATensorTile[c_type, Int(3), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(3), Layout[*(), *()]]()], sfa_tma_op: TMATensorTile[sfa_dtype, Int(5), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(5), Layout[*(), *()]]()], sfb_tma_op: TMATensorTile[sfb_dtype, Int(5), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(5), Layout[*(), *()]]()], alpha: Float32, cluster_dim: StaticTuple[Int32, Int(3)], mnk: StaticTuple[UInt32, Int(3)], workspace: Span[UInt64, MutAnyOrigin])

Kernel entry point - ported from legacy kernel.

Args: