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
persistent_topk_block
def persistent_topk_block[in_dtype: DType, //](ctx: DeviceContext, in_scores: Pointer[Scalar[in_dtype], ImmutAnyOrigin], out_idxs: Pointer[Int32, MutAnyOrigin], N: Int, K: Int, total_seq_len: Int, row_bounds: Optional[Pointer[Int32, ImmutAnyOrigin]] = None)
Launch block-wide bitonic top-k for total_seq_len score rows.
For N ≤ PERSISTENT_TOPK_MAX_N (= 2048) a single block sorts the whole row.
For N > PERSISTENT_TOPK_MAX_N a streaming variant folds _TILE-wide tiles
into a running top-_TILE champion; this requires K ≤ PERSISTENT_TOPK_MAX_N
(the champion width). Call sites needing K > PERSISTENT_TOPK_MAX_N must
use topk_gpu.
Each row of N scores yields the K highest-scoring column indices
(as int32) in descending score order in out_idxs.
Scores are read as in_dtype, f32 or bf16, and widened exactly on the way
in (see _widen_scores); everything below the load -- the keys, the digit
widths, the comparators, the champion buffers -- is f32 either way. The read
width is therefore a bandwidth decision, not a numerical one: whichever
dtype filled the buffer decided the scores' precision, and this selects the
K largest of exactly the values it is handed.
in_scores must be at least 16 B aligned -- every device allocator satisfies
this, but an offset pointer need not. The row stride N should be a multiple
of _elems_per_16b[in_dtype]() (4 at f32, 8 at bf16) or rows whose base
misses it fall to the scalar path; mla_index_fp8 pads its score stride for
exactly this reason. Unaligned strides stay correct, only slower.
Parameters:
- in_dtype (
DType): Element type of the score buffer, f32 or bf16 (inferred).
Args:
- ctx (
DeviceContext): Device context. - in_scores (
Pointer[Scalar[in_dtype], ImmutAnyOrigin]): Flat score buffer[total_seq_len × N]row-major. - out_idxs (
Pointer[Int32, MutAnyOrigin]): Output buffer[total_seq_len × K]row-major (int32). - N (
Int): Score columns per token (the row stride). - K (
Int): Top-k count per token (≤ N, and ≤ PERSISTENT_TOPK_MAX_N when N > 2048). - total_seq_len (
Int): Number of rows (one block per row). - row_bounds (
Optional[Pointer[Int32, ImmutAnyOrigin]]): Optional[total_seq_len]int32 per-row live-column counts. When set, rowrreads only its firstrow_bounds[r]columns; columns past the bound are never selected, never read (they may be uninitialized), and pad the output with-1. Scan cost tracks the real row lengths. Only supported forN ≤ PERSISTENT_TOPK_MAX_Nhere; wider rows go throughpersistent_topk_block_split.