For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /max/get-started.md).
Mojo struct
EpilogueApplier
struct EpilogueApplier[MMA_M: Int, stageN: Int, num_stages: Int, repeats: Int, cta_group: Int, transpose_c: Bool]
Apply element-wise epilogue lambda to register fragments.
Fieldsβ
- βcoords (
EpilogueApplier[MMA_M, stageN, num_stages, repeats, cta_group, transpose_c].Coords): - βwarp_id (
UInt32): - βlane_id (
UInt32): - βM (
UInt32): - βN (
UInt32):
Implemented traitsβ
AnyType,
Copyable,
ImplicitlyCopyable,
ImplicitlyDeletable,
Movable,
RegisterPassable,
TrivialRegisterPassable
comptime membersβ
Coordsβ
comptime Coords = FragmentCoords[stageN, repeats]
Methodsβ
__init__β
def __init__(warp_id: UInt32, lane_id: UInt32, c_shape: Tuple[UInt32, UInt32]) -> Self
compute_staged_coordsβ
def compute_staged_coords(self, stage: UInt32, c_row: UInt32, c_col: UInt32) -> Tuple[UInt32, UInt32]
Compute global coords with warp and stage offsets (layout-dependent).
Returns:
apply_to_fragmentβ
def apply_to_fragment[epilogue_dtype: DType, frag_size: Int, compute_lambda_fn: def[dtype: DType, width: SIMDLength, *, alignment: Int = Int(1)](IndexList[Int(2)], SIMD[dtype, width]) capturing thin -> SIMD[dtype, width], is_in_bounds: Bool = False](self, mut frag: Array[Scalar[epilogue_dtype], frag_size], staged_row: UInt32, staged_col: UInt32, is_upper: Bool)
Apply epilogue lambda to fragment elements with global coords.
is_in_bounds=True: caller asserts the whole tile fits in
(self.M, self.N); per-position checks are elided. Default
False keeps them β TMA masks the OOB write but the lambda
may dereference operands at the supplied (row, col).
transpose_c=True: per-position β top.row/bot.row in a (top, bot)
pair are 8 apart (upper +0/+8 or lower +16/+24) and can straddle
self.N. transpose_c=False: early return on top_col is safe
(top.col == bot.col and self.N is alignment-bound).
apply_to_both_fragmentsβ
def apply_to_both_fragments[epilogue_dtype: DType, frag_size: Int, compute_lambda_fn: def[dtype: DType, width: SIMDLength, *, alignment: Int = Int(1)](IndexList[Int(2)], SIMD[dtype, width]) capturing thin -> SIMD[dtype, width], is_lower_frag_required: Bool, is_in_bounds: Bool = False](self, mut upper_frag: Array[Scalar[epilogue_dtype], frag_size], mut lower_frag: Array[Scalar[epilogue_dtype], frag_size], stage: UInt32, c_row: UInt32, c_col: UInt32) -> Tuple[Array[Scalar[epilogue_dtype], frag_size], Array[Scalar[epilogue_dtype], frag_size]]
Apply epilogue to both fragments (main entry point).
Returns:
Tuple[Array[Scalar[epilogue_dtype], frag_size], Array[Scalar[epilogue_dtype], frag_size]]
apply_elementwise_epilogue_to_fragmentβ
def apply_elementwise_epilogue_to_fragment[epilogue_dtype: DType, frag_size: Int, elementwise_lambda_fn: def[dtype: DType, width: SIMDLength, *, alignment: Int = Int(1)](IndexList[Int(2)], SIMD[dtype, width]) capturing thin -> None, is_in_bounds: Bool = False](self, frag: SIMD[epilogue_dtype, frag_size], staged_row: UInt32, staged_col: UInt32, is_upper: Bool)
Apply elementwise epilogue lambda to fragment elements with global coords.
Unlike apply_to_fragment which uses a compute lambda that returns modified values, this calls an elementwise epilogue (returns None) that stores directly to global memory.
is_in_bounds=True: caller asserts the whole tile fits in
(self.M, self.N); the per-position row/column checks are elided and
the lambda is called for every element. Default False keeps the
checks (needed for border tiles, where an unchecked direct store would
write out of bounds). Mirrors apply_to_fragment's is_in_bounds.
apply_elementwise_epilogue_to_both_fragmentsβ
def apply_elementwise_epilogue_to_both_fragments[epilogue_dtype: DType, frag_size: Int, elementwise_lambda_fn: def[dtype: DType, width: SIMDLength, *, alignment: Int = Int(1)](IndexList[Int(2)], SIMD[dtype, width]) capturing thin -> None, is_lower_frag_required: Bool, is_in_bounds: Bool = False](self, upper_frag: SIMD[epilogue_dtype, frag_size], lower_frag: SIMD[epilogue_dtype, frag_size], stage: UInt32, c_row: UInt32, c_col: UInt32)
Apply elementwise epilogue to both fragments.
Similar to apply_to_both_fragments but uses elementwise_epilogue_type which writes directly to global memory and returns None.
is_in_bounds is threaded to the per-fragment path: when True the
caller has asserted the whole tile fits in (M, N) and the
per-position bounds checks are elided.
add_residual_to_fragmentβ
def add_residual_to_fragment[epilogue_dtype: DType, frag_size: Int, c_type: DType, c_smem_stride: Int, swizzle: Swizzle](self, mut frag: Array[Scalar[epilogue_dtype], frag_size], local_row: UInt32, local_col: UInt32, is_upper: Bool, src_ptr: Pointer[Scalar[c_type], address_space=AddressSpace.SHARED, _safe=False], beta: Scalar[epilogue_dtype])
Add beta * C to fragment elements by loading C from swizzled SMEM.
Uses the same per-lane coordinate mapping as apply_to_fragment, but instead of applying a lambda, loads source C values from SMEM at the matching swizzled addresses and adds beta * C to each element.
Args:
- βfrag (
Array[Scalar[epilogue_dtype], frag_size]): Fragment register values to modify in-place. - βlocal_row (
UInt32): Tile-local row offset (warp offset within tile). - βlocal_col (
UInt32): Tile-local column offset (stage offset within tile). - βis_upper (
Bool): Whether this is the upper (rows 0-15) or lower (16-31) fragment half. - βsrc_ptr (
Pointer[Scalar[c_type], address_space=AddressSpace.SHARED, _safe=False]): Pointer to source C SMEM tile (same TMA swizzle as output). - βbeta (
Scalar[epilogue_dtype]): Residual scale factor.
add_residual_to_both_fragmentsβ
def add_residual_to_both_fragments[epilogue_dtype: DType, frag_size: Int, is_lower_frag_required: Bool, c_type: DType, c_smem_stride: Int, swizzle: Swizzle](self, mut upper_frag: Array[Scalar[epilogue_dtype], frag_size], mut lower_frag: Array[Scalar[epilogue_dtype], frag_size], stage: UInt32, src_ptr: Pointer[Scalar[c_type], address_space=AddressSpace.SHARED, _safe=False], beta: Scalar[epilogue_dtype]) -> Tuple[Array[Scalar[epilogue_dtype], frag_size], Array[Scalar[epilogue_dtype], frag_size]]
Add beta * C to both fragment halves from swizzled SMEM.
Computes tile-local coordinates from stage and warp ID, then loads source C from SMEM and adds beta * C to each fragment element.
Args:
- βupper_frag (
Array[Scalar[epilogue_dtype], frag_size]): Upper fragment (rows 0-15 within warp tile). - βlower_frag (
Array[Scalar[epilogue_dtype], frag_size]): Lower fragment (rows 16-31 within warp tile). - βstage (
UInt32): Output stage index (for column offset computation). - βsrc_ptr (
Pointer[Scalar[c_type], address_space=AddressSpace.SHARED, _safe=False]): Pointer to source C SMEM tile. - βbeta (
Scalar[epilogue_dtype]): Residual scale factor.
Returns:
Tuple[Array[Scalar[epilogue_dtype], frag_size], Array[Scalar[epilogue_dtype], frag_size]]: Updated (upper_frag, lower_frag) tuple.
Was this page helpful?
Thank you! We'll create more content like this.
Thank you for helping us improve!