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
mla_indexer_ragged_float8_paged
def mla_indexer_ragged_float8_paged[oi_layout: TensorLayout, q_layout: TensorLayout, qs_layout: TensorLayout, iro_layout: TensorLayout, //, dtype: DType, KCollectionT: KVCollectionT, num_heads: Int, depth: Int, top_k: Int, mask_str: StringSpan[ImmStaticOrigin], scores_dtype: DType = .bfloat16, kpool: Int = Int(1)](output_indices: TileTensor[.int32, oi_layout], q: TileTensor[dtype, q_layout], q_s: TileTensor[.float32, qs_layout], input_row_offsets: TileTensor[.uint32, iro_layout], k_collection: KCollectionT, layer_idx: UInt32, ctx: DeviceContext, scores_budget_bytes: Int = Int((mul get_defined_int[StringSpan("MLA_INDEX_SCORES_BUDGET_MB"), Int(512)](), 1048576)))
Compute FP8 indexed attention scores using paged KV cache and return top-k indices.
This function:
- Computes FP8 matmul between q and cached k (with scales), aggregated across heads
- Applies the specified mask (causal, etc.)
- Computes top-k indices per token (scores are summed across all heads)
Parameters:
- oi_layout (
TensorLayout): Layout of the top-k index output. - q_layout (
TensorLayout): Layout of the query tensor. - qs_layout (
TensorLayout): Layout of the query scales. - iro_layout (
TensorLayout): Layout of the ragged query row offsets. - dtype (
DType): Element type of theqquery tensor, an FP8 dtype. - KCollectionT (
KVCollectionT): Type of the KV collection holding cached K values and K scales. - num_heads (
Int): Number of attention heads per token. - depth (
Int): Per-head key dimension (head size) in elements. - top_k (
Int): Requested number of top-scoring key indices to select per token. - mask_str (
StringSpan[ImmStaticOrigin]): Name of the mask to apply, eitherMaskName.NULLorMaskName.CAUSAL. - scores_dtype (
DType): Element type of the transient score matrix.float32orbfloat16; the latter halves the buffer, so a row window under a fixed byte budget holds twice the rows. Honoured only on the SM100 scorers -- the scalar fallback is f32-only and resolves to it, which no caller can observe because the matrix is internal. - kpool (
Int): Tokens per pooled cache row.1scores one row per token;k > 1scores one pooled key perkconsecutive tokens, so every candidate count and the caller'stop_kare pool-granular.
Args:
- output_indices (
TileTensor[.int32, oi_layout]): Dense output tensor for top-k indices [total_seq_len, top_k]. Invalid positions (where there are fewer than top_k valid keys due to causal masking or shorter sequences) are filled with -1. - q (
TileTensor[dtype, q_layout]): Query tensor [total_seq_len, num_heads, head_dim] in FP8. - q_s (
TileTensor[.float32, qs_layout]): Query scales [total_seq_len, num_heads] in float32. - input_row_offsets (
TileTensor[.uint32, iro_layout]): Ragged row offsets for queries [batch_size + 1]. - k_collection (
KCollectionT): KV collection containing cached K values and K scales. K scales are accessed via k_cache.scales (quantization_granularity=head_size). - layer_idx (
UInt32): Layer index for retrieving cache. - ctx (
DeviceContext): Device context. - scores_budget_bytes (
Int): Peak bytes the transient score matrix may occupy. Longer batches are scored a row-window at a time to stay under it (see the chunking below). Runtime rather than comptime because nothing below it is specialized on the value -- it reaches only therows_per_chunkarithmetic -- so a sweep over budgets costs no recompiles. Exposed so tests can force a window small enough to exercise the multi-chunk path on toy shapes.