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 struct

SplitParallelCombineParams

struct SplitParallelCombineParams[output_type: DType, accum_type: DType, num_splits: Int, ragged: Bool = False, has_attn_sink: Bool = False]

Holds the pointers, strides, and shape metadata for the split-parallel combine kernel.

Carries the per-split partial output and LSE accumulators, the final output tensor, optional ragged row offsets and attention-sink pointer, and the derived strides used by mla_combine_kernel_split_parallel. All 8 warps cooperate on a single head, so heads_per_block is fixed at 1.

Parameters​

  • ​output_type (DType): The element type of the per-split partial output accumulator and the final combined output tensor.
  • ​accum_type (DType): The element type of the per-split LSE accumulator values.
  • ​num_splits (Int): The compile-time number of KV cache splits to combine. Used to partition splits across the 8 warps.
  • ​ragged (Bool): Whether batches have variable-length sequences in a packed layout (defaults to False). When true, input_row_offsets_ptr selects each batch's token range.
  • ​has_attn_sink (Bool): Whether to apply attention-sink correction to the global LSE (defaults to False). When true, attn_sink_ptr must be provided.

Fields​

  • ​out_accum_split_ptr (Pointer[Scalar[output_type], MutAnyOrigin, _safe=False]):
  • ​lse_accum_split_ptr (Pointer[Scalar[accum_type], MutAnyOrigin, _safe=False]):
  • ​output_ptr (Pointer[Scalar[output_type], MutAnyOrigin, _safe=False]):
  • ​input_row_offsets_ptr (Pointer[UInt32, MutAnyOrigin, _safe=False]):
  • ​attn_sink_ptr (OptionalReg[Pointer[Float32, MutAnyOrigin, _safe=False]]):
  • ​batch_size (Int):
  • ​seq_len (Int):
  • ​num_heads (Int):
  • ​head_dim (Int):
  • ​lse_stride_split (Int):
  • ​lse_stride_batch (Int):
  • ​lse_stride_seq (Int):
  • ​out_accum_stride_split (Int):
  • ​out_accum_stride_head (Int):
  • ​out_stride_row (Int):

Implemented traits​

AnyType, Copyable, DevicePassable, ImplicitlyCopyable, ImplicitlyDeletable, Movable, RegisterPassable, TrivialRegisterPassable

comptime members​

device_type​

comptime device_type = SplitParallelCombineParams[output_type, accum_type, num_splits, ragged, has_attn_sink]

heads_per_block​

comptime heads_per_block = 1

num_threads​

comptime num_threads = (Int(8) * _resolve_warp_size())

num_warps​

comptime num_warps = 8

Methods​

__init__​

def __init__(out_accum_split_ptr: Pointer[Scalar[output_type], MutAnyOrigin, _safe=False], lse_accum_split_ptr: Pointer[Scalar[accum_type], MutAnyOrigin, _safe=False], output_ptr: Pointer[Scalar[output_type], MutAnyOrigin, _safe=False], input_row_offsets_ptr: Pointer[UInt32, MutAnyOrigin, _safe=False], attn_sink_ptr: OptionalReg[Pointer[Float32, MutAnyOrigin, _safe=False]], batch_size: Int, seq_len: Int, num_heads: Int, head_dim: Int) -> Self

get_type_name​

static def get_type_name() -> String

Returns:

String

get_device_type_name​

static def get_device_type_name() -> String

Returns:

String

Was this page helpful?