IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /max/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. /max/get-started.md).

Mojo function

saturated_reduce_kernel

def saturated_reduce_kernel[rank: Int, axis: Int, num_reductions: Int, BLOCK_SIZE: Int, input_fn: def[dtype: DType, width: Int, rank: Int](IndexList[rank]) capturing thin -> SIMD[dtype, width], output_fn: def[dtype: DType, width: SIMDLength, rank: Int](IndexList[rank], StaticTuple[SIMD[dtype, width], num_reductions]) capturing thin -> None, reduce_fn: def[ty: DType, width: SIMDLength, reduction_idx: Int](SIMD[ty, width], SIMD[ty, width]) capturing thin -> SIMD[ty, width], dtype: DType, simd_width: Int, accum_type: DType = get_accum_type[dtype]()](shape: IndexList[rank], init: StaticTuple[Scalar[dtype], num_reductions])

GPU kernel for reductions when the device is saturated with enough rows. Each thread independently reduces an entire row using SIMD packing, avoiding shared-memory synchronization entirely. Used when reducing along a non-contiguous axis.

Parameters:

  • ​rank (Int): The tensor rank.
  • ​axis (Int): The axis along which to reduce.
  • ​num_reductions (Int): The number of fused reductions to perform.
  • ​BLOCK_SIZE (Int): The number of threads per block.
  • ​input_fn (def[dtype: DType, width: Int, rank: Int](IndexList[rank]) capturing thin -> SIMD[dtype, width]): The lambda to load input elements.
  • ​output_fn (def[dtype: DType, width: SIMDLength, rank: Int](IndexList[rank], StaticTuple[SIMD[dtype, width], num_reductions]) capturing thin -> None): The lambda to store output elements.
  • ​reduce_fn (def[ty: DType, width: SIMDLength, reduction_idx: Int](SIMD[ty, width], SIMD[ty, width]) capturing thin -> SIMD[ty, width]): The binary reduction function.
  • ​dtype (DType): The data type of the elements.
  • ​simd_width (Int): The SIMD vector width.
  • ​accum_type (DType): The accumulator data type.

Args:

Was this page helpful?