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
mha_decoding_single_batch
def mha_decoding_single_batch[q_type: DType, k_t: MHAOperand, v_t: MHAOperand, output_type: DType, mask_t: MHAMask, *, BM: Int, BN: Int, BK: Int, WM: Int, WN: Int, depth: Int, num_heads: Int, num_threads: Int, num_pipeline_stages: Int, group: Int = Int(1), decoding_warp_split_k: Bool = False, sink: Bool = False](q_ptr: Pointer[Scalar[q_type], ImmutAnyOrigin, _safe=False], k: k_t, v: v_t, output_ptr: Pointer[Scalar[output_type], MutAnyOrigin, _safe=False], exp_sum_ptr: Pointer[Scalar[get_accum_type[q_type]()], MutAnyOrigin, _safe=False], qk_max_ptr: Pointer[Scalar[get_accum_type[q_type]()], MutAnyOrigin, _safe=False], scale: Float32, num_keys: Int, num_partitions: Int, mask: mask_t, batch_idx: Int, sink_weights: OptionalReg[LayoutTensor[q_type, Layout.row_major(Int(-1)), ImmutAnyOrigin]])
Flash attention v2 algorithm.
Parameters:
- βq_type (
DType): Element type of the query tensor. - βk_t (
MHAOperand): Key operand type (KV cache or dense tensor). - βv_t (
MHAOperand): Value operand type (KV cache or dense tensor). - βoutput_type (
DType): Element type of the output tensor. - βmask_t (
MHAMask): Attention mask type implementingMHAMask. - βBM (
Int): Number of query rows per thread block. - βBN (
Int): Number of key columns per thread block. - βBK (
Int): Tile size in the depth dimension for shared-memory tiles. - βWM (
Int): Warp tile height in the query (M) dimension. - βWN (
Int): Warp tile width in the key (N) dimension. - βdepth (
Int): Attention head depth (key/value dimension per head). - βnum_heads (
Int): Total number of query heads. - βnum_threads (
Int): Number of threads per thread block. - βnum_pipeline_stages (
Int): Number of software-pipeline stages for async copies. - βgroup (
Int): GQA group size, query heads per key/value head (defaults to 1). - βdecoding_warp_split_k (
Bool): Enable warp-level split-K reduction (defaults toFalse). - βsink (
Bool): Enable attention-sink mode where the first tokens always attend (defaults toFalse).
Args:
- βq_ptr (
Pointer[Scalar[q_type], ImmutAnyOrigin, _safe=False]): Pointer to the query tensor in global memory. - βk (
k_t): Key operand backed by a KV cache or dense tensor. - βv (
v_t): Value operand backed by a KV cache or dense tensor. - βoutput_ptr (
Pointer[Scalar[output_type], MutAnyOrigin, _safe=False]): Pointer to the output tensor in global memory. - βexp_sum_ptr (
Pointer[Scalar[get_accum_type[q_type]()], MutAnyOrigin, _safe=False]): Pointer to the per-head online-softmax denominator (sum of exponentials) for cross-partition reduction. - βqk_max_ptr (
Pointer[Scalar[get_accum_type[q_type]()], MutAnyOrigin, _safe=False]): Pointer to the per-head online-softmax running maximum for cross-partition reduction. - βscale (
Float32): Softmax temperature scale applied to QΒ·Kα΅. - βnum_keys (
Int): Number of valid key/value entries (cache length). - βnum_partitions (
Int): Number of split-K partitions along the key dimension. - βmask (
mask_t): Mask instance used to apply the attention mask. - βbatch_idx (
Int): Index of the sequence within the batch. - βsink_weights (
OptionalReg[LayoutTensor[q_type, Layout.row_major(Int(-1)), ImmutAnyOrigin]]): Optional sink-token weight tensor for attention sinks.
Was this page helpful?
Thank you! We'll create more content like this.
Thank you for helping us improve!