IMPORTANT: To view this page as Markdown, append `.md` to the URL (e.g. /get-started.md). For the complete documentation index, see llms.txt.
Skip to main content
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, row r reads only its first row_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 for N ≤ PERSISTENT_TOPK_MAX_N here; wider rows go through persistent_topk_block_split.

Was this page helpful?