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_split

def persistent_topk_block_split[in_dtype: DType, //, ordered: Bool = False, deterministic: Bool = False, sig_bits: Int = _hsel_sig_bits[in_dtype](), phi_dtype: DType = _hsel_phi_dtype[in_dtype](), scan_items: Int = _hsel_prefetch_scan_items[in_dtype, phi_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 bitonic top-k, splitting the N dimension when rows under-fill GPU.

Same contract and output as persistent_topk_block. When the row count is small relative to the SM count and N spans many tiles (the long-context decode regime — a handful of blocks would otherwise each fold the whole row serially), the streaming fold is split across rows * S blocks (phase 1), each producing a sorted top-_TILE partial, then merged per row (phase 2). All other shapes fall back to persistent_topk_block unchanged.

The score-dtype and pointer-alignment contract is persistent_topk_block's, unchanged.

Parameters:

  • ​in_dtype (DType): Element type of the score buffer, f32 or bf16 (inferred).
  • ​ordered (Bool): Whether slot q must hold the q-th largest score. True is the strong contract: descending score, ties by ascending column. False promises the same K columns, deterministically, in an unspecified order -- which lets the short-row path skip its ranking pass, most of its work. The set does not depend on this, because the tie-break lives in the key the select already compares. Only shapes that have a cheaper path take one; the rest stay ordered, which satisfies the weaker promise too.
  • ​deterministic (Bool): Whether one input must always give one output. True, the default, is the guarantee above: the same K columns every run. False drops it, and drops the set guarantee with it -- when more columns tie at the threshold than there are slots left, which of them survive is decided by the order the block's warps reach a cursor. In exchange the cheap contract's tail needs no ordering scan at all. Only meaningful together with ordered=False; the ordered path is deterministic by construction.
  • ​sig_bits (Int): How many of a key half's register a score can occupy, which is what sets the radix select's round count. Defaults to _hsel_sig_bits[in_dtype]() -- 16 at bf16 and 32 at f32 -- so a caller never needs to pass it. Forcing 32 at bf16 asks for the f32 round schedule on a bf16 buffer, which is a slower way to the same answer and exists so the two schedules can be compared inside one process. It reaches the schedule only, never the choice of kernel, and is irrelevant below PERSISTENT_TOPK_MAX_N, where no radix select runs.
  • ​phi_dtype (DType): The width the in-register phi payload is carried at, which is what sets how many columns a scan group or a resident row costs. Defaults to _hsel_phi_dtype[in_dtype](), which is uint32 at every score dtype, so a caller never needs to pass it. Forcing uint16 at bf16 gives the narrow payload; it must be at least as wide as the score and at least sig_bits wide. Like sig_bits it reaches the kernels only, never the dispatch, so both arms run the same instantiation family on the same grid.
  • ​scan_items (Int): Columns a thread carries per scan step on the arms that prefetch. Defaults to _hsel_prefetch_scan_items[in_dtype, phi_dtype]() -- 16 where the payload is narrow enough to hold them, 8 otherwise, which at the default phi_dtype is always 8. Must be a multiple of _elems_per_16b[in_dtype](), or a thread's group base lands mid-vector for some threads and not others. The arms that do not prefetch always take the default width; they hold no group across the loop, so a wider one would cost registers for nothing.

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.
  • ​row_bounds (Optional[Pointer[Int32, ImmutAnyOrigin]]): Optional [total_seq_len] int32 per-row live-column counts; see persistent_topk_block. Row r scans only [0, row_bounds[r]) of its N-wide stripe, so under capture-frozen metadata (where N is a worst-case bound) the scan cost tracks each row's real length. Supported on every path this launcher selects for N > 2048 (the histogram-select family) and on the 2048 single-block path.

Was this page helpful?