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 slotqmust hold theq-th largest score.Trueis the strong contract: descending score, ties by ascending column.Falsepromises the sameKcolumns, 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 sameKcolumns every run.Falsedrops 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 withordered=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 belowPERSISTENT_TOPK_MAX_N, where no radix select runs. - phi_dtype (
DType): The width the in-registerphipayload 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 isuint32at every score dtype, so a caller never needs to pass it. Forcinguint16at bf16 gives the narrow payload; it must be at least as wide as the score and at leastsig_bitswide. Likesig_bitsit 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 defaultphi_dtypeis 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; seepersistent_topk_block. Rowrscans only[0, row_bounds[r])of itsN-wide stripe, so under capture-frozen metadata (whereNis a worst-case bound) the scan cost tracks each row's real length. Supported on every path this launcher selects forN > 2048(the histogram-select family) and on the 2048 single-block path.