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
generic_flash_attention_kv_cache_ragged
def generic_flash_attention_kv_cache_ragged[collection_t: KVCollectionT, dtype: DType, //, *, target: StringSlice[ImmStaticOrigin], mask_str: StringSlice[ImmStaticOrigin], local_window_size: Int = Int(-1), output_dtype: DType = dtype](q: LayoutTensor[dtype, element_layout=q.element_layout, layout_int_type=q.layout_int_type, linear_idx_type=q.linear_idx_type, masked=q.masked, alignment=q.alignment], input_row_offsets: LayoutTensor[DType.uint32, Layout.row_major(Int(-1)), ImmutAnyOrigin], kv_collection: collection_t, layer_idx: UInt32, scale: Float32, output: LayoutTensor[output_dtype, element_layout=output.element_layout, layout_int_type=output.layout_int_type, linear_idx_type=output.linear_idx_type, masked=output.masked, alignment=output.alignment], context: DeviceContext, decode_dispatch_metadata: MHADecodeDispatchMetadata)
Dispatches flash attention over a ragged batch against a paged KV cache.
Parameters:
- βcollection_t (
KVCollectionT): The KV cache collection type storing the K and V caches for this layer (inferred). - βdtype (
DType): Data type of the query tensor (inferred). - βtarget (
StringSlice[ImmStaticOrigin]): Target device string for kernel dispatch. - βmask_str (
StringSlice[ImmStaticOrigin]): Attention mask name selecting the masking strategy, such as "causal", "null", or "sliding_window_causal". - βlocal_window_size (
Int): Sliding-window size in tokens for windowed masks; -1 for masks that ignore it (defaults to -1). - βoutput_dtype (
DType): Data type of theoutputtensor (defaults todtype).
Args:
- βq (
LayoutTensor[dtype, element_layout=q.element_layout, layout_int_type=q.layout_int_type, linear_idx_type=q.linear_idx_type, masked=q.masked, alignment=q.alignment]): Query tensor with shape (sum(seq_lens), num_heads, head_size). - βinput_row_offsets (
LayoutTensor[DType.uint32, Layout.row_major(Int(-1)), ImmutAnyOrigin]): Tensor with shape (batch_size + 1,) denoting the start of each sequence along the ragged sequence dimension. - βkv_collection (
collection_t): The collection storing the KVCache entries for this layer, retrieved via layer_idx. - βlayer_idx (
UInt32): The index of the layer being executed, used to retrieve the KVCache objects from kv_collection. - βscale (
Float32): The scaling factor in scaled dot-product attention, usually rsqrt(head_size). - βoutput (
LayoutTensor[output_dtype, element_layout=output.element_layout, layout_int_type=output.layout_int_type, linear_idx_type=output.linear_idx_type, masked=output.masked, alignment=output.alignment]): The pre-allocated output buffer to write results to, with shape (sum(seq_lens), num_heads, head_size). - βcontext (
DeviceContext): The call context pointer, passed by the graph compiler. - βdecode_dispatch_metadata (
MHADecodeDispatchMetadata): Precomputed dispatch metadata used to select decode kernels for the GPU target.
Was this page helpful?
Thank you! We'll create more content like this.
Thank you for helping us improve!