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

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), c_store_dead: Bool = False, p5_direct_scatter: Bool = False, p5_row_cache: Bool = False, SinkT: EpiloguePeerSink = NullPeerSink]

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).
  • ​c_store_dead (Bool): Whether nothing reads the local C output, so the store is dead for the whole launch. The caller then supplies an empty C descriptor and an unbacked C tensor, so neither may be touched here (defaults to False).
  • ​p5_direct_scatter (Bool): Whether to compile in the in-epilogue EP-combine peer scatter-send. Still gated at runtime by p3_control (defaults to False).
  • ​p5_row_cache (Bool): Whether the peer send resolves each row's destination once per launch through P3PeerSendConfig.row_base_ptr instead of once per (row, tile) (defaults to False).
  • ​SinkT (EpiloguePeerSink): A second destination for the output tile, replacing the local store when its Enabled is set. Stateless, so it costs nothing to carry; defaults to NullPeerSink, which keeps the local store.

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, c_store_dead, p5_direct_scatter, p5_row_cache, SinkT].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, c_store_dead, p5_direct_scatter, p5_row_cache, SinkT].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, c_store_dead, p5_direct_scatter, p5_row_cache, SinkT].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, c_store_dead, p5_direct_scatter, p5_row_cache, SinkT].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, c_store_dead, p5_direct_scatter, p5_row_cache, SinkT].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, c_store_dead, p5_direct_scatter, p5_row_cache, SinkT].epc.num_stages

P5_SEND_DRAIN_BARRIER_ID​

comptime P5_SEND_DRAIN_BARRIER_ID = 6

p5_send_enabled​

comptime p5_send_enabled = p5_direct_scatter if (get_defined_int[StringSpan("P5_SEND_WARPS"), Int(0)]() > Int(0)) else (get_defined_int[StringSpan("P5_SEND_WARPS"), Int(0)]() > Int(0))

P5_SEND_RELEASE_BARRIER_ID​

comptime P5_SEND_RELEASE_BARRIER_ID = 5

P5_SEND_THREADS​

comptime P5_SEND_THREADS = (get_defined_int[StringSpan("P5_SEND_WARPS"), Int(0)]() * _resolve_warp_size())

p5_send_warps​

comptime p5_send_warps = get_defined_int[StringSpan("P5_SEND_WARPS"), Int(0)]()

p5_stamp_window​

comptime p5_stamp_window = get_defined_bool[StringSpan("P5_STAMP_WINDOW")]()

p5_trace_epi​

comptime p5_trace_epi = get_defined_bool[StringSpan("P5_TRACE_EPI")]()

P5SendDrain​

comptime P5SendDrain = WarpGroupBarrier[(Int((mul _resolve_warp_size(), num_output_warps)) + Int((mul get_defined_int[StringSpan("P5_SEND_WARPS"), Int(0)](), _resolve_warp_size()))), Int(6)]

P5SendRelease​

comptime P5SendRelease = WarpGroupBarrier[(Int((mul _resolve_warp_size(), num_output_warps)) + Int((mul get_defined_int[StringSpan("P5_SEND_WARPS"), Int(0)](), _resolve_warp_size()))), Int(5)]

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, c_store_dead, p5_direct_scatter, p5_row_cache, SinkT].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, c_store_dead, p5_direct_scatter, p5_row_cache, SinkT].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]

WarpRole​

comptime WarpRole = WarpRole1D1D[num_epi_warps=num_output_warps]

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:

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:

get_epilogue_dtype​

static def get_epilogue_dtype() -> DType

Returns:

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:

def write[ComputeFnType: ElementwiseComputeFn](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, compute_fn: ComputeFnType)

Write accumulated results to global memory, applying compute_fn.

Takes the compute epilogue as a runtime closure instead of the elementwise_compute_lambda_fn parameter, which must be unset. Only the register-based epilogue is supported.

Parameters:

Args:

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:

def write_batched[ComputeFnType: ElementwiseComputeFn](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], compute_fn: ComputeFnType, alpha: Float32 = 1)

Write accumulated results to global memory (3D batched coords), applying compute_fn.

Takes the compute epilogue as a runtime closure instead of the elementwise_compute_lambda_fn parameter, which must be unset.

Parameters:

Args:

write_splitk​

def write_splitk[reduction_layout: TensorLayout, reduction_engine: TensorEngine](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, Engine=reduction_engine], 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, c_tensor_engine: TensorEngine, RowScalesT: RowScales = NullRowScales](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, Engine=c_tensor_engine, address_space=c_tensor.address_space, linear_idx_type=c_tensor.linear_idx_type], p3_control: Int = Int(-1), p3_expert_id: Int32 = Int32(0), p3_cfg: P3PeerSendConfig = P3PeerSendConfig.disabled(), row_scales: RowScalesT = NullRowScales())

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).
  • ​c_tensor_engine (TensorEngine): Engine of the C tensor in GMEM (inferred).
  • ​RowScalesT (RowScales): Per-row output scales type. Defaults to the no-op NullRowScales.

Args:

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:

  1. Load accum from TMEM to registers (epilogue dtype).
  2. Apply elementwise_compute_lambda_fn (pre-residual fusion).
  3. Wait for source[k] via src_pipeline.consume(); compute D = accum + beta * C reading from the SMEM buffer at the pipeline's current stage index; release source[k] on context exit.
  4. 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:

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.

Was this page helpful?