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
TileWriter
struct TileWriter[tma_origin: ImmOrigin, c_type: DType, c_rank: Int, c_tile_shape: IndexList[c_rank], c_desc_shape: IndexList[c_rank], //, a_type: DType, accum_type: DType, block_tile_shape: IndexList[Int(3)], mma_shape: IndexList[Int(3)], opc: OutputPipelineConfig, c_swizzle: TensorMapSwizzle, transpose_c: Bool, c_smem_dim0: Int, c_smem_dim1: Int, num_output_stages: Int, num_output_warps: Int, elementwise_lambda_fn: Optional[def[dtype: DType, width: SIMDLength, *, alignment: Int = Int(1)](IndexList[Int(2)], SIMD[dtype, width]) capturing thin -> None] = None, 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, register_based_epilogue: Bool = True, batched: Bool = False, problem_n: Int = Int(0), num_peers: Int = Int(1)]
Output tile writer for SM100 matmul epilogue.
Stores pointer to TMA descriptor. SMEM tiles passed per-call.
Parameters are passed explicitly to work with both MatmulConfig and BlockScaledMatmulConfig.
The opc (OutputPipelineConfig) parameter must match the config used when constructing the OutputTilePipeline that provides OutputStage instances to the write() method.
Parameters
- tma_origin (
ImmOrigin): Memory origin of the TMA descriptor pointer (inferred). - c_type (
DType): Element dtype of the C output tensor (inferred). - c_rank (
Int): Rank of the C output tensor (inferred). - c_tile_shape (
IndexList[c_rank]): Per-tile shape of the C output (inferred). - c_desc_shape (
IndexList[c_rank]): TMA descriptor shape for C (inferred). - a_type (
DType): Element dtype of the A input matrix. - accum_type (
DType): Accumulator dtype stored in TMEM. - block_tile_shape (
IndexList[Int(3)]): Block tile shape as (BM, BN, BK). - mma_shape (
IndexList[Int(3)]): MMA instruction shape as (MMA_M, MMA_N, MMA_K). - opc (
OutputPipelineConfig): Output pipeline config bundling accumulator stages, stage stride, and CTA group. - c_swizzle (
TensorMapSwizzle): TMA swizzle pattern for the C SMEM layout. - transpose_c (
Bool): Whether C is stored transposed. - c_smem_dim0 (
Int): Row dimension of the C SMEM tile. - c_smem_dim1 (
Int): Column dimension of the C SMEM tile. - num_output_stages (
Int): Number of C SMEM pipeline stages. - num_output_warps (
Int): Number of warps driving the output pipeline. - elementwise_lambda_fn (
Optional[def[dtype: DType, width: SIMDLength, *, alignment: Int = Int(1)](IndexList[Int(2)], SIMD[dtype, width]) capturing thin -> None]): Optional elementwise epilogue applied to fragments before the store (defaults to None). - 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]]): Optional compute epilogue fused into the register path (defaults to None). - register_based_epilogue (
Bool): Whether the compute epilogue runs in registers (true) or SMEM (false) (defaults to True). - batched (
Bool): Whether the output uses 3D batched coordinates with a batch index (defaults to False). - problem_n (
Int): Logical N dimension used for row-major bounds checking in the slow path; 0 disables the N check (defaults to 0). - num_peers (
Int): Number of TMA store descriptors in the array; 1 for a local epilogue (defaults to 1).
Fields
- c_tma_op (
TileWriter[a_type, accum_type, block_tile_shape, mma_shape, opc, c_swizzle, transpose_c, c_smem_dim0, c_smem_dim1, num_output_stages, num_output_warps, elementwise_lambda_fn, elementwise_compute_lambda_fn, register_based_epilogue, batched, problem_n, num_peers].TmaOpPtr):
Implemented traits
AnyType,
Copyable,
Deinitable,
ImplicitlyCopyable,
Movable,
RegisterPassable,
TrivialRegisterPassable
comptime members
accum_tile_layout
comptime accum_tile_layout = Layout.row_major(block_tile_shape[Int(0)], TileWriter[a_type, accum_type, block_tile_shape, mma_shape, opc, c_swizzle, transpose_c, c_smem_dim0, c_smem_dim1, num_output_stages, num_output_warps, elementwise_lambda_fn, elementwise_compute_lambda_fn, register_based_epilogue, batched, problem_n, num_peers].stageN)
AccumTmemArray
comptime AccumTmemArray = TmemArrayType[accum_type, TileWriter[a_type, accum_type, block_tile_shape, mma_shape, opc, c_swizzle, transpose_c, c_smem_dim0, c_smem_dim1, num_output_stages, num_output_warps, elementwise_lambda_fn, elementwise_compute_lambda_fn, register_based_epilogue, batched, problem_n, num_peers].accum_tile_layout, EpilogueConfig.create(MMA_M=mma_shape[Int(0)], MMA_N=mma_shape[Int(1)], stageN=c_smem_dim0 if transpose_c else c_smem_dim1, cta_group=opc.cta_group, transpose_c=transpose_c, BM=block_tile_shape[Int(0)], BN=block_tile_shape[Int(1)]).num_stages, cta_group=opc.cta_group]
bits
comptime bits = 256
BM
comptime BM = block_tile_shape[Int(0)]
BN
comptime BN = block_tile_shape[Int(1)]
c_smem_layout
comptime c_smem_layout = Layout.row_major(c_smem_dim0, c_smem_dim1)
cta_group
comptime cta_group = opc.cta_group
CTileArray
comptime CTileArray = SMemTileArray2DRowMajor[c_type, c_smem_dim0, c_smem_dim1, num_output_stages]
data_paths
comptime data_paths = 16
epc
comptime epc = EpilogueConfig.create(MMA_M=mma_shape[Int(0)], MMA_N=mma_shape[Int(1)], stageN=TileWriter[a_type, accum_type, block_tile_shape, mma_shape, opc, c_swizzle, transpose_c, c_smem_dim0, c_smem_dim1, num_output_stages, num_output_warps, elementwise_lambda_fn, elementwise_compute_lambda_fn, register_based_epilogue, batched, problem_n, num_peers].stageN, cta_group=opc.cta_group, transpose_c=transpose_c, BM=block_tile_shape[Int(0)], BN=block_tile_shape[Int(1)])
epilogue_dtype
comptime epilogue_dtype = TileWriter.get_epilogue_dtype()
fragment_size
comptime fragment_size = (Int(128) // _resolve_warp_size())
is_lower_frag_required
comptime is_lower_frag_required = TileWriter[a_type, accum_type, block_tile_shape, mma_shape, opc, c_swizzle, transpose_c, c_smem_dim0, c_smem_dim1, num_output_stages, num_output_warps, elementwise_lambda_fn, elementwise_compute_lambda_fn, register_based_epilogue, batched, problem_n, num_peers].epc.is_lower_frag_required
MMA_M
comptime MMA_M = mma_shape[Int(0)]
MMA_N
comptime MMA_N = mma_shape[Int(1)]
N_dim
comptime N_dim = Int(0) if transpose_c else Int(1)
needs_sync
comptime needs_sync = False
num_accum_pipeline_stages
comptime num_accum_pipeline_stages = opc.num_stages
num_stages
comptime num_stages = TileWriter[a_type, accum_type, block_tile_shape, mma_shape, opc, c_swizzle, transpose_c, c_smem_dim0, c_smem_dim1, num_output_stages, num_output_warps, elementwise_lambda_fn, elementwise_compute_lambda_fn, register_based_epilogue, batched, problem_n, num_peers].epc.num_stages
rep
comptime rep = (TileWriter[a_type, accum_type, block_tile_shape, mma_shape, opc, c_swizzle, transpose_c, c_smem_dim0, c_smem_dim1, num_output_stages, num_output_warps, elementwise_lambda_fn, elementwise_compute_lambda_fn, register_based_epilogue, batched, problem_n, num_peers].stageN // Int(8))
rep_frag_size
comptime rep_frag_size = ((Int(128) // _resolve_warp_size()) * TileWriter[a_type, accum_type, block_tile_shape, mma_shape, opc, c_swizzle, transpose_c, c_smem_dim0, c_smem_dim1, num_output_stages, num_output_warps, elementwise_lambda_fn, elementwise_compute_lambda_fn, register_based_epilogue, batched, problem_n, num_peers].rep)
Stage
comptime Stage = OutputStage[opc]
stage_contiguous_size
comptime stage_contiguous_size = c_smem_dim1
stage_stride_cols
comptime stage_stride_cols = opc.stage_stride_cols
stageN
comptime stageN = c_smem_dim0 if transpose_c else c_smem_dim1
TmaOp
comptime TmaOp = TMATensorTile[c_type, c_rank, c_tile_shape, c_desc_shape]
TmaOpArray
comptime TmaOpArray = Array[TMATensorTile[c_type, c_rank, c_tile_shape, c_desc_shape], num_peers]
TmaOpArrayPtr
comptime TmaOpArrayPtr = Pointer[Array[TMATensorTile[c_type, c_rank, c_tile_shape, c_desc_shape], num_peers], tma_origin]
TmaOpPtr
comptime TmaOpPtr = Pointer[TMATensorTile[c_type, c_rank, c_tile_shape, c_desc_shape], tma_origin]
Methods
__init__
def __init__(c_tma_op: Pointer[TMATensorTile[c_type, c_rank, c_tile_shape, c_desc_shape], tma_origin]) -> Self
Initialize with pointer to TMA descriptor.
Args:
- c_tma_op (
Pointer[TMATensorTile[c_type, c_rank, c_tile_shape, c_desc_shape], tma_origin]): Pointer to the TMA store descriptor for C.
def __init__(c_tma_ops: Pointer[Array[TMATensorTile[c_type, c_rank, c_tile_shape, c_desc_shape], num_peers], tma_origin]) -> Self
Initialize from the c_tma_ops array pointer (TileWriterLike).
The standard local store targets a single descriptor, so this uses
element [0] of the array. Unifies construction with the
reduce-scatter writer, which retains all num_peers descriptors.
Args:
- c_tma_ops (
Pointer[Array[TMATensorTile[c_type, c_rank, c_tile_shape, c_desc_shape], num_peers], tma_origin]): Pointer to the array of TMA store descriptors for C; element[0]is used for the local store.
get_epilogue_dtype
write
def write(self, c_tiles: SMemTileArray2DRowMajor[c_type, c_smem_dim0, c_smem_dim1, num_output_stages], stage: OutputStage[opc], tile_coord: Tuple[UInt32, UInt32], shape: Tuple[UInt32, UInt32], elect_one_warp: Bool)
Write accumulated results to global memory (2D coords).
Args:
- c_tiles (
SMemTileArray2DRowMajor[c_type, c_smem_dim0, c_smem_dim1, num_output_stages]): SMEM tile array for the C output. - stage (
OutputStage[opc]): OutputStage with pipeline, index, and TMEM handle. - tile_coord (
Tuple[UInt32, UInt32]): (m_tile, n_tile) tile coordinates. - shape (
Tuple[UInt32, UInt32]): (M, N) problem dimensions. - elect_one_warp (
Bool): Whether this warp is elected for coordination.
write_batched
def write_batched(self, c_tiles: SMemTileArray2DRowMajor[c_type, c_smem_dim0, c_smem_dim1, num_output_stages], stage: OutputStage[opc], tile_coord: Tuple[UInt32, UInt32, UInt32], shape: Tuple[UInt32, UInt32], alpha: Float32 = 1)
Write accumulated results to global memory (3D batched coords).
Args:
- c_tiles (
SMemTileArray2DRowMajor[c_type, c_smem_dim0, c_smem_dim1, num_output_stages]): TileTensor-based SMEM tile array for C output. - stage (
OutputStage[opc]): OutputStage with pipeline, index, and TMEM handle. - tile_coord (
Tuple[UInt32, UInt32, UInt32]): (m_tile, n_tile, batch) coordinates. - shape (
Tuple[UInt32, UInt32]): (M, N) problem dimensions. - alpha (
Float32): Tensor scale factor (scalar).
write_splitk
def write_splitk[reduction_layout: TensorLayout](self, c_tiles: SMemTileArray2DRowMajor[c_type, c_smem_dim0, c_smem_dim1, num_output_stages], stage: OutputStage[opc], scheduler: TileScheduler, reduction_tensor: TileTensor[accum_type, reduction_layout, MutAnyOrigin], work_info: WorkInfo, shape: Tuple[UInt32, UInt32], elect_one_warp: Bool)
Write with split-K reduction. Only last split writes to GMEM.
write_absolute_with_bounds_check
def write_absolute_with_bounds_check[c_tensor_layout: TensorLayout](self, c_tiles: SMemTileArray2DRowMajor[c_type, c_smem_dim0, c_smem_dim1, num_output_stages], output_stage: OutputStage[opc], m_abs: UInt32, n_abs: UInt32, m_end: UInt32, expert_scale: Float32, c_tensor: TileTensor[c_type, c_tensor_layout, MutAnyOrigin])
Write with absolute coordinates and bounds checking.
For 1D-1D grouped kernels where M coordinate is absolute.
Parameters:
- c_tensor_layout (
TensorLayout): Layout of the C tensor in GMEM (inferred).
Args:
- c_tiles (
SMemTileArray2DRowMajor[c_type, c_smem_dim0, c_smem_dim1, num_output_stages]): SMEM tile array for the C output. - output_stage (
OutputStage[opc]): OutputStage with pipeline, index, and TMEM handle. - m_abs (
UInt32): Absolute M coordinate (start of tile in token space). - n_abs (
UInt32): Absolute N coordinate (start of tile). - m_end (
UInt32): End offset for bounds checking (exclusive). - expert_scale (
Float32): Per-expert output scaling factor. - c_tensor (
TileTensor[c_type, c_tensor_layout, MutAnyOrigin]): C tensor in GMEM for bounds-checked stores.
write_with_residual
def write_with_residual[pipeline_origin: MutOrigin, //, num_src_stages: Int](self, out_tiles: SMemTileArray2DRowMajor[c_type, c_smem_dim0, c_smem_dim1, num_output_stages], stage: OutputStage[opc], src_tile: SMemTileArray2DRowMajor[c_type, c_smem_dim0, c_smem_dim1, num_src_stages], src_pipeline: Pointer[ProducerConsumerPipeline[num_src_stages], pipeline_origin], beta: Scalar[c_type], tile_coord: Tuple[UInt32, UInt32], shape: Tuple[UInt32, UInt32], elect_one_warp: Bool)
Write with residual: D = lambda(accum) + beta * C.
Matches the CUTLASS sm100_epilogue_tma_warpspecialized lockstep
pattern: the epilogue load warp pre-fetches one source sub-tile per
inner epilogue stage into a num_src_stages-deep SMEM pipeline; this
method drives one wait_producer / use / consumer_release / step
cycle on src_pipeline per inner stage. The buffer index is read from
the pipeline's consumer_stage() rather than computed offline, so
producer and consumer stay synchronized exactly as in CUTLASS's
consumer_wait → copy(sC) → consumer_release per epi sub-tile.
Pipeline per inner stage:
- Load accum from TMEM to registers (epilogue dtype).
- Apply
elementwise_compute_lambda_fn(pre-residual fusion). - Wait for source[k] via
src_pipeline.consume(); computeD = accum + beta * Creading from the SMEM buffer at the pipeline's current stage index; release source[k] on context exit. - Apply
elementwise_lambda_fn(post-residual, owns the GMEM store) OR stage to output SMEM and TMA-store to GMEM.
Parameters:
- pipeline_origin (
MutOrigin): Mutability origin of the source pipeline ref. - num_src_stages (
Int): Number of source SMEM buffers; must equal the epi-load pipeline's stage count in the kernel.
Args:
- out_tiles (
SMemTileArray2DRowMajor[c_type, c_smem_dim0, c_smem_dim1, num_output_stages]): Output SMEM tile array (for D output). - stage (
OutputStage[opc]): OutputStage with pipeline, index, and TMEM handle. - src_tile (
SMemTileArray2DRowMajor[c_type, c_smem_dim0, c_smem_dim1, num_src_stages]): Source C SMEM tile array (num_src_stages buffers). - src_pipeline (
Pointer[ProducerConsumerPipeline[num_src_stages], pipeline_origin]): Pointer to the source producer/consumer pipeline. One acquire/release cycle is driven per inner epilogue stage. - beta (
Scalar[c_type]): Residual scale factor. - tile_coord (
Tuple[UInt32, UInt32]): (m_tile, n_tile) coordinates. - shape (
Tuple[UInt32, UInt32]): (M, N) problem dimensions. - elect_one_warp (
Bool): Whether this warp is elected for coordination.
write_batched_with_tma_epilogue_load
def write_batched_with_tma_epilogue_load[epi_load_swizzle: TensorMapSwizzle, epilogue_layout: TensorLayout](self, c_tiles: SMemTileArray2DRowMajor[c_type, c_smem_dim0, c_smem_dim1, num_output_stages], output_stage: OutputStage[opc], epilogue_tile: TileTensor[c_type, epilogue_layout, MutAnyOrigin, address_space=AddressSpace.SHARED], tile_coord: Tuple[UInt32, UInt32, UInt32], c_shape: Tuple[UInt32, UInt32])
Write accumulated results with epilogue tensor addition to global memory.
Pipeline: TMEM → Registers → (+epilogue from SMEM) → SMEM → GMEM (TMA).
write_batched_with_1d_bias
def write_batched_with_1d_bias[epilogue_layout: TensorLayout](self, c_tiles: SMemTileArray2DRowMajor[c_type, c_smem_dim0, c_smem_dim1, num_output_stages], output_stage: OutputStage[opc], epilogue_tile: TileTensor[c_type, epilogue_layout, MutAnyOrigin, address_space=AddressSpace.SHARED], tile_coord: Tuple[UInt32, UInt32, UInt32], c_shape: Tuple[UInt32, UInt32])
Write accumulated results with 1D bias addition to global memory.
Pipeline: TMEM -> Registers -> (+1D bias broadcast from SMEM) -> SMEM -> GMEM (TMA).
The bias SMEM tile is 1×MMA_N loaded via cp.async (linear layout, no swizzle) and then broadcast across all M rows.
write_batched_with_tma_epilogue_load_strips
def write_batched_with_tma_epilogue_load_strips[epi_load_swizzle: TensorMapSwizzle, num_epi_stages: Int](self, c_tiles: SMemTileArray2DRowMajor[c_type, c_smem_dim0, c_smem_dim1, num_output_stages], output_stage: OutputStage[opc], mut epilogue_pipeline: ProducerConsumerPipeline[num_epi_stages], epilogue_tiles_base: Pointer[Scalar[c_type], MutAnyOrigin, address_space=AddressSpace.SHARED], epilogue_tile_elems: Int, tile_coord: Tuple[UInt32, UInt32, UInt32], c_shape: Tuple[UInt32, UInt32])
Write accumulated results with BM×stageN pipelined epilogue addition.
For non-AB_swapped configs. Each epilogue pipeline stage is one BM×stageN tile. Producer sends tiles in stage-outer / col_wg-inner order; consumer mirrors that structure so each TMEM stage is fully processed (load → add epilogue → write) before advancing to the next.