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

gemv_split_k

def gemv_split_k[c_type: DType, a_type: DType, b_type: DType, c_layout: TensorLayout, a_layout: TensorLayout, b_layout: TensorLayout, c_storage: TensorStorage, a_storage: TensorStorage, b_storage: TensorStorage, simd_width: Int, tile_m: Int, tile_n: Int, num_threads: Int, unroll_factor: Int = Int(2), weight_non_temporal: Bool = True, elementwise_lambda_fn: Optional[def[dtype: DType, width: SIMDLength, *, alignment: Int = Int(1)](IndexList[Int(2)], SIMD[dtype, width]) capturing thin -> None] = None, accum_type: DType = get_accum_type[c_type](), check_bounds_m: Bool = True, check_bounds_n: Bool = True, pdl_level: PDLLevel = PDLLevel()](output: TileTensor[c_type, c_layout, MutAnyOrigin, Storage=c_storage], act: TileTensor[a_type, a_layout, ImmutAnyOrigin, Storage=a_storage], weight: TileTensor[b_type, b_layout, ImmutAnyOrigin, Storage=b_storage], m: Int, n: Int, k: Int)

GEMV with tiling in K dimension. Assuming the B (weight) matrix is transposed i.e. row major N x K, this kernel implements a vector (1 x K) times a matrix (N x K). The impl can actually handle M > 1 but it's only optimal for tiny M. We use it for M = 1 only.

The launch grid covers ceildiv(m, tile_m) * tile_m rows and ceildiv(n, tile_n) * tile_n columns, so the final blocks read and write past the buffers unless the bounds guards are on: check_bounds_m=False is only safe when the launcher guarantees m % tile_m == 0 (m is a runtime value, so tile_m == 1 is the usual way to guarantee it), and check_bounds_n=False is only safe when n % tile_n == 0.

Parameters:

  • ​c_type (DType): Output element type.
  • ​a_type (DType): Activation matrix element type.
  • ​b_type (DType): Weight matrix element type.
  • ​c_layout (TensorLayout): Layout descriptor for the output tensor.
  • ​a_layout (TensorLayout): Layout descriptor for the activation matrix.
  • ​b_layout (TensorLayout): Layout descriptor for the weight matrix.
  • ​c_storage (TensorStorage): Storage kind for the output tensor.
  • ​a_storage (TensorStorage): Storage kind for the activation matrix.
  • ​b_storage (TensorStorage): Storage kind for the weight matrix.
  • ​simd_width (Int): Number of elements per vectorized load; sets tile_k with num_threads.
  • ​tile_m (Int): Number of output rows each thread accumulates.
  • ​tile_n (Int): Number of weight rows each thread accumulates.
  • ​num_threads (Int): Threads per block; sets K-tile width and cross-warp reduction count.
  • ​unroll_factor (Int): K-loop unroll factor for instruction-level parallelism (defaults to 2).
  • ​weight_non_temporal (Bool): When True, load the weight matrix with non-temporal (streaming) hints to bypass the cache (defaults to True).
  • ​elementwise_lambda_fn (Optional[def[dtype: DType, width: SIMDLength, *, alignment: Int = Int(1)](IndexList[Int(2)], SIMD[dtype, width]) capturing thin -> None]): Optional epilogue applied to each output element.
  • ​accum_type (DType): Accumulation precision type.
  • ​check_bounds_m (Bool): When True, guards M-tail rows when m is not a multiple of tile_m.
  • ​check_bounds_n (Bool): When True, guards N-tail columns when n is not a multiple of tile_n.
  • ​pdl_level (PDLLevel): Programmatic dependent launch level for PDL barriers.

Args: