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
BlackwellBlockwiseFP8MatmulKernel
struct BlackwellBlockwiseFP8MatmulKernel[a_type: DType, b_type: DType, c_type: DType, a_scales_type: DType, b_scales_type: DType, b_scales_layout: TensorLayout, transpose_b: Bool, config: MatmulConfig[a_type, b_type, c_type, transpose_b], cluster_shape: StaticTuple[Int32, Int(3)] = StaticTuple(Int32(1)), n_scale_granularity: Int = Int(128)]
Blockwise FP8 matmul kernel with register-based accumulation.
This kernel implements per-K-iteration scaling in CUDA cores:
- Load warp: TMA loads A, B, A-scales to SMEM
- MMA warp: Standard MMA (partial to TMEM)
- Epilogue warp: TMEM read → scale → register accumulate → output
Parameters
- a_type (
DType): Element type of the A matrix tiles. - b_type (
DType): Element type of the B matrix tiles. - c_type (
DType): Element type of the C output matrix tiles. - a_scales_type (
DType): A-scales element type (must equalb_scales_type). - b_scales_type (
DType): B-scales element type (must equala_scales_type). - b_scales_layout (
TensorLayout): Memory layout of the B-scales tensor. - transpose_b (
Bool): Whether B is stored transposed (must beTrue). - config (
MatmulConfig[a_type, b_type, c_type, transpose_b]): Matmul tile, MMA, pipeline, and cluster configuration. - cluster_shape (
StaticTuple[Int32, Int(3)]): CTA cluster shape(x, y, z)for LLVM metadata (defaults to(1, 1, 1)). - n_scale_granularity (
Int): B-scales N-direction block size in elements (defaults to 128).
Implemented traits
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_scales_expected_bytes
comptime a_scales_expected_bytes = (config.block_tile_shape[Int(0)] * size_of[a_scales_type]())
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]()[Int(0)]
a_tma_load_size
comptime a_tma_load_size = (BlackwellBlockwiseFP8MatmulKernel[a_type, b_type, c_type, a_scales_type, b_scales_type, b_scales_layout, transpose_b, config, cluster_shape, n_scale_granularity].a_tile_dim0 * (config.a_swizzle.bytes() // size_of[a_type]()))
a_tma_rows
comptime a_tma_rows = BlackwellBlockwiseFP8MatmulKernel[a_type, b_type, c_type, a_scales_type, b_scales_type, b_scales_layout, transpose_b, config, cluster_shape, n_scale_granularity].a_tile_dim0
accum_dims
comptime accum_dims = get_accumulator_dims[c_smem_dim1=config.output_tile_shape[Int(1)], block_tile_shape=config.block_tile_shape, mma_shape=config.mma_shape, cta_group=config.cta_group]()
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
AccumTensor
comptime AccumTensor = TmemTensor[DType.float32, BlackwellBlockwiseFP8MatmulKernel[a_type, b_type, c_type, a_scales_type, b_scales_type, b_scales_layout, transpose_b, config, cluster_shape, n_scale_granularity].tmem_accum_layout, cta_group=config.cta_group]
Accumulator
comptime Accumulator = BlockwiseFP8Accumulator[DType.float32, get_accumulator_dims[c_smem_dim1=config.output_tile_shape[Int(1)], block_tile_shape=config.block_tile_shape, mma_shape=config.mma_shape, cta_group=config.cta_group]()[Int(0)], get_accumulator_dims[c_smem_dim1=config.output_tile_shape[Int(1)], block_tile_shape=config.block_tile_shape, mma_shape=config.mma_shape, cta_group=config.cta_group]()[Int(1)], BlackwellBlockwiseFP8MatmulKernel[a_type, b_type, c_type, a_scales_type, b_scales_type, b_scales_layout, transpose_b, config, cluster_shape, n_scale_granularity].is_lower_required, config.block_tile_shape, config.mma_shape, Int((mul config.cluster_shape[Int(0)], config.cluster_shape[Int(1)])), n_scale_granularity]
ADescLayout
comptime ADescLayout = Layout[*(), *()]
AScalesLayout
comptime AScalesLayout = Layout[*(), *()]
AScalesTmaOp
comptime AScalesTmaOp = TMATensorTile[a_scales_type, Int(2), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(2), Layout[*(), *()]]()]
ATileLayout
comptime ATileLayout = Layout[*(), *()]
ATmaOp
comptime ATmaOp = TMATensorTile[a_type, Int(2), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(2), 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_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]()[Int(1)]
b_tma_load_size
comptime b_tma_load_size = (BlackwellBlockwiseFP8MatmulKernel[a_type, b_type, c_type, a_scales_type, b_scales_type, b_scales_layout, transpose_b, config, cluster_shape, n_scale_granularity].b_tile_dim0 * (config.b_swizzle.bytes() // size_of[b_type]()))
b_tma_rows
comptime b_tma_rows = BlackwellBlockwiseFP8MatmulKernel[a_type, b_type, c_type, a_scales_type, b_scales_type, b_scales_layout, transpose_b, config, cluster_shape, n_scale_granularity].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)]
BScalesTile
comptime BScalesTile = TileTensor[b_scales_type, b_scales_layout, ImmutAnyOrigin]
BTileLayout
comptime BTileLayout = Layout[*(), *()]
BTmaOp
comptime BTmaOp = TMATensorTile[b_type, Int(2), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(2), Layout[*(), *()]]()]
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]()[Int(2)]
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(2), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(2), Layout[*(), *()]]()]
EPILOGUE_THREADS
comptime EPILOGUE_THREADS = (Int(4) * _resolve_warp_size())
EpilogueCtx
comptime EpilogueCtx = EpilogueWarpContext[BlackwellBlockwiseFP8MatmulKernel[a_type, b_type, c_type, a_scales_type, b_scales_type, b_scales_layout, transpose_b, config, cluster_shape, n_scale_granularity].opc, _resolve_warp_size(), Int((mul _resolve_warp_size(), 4))]
EpilogueHandle
comptime EpilogueHandle = EpilogueWarp[BlackwellBlockwiseFP8MatmulKernel[a_type, b_type, c_type, a_scales_type, b_scales_type, b_scales_layout, transpose_b, config, cluster_shape, n_scale_granularity].opc, _resolve_warp_size(), Int((mul _resolve_warp_size(), 4))]
input_expected_bytes
comptime input_expected_bytes = (config * Int((add (mul config.block_tile_shape[Int(0)], config.block_tile_shape[Int(2)], size_of[a_type]()), (mul config.block_tile_shape[Int(1)], config.block_tile_shape[Int(2)], size_of[b_type]()), (mul config.block_tile_shape[Int(0)], size_of[a_scales_type]()))))
InputTilePipeline
comptime InputTilePipeline = InputTilePipeline[BlockwiseFP8TilePayload[a_type, b_type, a_scales_type, 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(Int(1), config.block_tile_shape[Int(0)], __list_literal__=NoneType(None)), config.num_pipeline_stages], (config // config), config.k_group_size]
is_lower_required
comptime is_lower_required = is_lower_fragment_required[config.cta_group, config.block_tile_shape]()
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[BlackwellBlockwiseFP8MatmulKernel[a_type, b_type, c_type, a_scales_type, b_scales_type, b_scales_layout, transpose_b, config, cluster_shape, n_scale_granularity].opc, _resolve_warp_size(), Int((mul _resolve_warp_size(), 4))]
MmaHandle
comptime MmaHandle = MmaWarp[BlackwellBlockwiseFP8MatmulKernel[a_type, b_type, c_type, a_scales_type, b_scales_type, b_scales_layout, transpose_b, config, cluster_shape, n_scale_granularity].opc, _resolve_warp_size(), Int((mul _resolve_warp_size(), 4))]
MmaOp
comptime MmaOp = MmaOpSM100_SS[c_type, a_type, b_type, 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(BlackwellBlockwiseFP8MatmulKernel[a_type, b_type, c_type, a_scales_type, b_scales_type, b_scales_layout, transpose_b, config, cluster_shape, n_scale_granularity].num_accum_pipeline_stages, BlackwellBlockwiseFP8MatmulKernel[a_type, b_type, c_type, a_scales_type, b_scales_type, b_scales_layout, transpose_b, config, cluster_shape, n_scale_granularity].stage_stride_cols, BlackwellBlockwiseFP8MatmulKernel[a_type, b_type, c_type, a_scales_type, b_scales_type, b_scales_layout, transpose_b, config, cluster_shape, n_scale_granularity].cta_group)
OutputM
comptime OutputM = config.output_tile_shape[Int(0)]
OutputN
comptime OutputN = config.output_tile_shape[Int(1)]
OutputPipeline
comptime OutputPipeline = OutputTilePipeline[BlackwellBlockwiseFP8MatmulKernel[a_type, b_type, c_type, a_scales_type, b_scales_type, b_scales_layout, transpose_b, config, cluster_shape, n_scale_granularity].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
SmemType
comptime SmemType = BlockwiseFP8Smem[a_type, b_type, c_type, a_scales_type, transpose_b, config=config]
stage_stride_cols
comptime stage_stride_cols = (Int(512) // config)
TilePayload
comptime TilePayload = BlockwiseFP8TilePayload[a_type, b_type, a_scales_type, 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(Int(1), config.block_tile_shape[Int(0)], __list_literal__=NoneType(None)), config.num_pipeline_stages]
TileWriterType
comptime TileWriterType = BlockwiseFP8TileWriter[c_type, config.output_tile_shape[Int(0)], config.output_tile_shape[Int(1)], DType.float32, get_accumulator_dims[c_smem_dim1=config.output_tile_shape[Int(1)], block_tile_shape=config.block_tile_shape, mma_shape=config.mma_shape, cta_group=config.cta_group]()[Int(0)], get_accumulator_dims[c_smem_dim1=config.output_tile_shape[Int(1)], block_tile_shape=config.block_tile_shape, mma_shape=config.mma_shape, cta_group=config.cta_group]()[Int(1)], block_tile_shape=config.block_tile_shape, mma_shape=config.mma_shape, is_lower_frag_required=BlackwellBlockwiseFP8MatmulKernel[a_type, b_type, c_type, a_scales_type, b_scales_type, b_scales_layout, transpose_b, config, cluster_shape, n_scale_granularity].is_lower_required, cta_group=config.cta_group, num_output_stages=config.num_output_stages, num_output_warps=Int(4), c_swizzle=config.c_swizzle]
TMA_LOAD_THREADS
comptime TMA_LOAD_THREADS = WARP_SIZE
Tmem
comptime Tmem = TmemAllocation[OutputPipelineConfig(config.num_accum_pipeline_stages, (Int(512) // config), config.cta_group).cta_group]
tmem_accum_layout
comptime tmem_accum_layout = Layout.row_major(config.mma_shape[Int(0)], config.mma_shape[Int(1)])
TmemDealloc
comptime TmemDealloc = TmemDeallocBarrier[OutputPipelineConfig(config.num_accum_pipeline_stages, (Int(512) // config), config.cta_group).cta_group]
Methods
load_input_tiles
static def load_input_tiles[a_tma_origin: ImmOrigin, b_tma_origin: ImmOrigin, a_scales_tma_origin: ImmOrigin, tiles_origin: MutOrigin, //](a_loader: TileLoader[a_tma_origin, a_type, Layout[*(), *()], Layout[*(), *()], cta_group=config.cta_group], b_loader: TileLoader[b_tma_origin, b_type, Layout[*(), *()], Layout[*(), *()], cta_group=config.cta_group], a_scales_loader: ScalesLoader[a_scales_tma_origin, a_scales_type, Layout[*(), *()], Layout[*(), *()], cta_group=config.cta_group], tiles: InputProducerStage[tiles_origin, BlockwiseFP8TilePayload[a_type, b_type, a_scales_type, 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(Int(1), config.block_tile_shape[Int(0)], __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], iter_idx: Int, elect_one_cta: Bool)
Load A, B, and A-scales tiles using TMA.
Parameters:
- a_tma_origin (
ImmOrigin): Immutable origin of the A TMA descriptor (inferred). - b_tma_origin (
ImmOrigin): Immutable origin of the B TMA descriptor (inferred). - a_scales_tma_origin (
ImmOrigin): Immutable origin of A-scales TMA descriptor (inferred). - tiles_origin (
MutOrigin): Mutable origin of the producer tiles (inferred).
Args:
- a_loader (
TileLoader[a_tma_origin, a_type, Layout[*(), *()], Layout[*(), *()], cta_group=config.cta_group]): TileLoader for A matrix. - b_loader (
TileLoader[b_tma_origin, b_type, Layout[*(), *()], Layout[*(), *()], cta_group=config.cta_group]): TileLoader for B matrix. - a_scales_loader (
ScalesLoader[a_scales_tma_origin, a_scales_type, Layout[*(), *()], Layout[*(), *()], cta_group=config.cta_group]): ScalesLoader for A-scales. - tiles (
InputProducerStage[tiles_origin, BlockwiseFP8TilePayload[a_type, b_type, a_scales_type, 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(Int(1), config.block_tile_shape[Int(0)], __list_literal__=NoneType(None)), config.num_pipeline_stages], (config // config), config.k_group_size]): InputProducerStage context with encapsulated tile access. - peer_cta_coord (
Tuple[Int, Int, Int]): Peer CTA coordinates for multicast. - work_tile_coord (
Tuple[Int, Int]): Current work tile M/N coordinates. - iter_idx (
Int): K iteration index. - elect_one_cta (
Bool): Whether this is the elected CTA in the cluster.
mma
static def mma[tiles_origin: MutOrigin, //](tiles: InputConsumerStage[tiles_origin, BlockwiseFP8TilePayload[a_type, b_type, a_scales_type, 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(Int(1), config.block_tile_shape[Int(0)], __list_literal__=NoneType(None)), config.num_pipeline_stages], (config // config), config.k_group_size], mma_op: MmaOpSM100_SS[c_type, a_type, b_type, 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], accum_tensor: TmemTensor[DType.float32, Self.tmem_accum_layout, cta_group=config.cta_group])
Execute standard MMA operations (partial results to TMEM).
For blockwise FP8, each K iteration writes a fresh partial to TMEM. The epilogue accumulates across K in registers, not TMEM. Therefore init_c is always True (unlike standard matmul).
Parameters:
- tiles_origin (
MutOrigin): Mutable origin of the consumer tiles (inferred).
Args:
- tiles (
InputConsumerStage[tiles_origin, BlockwiseFP8TilePayload[a_type, b_type, a_scales_type, 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(Int(1), config.block_tile_shape[Int(0)], __list_literal__=NoneType(None)), config.num_pipeline_stages], (config // config), config.k_group_size]): Input consumer stage with A, B, A-scales tiles. - mma_op (
MmaOpSM100_SS[c_type, a_type, b_type, 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]): The MMA operator. - accum_tensor (
TmemTensor[DType.float32, Self.tmem_accum_layout, cta_group=config.cta_group]): Typed TMEM tensor view for the accumulator stage.
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(2), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(2), Layout[*(), *()]]()], b_tma_op: TMATensorTile[b_type, Int(2), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(2), Layout[*(), *()]]()], c_tma_op: TMATensorTile[c_type, Int(2), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(2), Layout[*(), *()]]()], a_scales_tma_op: TMATensorTile[a_scales_type, Int(2), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(2), 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:
- ctx (
KernelContext[config.num_clc_pipeline_stages, config.cta_group, config.cluster_shape[Int(0)], config.cluster_shape[Int(1)]]): Kernel context with warp and CTA role and multicast masks. - a_tma_op (
TMATensorTile[a_type, Int(2), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(2), Layout[*(), *()]]()]): TMA descriptor op for A matrix tiles. - b_tma_op (
TMATensorTile[b_type, Int(2), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(2), Layout[*(), *()]]()]): TMA descriptor op for B matrix tiles. - c_tma_op (
TMATensorTile[c_type, Int(2), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(2), Layout[*(), *()]]()]): TMA descriptor op for C output tiles. - a_scales_tma_op (
TMATensorTile[a_scales_type, Int(2), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(2), Layout[*(), *()]]()]): TMA descriptor op for A-scales tiles. - input_barriers (
SMemArray[SharedMemBarrier, ((config // config) * Int(2))]): Input pipeline barriers for producer and consumer. - accum_barriers (
SMemArray[SharedMemBarrier, (config * Int(2))]): Accumulator barriers for MMA to epilogue sync. - clc_throttle (
SMemArray[SharedMemBarrier, (config * Int(2))]): CLC throttle barriers for scheduler backpressure. - clc_full (
SMemArray[SharedMemBarrier, config.num_clc_pipeline_stages]): CLC full barriers signalling tile availability. - clc_empty (
SMemArray[SharedMemBarrier, config.num_clc_pipeline_stages]): CLC empty barriers signalling tile consumption. - tmem_dealloc (
SMemArray[SharedMemBarrier, Int(1)]): TMEM deallocation barrier for accumulator slot reuse.
run
static def run(a_tma_op: TMATensorTile[a_type, Int(2), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(2), Layout[*(), *()]]()], b_tma_op: TMATensorTile[b_type, Int(2), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(2), Layout[*(), *()]]()], c_tma_op: TMATensorTile[c_type, Int(2), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(2), Layout[*(), *()]]()], a_scales_tma_op: TMATensorTile[a_scales_type, Int(2), _to_index_list[Layout[*(), *()]](), _to_index_list[Int(2), Layout[*(), *()]]()], cluster_dim: StaticTuple[Int32, Int(3)], num_iters: Int32, b_scales: TileTensor[b_scales_type, b_scales_layout, ImmutAnyOrigin], problem_shape: StaticTuple[Int32, Int(3)])
Kernel entry point for blockwise FP8 matmul.