For the complete documentation index, see llms.txt. Markdown versions of all pages are available by appending .md to any URL (e.g. /get-started.md).
Mojo function
mha_decoding_single_batch_pipelined
def mha_decoding_single_batch_pipelined[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], k: k_t, v: v_t, output_ptr: Pointer[Scalar[output_type], MutAnyOrigin], exp_sum_ptr: Pointer[Scalar[get_accum_type[q_type]()], MutAnyOrigin], qk_max_ptr: Pointer[Scalar[get_accum_type[q_type]()], MutAnyOrigin], scale: Float32, num_keys: Int, num_partitions: Int, sink_weights: OptionalReg[LayoutTensor[q_type, Layout.row_major(Int(-1)), ImmutAnyOrigin]], mask: mask_t, batch_idx: Int)
Flash attention v2 decode kernel for a single batch element with pipelined multistage MMA.
Computes attention for the decoding (single-query) case using the FA2
online-softmax algorithm with multistage pipelining of K/V loads. When
num_partitions exceeds 1, each block processes a contiguous slice of
the key dimension and writes partial exp_sum and qk_max statistics
for a subsequent mha_splitk_reduce pass.
Parameters:
- q_type (
DType): Element type of the query tensor (inferred). - k_t (
MHAOperand): Key operand type backing the key tensor (inferred). - v_t (
MHAOperand): Value operand type backing the value tensor (inferred). - output_type (
DType): Element type of the output tensor (inferred). - mask_t (
MHAMask): Attention mask type implementingMHAMask(inferred). - BM (
Int): Number of query rows processed per thread block. - BN (
Int): Number of key columns per thread block tile. - BK (
Int): Tile size in the head-depth dimension. - 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 pipeline stages for the multistage MMA loads. - group (
Int): GQA group size, query heads per key/value head (defaults to 1). - decoding_warp_split_k (
Bool): Enable warp-level split-K for decode (defaults toFalse). - sink (
Bool): Enable attention-sink mode where the first tokens always attend (defaults toFalse).
Args:
- q_ptr (
Pointer[Scalar[q_type], ImmutAnyOrigin]): Pointer to the query tensor for this batch element. - k (
k_t): Key operand backed by a KV cache. - v (
v_t): Value operand backed by a KV cache. - output_ptr (
Pointer[Scalar[output_type], MutAnyOrigin]): Pointer to the output tensor for this batch element. - exp_sum_ptr (
Pointer[Scalar[get_accum_type[q_type]()], MutAnyOrigin]): Pointer to the online-softmax exponent sum buffer for this batch. - qk_max_ptr (
Pointer[Scalar[get_accum_type[q_type]()], MutAnyOrigin]): Pointer to the online-softmax running maximum buffer for this batch. - scale (
Float32): Softmax temperature scale applied to Q·Kᵀ. - num_keys (
Int): Total number of key/value entries (cache length) for this batch. - num_partitions (
Int): Number of split-K partitions dividing the key dimension. - sink_weights (
OptionalReg[LayoutTensor[q_type, Layout.row_major(Int(-1)), ImmutAnyOrigin]]): Optional sink-token weight tensor for attention sinks. - mask (
mask_t): Mask instance used to apply the attention mask. - batch_idx (
Int): Index of the batch element this block processes.