Skip to content

Binary heap top-k in C (optional)

Moduleds.04 · side · C · Pass 6 · 2 to 3 h
You buildc/src/ds/topk.c: tl_topk_f32, the kk largest values of a float array and their indices, best first, by a size-kk min-heap that lives in the caller’s output buffers
Contractcourse/contracts/c/include/tinyllm/topk.h · rules: c/ABI.md
TestsStandalone C test binary test_topk.c, under ASan and UBSan with a property test against a full sort; parity uses fixture files generated by the Python reference (what they check: section 4)
Needsrt.02 (C allocation and error support) · reading: ds.06 the first heap chapter and ds.01 (tl_vec; the heap itself allocates nothing)
Used byNo runtime integration. This optional C module is checked independently against shared fixtures.
MilestoneMS-L9, the optional standalone C module group
Optional depthCormen et al., Introduction to Algorithms (4th ed.), chapter 6 (heaps) and section 9.2 (selection in expected linear time); Knuth, TAOCP vol. 3, section 5.3.3 (minimum-comparison selection)
  • Top-kk needs only a heap of size kk: keep the best kk seen so far with the worst of them at the root, and each new value costs one comparison against the root, plus O(log⁡k)O(\log k) swaps only when it gets in (hand_example, matches_a_full_sort).
  • The order is total: a smaller value ranks lower, and among equal values the larger index ranks lower. That one rule settles both who wins a tie at the kk-th place (the lower index) and the output order of equal values (ascending index) (tie_at_the_cut_goes_to_the_lower_index, equal_values_are_listed_by_ascending_index).
  • NaN is rejected, −∞-\infty is a value. NaN has no place in any order; a masked logit is −∞-\infty and simply ranks last (nan_and_bad_k_are_einval_and_write_nothing, minus_infinity_is_an_ordinary_value).
  • No allocation: the heap is built inside idx and val, and a final in-place heap sort leaves them best first.
  • The C implementation and Python reference agree on the kept set through shared fixture files.
Terminal window
ol start ds.04 # stubs c/src/ds/topk.c into your repo
ol tests ds.04 # read the test catalog first
ol check ds.04 # exit code is the verdict
ol check ds.04 --ref-deps # only if rt.02 is not passing yet
ol diff ds.04 # after passing: your code against the reference

Your Python sampler (L8.1) specifies top-kk by sorting every logit: process_logits orders all VV ids by (logit descending, id ascending) and keeps the first kk. For SmolLM2, V=49152V = 49152. Sorting 49152 floats to keep 40 takes about Vlog⁡2V≈770,000V \log_2 V \approx 770{,}000 comparisons where roughly VV suffice. This optional module implements the same ordering in C as a standalone exercise. It has its own C test binary, and fixture files generated by the Python reference check parity; no engine calls into this C code.

SymbolMeaningType / shape
x∈Rnx \in \mathbb{R}^nthe input values (logits), xix_i at index iifloat[n]
nnnumber of values, VV for a logit rowint64_t
kknumber of values to keep, 0≤k≤n0 \le k \le nint64_t
(v,i)(v, i)an entry: a value and its index
a≺ba \prec bentry aa ranks below (“is worse than”) entry bb
idx,val\mathit{idx}, \mathit{val}the outputs: the kept indices and values, best firstint32_t[k], float[k]
hhthe heap’s current size, h≤kh \le k

Top-kk is defined by an order. Values alone do not give one: two logits can be equal, and then “the kk largest” is ambiguous. The sampling spec (spec/sampling.md step 6) breaks ties by index, so define

(va,ia)≺(vb,ib)  ⟺  va<vb  or  (va=vb and ia>ib).(v_a, i_a) \prec (v_b, i_b) \iff v_a < v_b \ \text{ or } \ (v_a = v_b \text{ and } i_a > i_b).

Read it as ”aa is worse than bb”. Indices are distinct, so for any two different entries exactly one of a≺ba \prec b and b≺ab \prec a holds: the order is total. The top-kk set is then unique: the kk entries that are worse than no more than k−1k - 1 others. Listing it best first gives values in descending order and, among equal values, ascending indices.

NaN breaks this. IEEE 754 makes every comparison with NaN false, so a NaN entry would be neither worse nor better than anything, and the result would depend on where in the array it happened to sit. The contract therefore rejects any NaN with TL_EINVAL. −∞-\infty is different: −∞<v-\infty < v for every other value, and −∞=−∞-\infty = -\infty, so it fits the order and ranks last. A logit masked by constrained decoding (L8.7) is −∞-\infty, so rejecting it would break every masked request.

2.2 The heap keeps the worst kept entry at the root

Section titled “2.2 The heap keeps the worst kept entry at the root”

A binary heap (the ds.06 chapter) is a complete binary tree stored in an array: the children of position pp are 2p+12p + 1 and 2p+22p + 2, its parent is ⌊(p−1)/2⌋\lfloor (p - 1)/2 \rfloor. Here the heap property is: no parent is better than its children, so the root is the worst entry in the heap. That is a min-heap under ≺\prec.

The algorithm scans xx once, keeping the best h≤kh \le k entries seen so far in the heap:

  1. While h<kh < k, append (xi,i)(x_i, i) at position hh and sift it up: swap it with its parent while it is worse than the parent.
  2. Once h=kh = k, compare (xi,i)(x_i, i) with the root, the worst of the kk kept entries. If the root is worse, the new entry replaces it and is sifted down: swapped with its worse child while that child is worse than it. Otherwise the new entry is not among the kk best seen so far, and it is dropped.

Invariant (S-M05 style): after processing x0,…,xix_0, \dots, x_{i}, the heap holds exactly the top-min⁡(k,i+1)\min(k, i+1) entries of that prefix. It holds after step 1 trivially. In step 2, if the new entry ranks below the root, it ranks below all kk kept entries and cannot be in the top kk; otherwise the root is now ranked below kk entries (the other k−1k - 1 and the new one) and must leave.

Ties need no special code. The scan goes in increasing index, so a later entry with the same value as the root has a larger index: it is worse, and it stays out. The lower index wins the last place, as the spec requires.

After the scan, idx and val hold the top kk in heap order. A heap sort turns that into best-first order without extra memory: swap the root (the worst) with the last position, shrink the heap by one, sift the new root down, and repeat. The worst entry ends at position k−1k - 1, the second worst at k−2k - 2, and position 0 ends with the best.

Each of the nn values costs one comparison with the root. An entry that gets in costs O(log⁡k)O(\log k) swaps, and the final sort costs O(klog⁡k)O(k \log k). In the worst case (an increasing array, where every value gets in) the total is O(nlog⁡k)O(n \log k). For a logit row the order is close to random, and then the ii-th value gets in with probability about k/ik / i, so the expected number of insertions is about kln⁡(n/k)k \ln(n/k): for n=49152n = 49152 and k=40k = 40, about 280 insertions and 49152 root comparisons, against nlog⁡2n≈770,000n \log_2 n \approx 770{,}000 comparisons for a full sort. Quickselect (Hoare) finds the kk-th value in expected O(n)O(n) but reorders the input, which the contract forbids (x is const), so it would need a copy of all nn values.

The contract allows 0 <= k <= n and requires idx and val untouched when it returns TL_EINVAL. So the function checks the ranges and pointers, then scans all of xx for NaN, and only then writes. A NaN at the end of the row is found after the whole scan; that costs one extra pass over xx, which is cheap next to the heap work and keeps the error path simple.

x=[1,3,2,3,−1]x = [1, 3, 2, 3, -1] (the logits of the sampling spec’s worked example), k=3k = 3. Entries are written (v,i)(v, i), the heap as an array, root first.

StepEntryActionHeap after
i=0i = 0(1,0)(1, 0)h=0<3h = 0 < 3: append, nothing to sift[(1,0)][(1,0)]
i=1i = 1(3,1)(3, 1)append at 1; parent (1,0)(1,0) is worse, so no swap[(1,0),(3,1)][(1,0), (3,1)]
i=2i = 2(2,2)(2, 2)append at 2; parent (1,0)(1,0) is worse, no swap[(1,0),(3,1),(2,2)][(1,0), (3,1), (2,2)]
i=3i = 3(3,3)(3, 3)full; root (1,0)≺(3,3)(1,0) \prec (3,3): replace root. Sift down: children (3,1)(3,1) and (2,2)(2,2); (2,2)≺(3,3)(2,2) \prec (3,3) and (3,1)(3,1) is not (equal value, smaller index), so swap with (2,2)(2,2)[(2,2),(3,1),(3,3)][(2,2), (3,1), (3,3)]
i=4i = 4(−1,4)(-1, 4)root (2,2)(2,2) is not worse than (−1,4)(-1, 4): dropunchanged

Heap sort: swap root and position 2, giving [(3,3),(3,1)∣(2,2)][(3,3), (3,1) \mid (2,2)]; sift down in a heap of 2: (3,1)(3,1) is not worse than (3,3)(3,3), stop. Swap root and position 1: [(3,1)∣(3,3),(2,2)][(3,1) \mid (3,3), (2,2)].

Result: idx = {1, 3, 2}, val = {3, 3, 2}. The two 3s tie, and the lower index comes first. The spec’s own trace keeps ids 1, 3, 2 in that order: the first test, hand_example.

tinyllm/topk.h
tl_status tl_topk_f32(const float *x, int64_t n, int64_t k, int32_t *idx, float *val);
/* The k largest values of x[0..n) and their indices, best first; equal values by
ascending index (and the lower index wins the last place). -inf is a value.
TL_EINVAL, idx and val untouched, for k < 0 or k > n, a NULL pointer that would
be used, or any NaN in x. k == 0 writes nothing. Allocates nothing. */

The C test binary calls this contract directly. Python produces the reference fixture rows for parity checks; it does not load the C implementation.

TestKINDChecksWhy it matters downstream
hand_exampleunit, smokesection 3 exactlyyou and the tests agree on the definition
tie_at_the_cut_goes_to_the_lower_indexboundary[5,1,1,1,1,0][5, 1, 1, 1, 1, 0], k=2k = 2 keeps ids 0 and 1the ordering is deterministic
equal_values_are_listed_by_ascending_indexboundaryall-equal input comes out in index ordertop-p walks the kept ids in this order
k_zero_and_k_equal_nboundaryk=0k = 0 writes nothing; k=nk = n is a full sort; n=0n = 0 is finetop_k = 0 means off
minus_infinity_is_an_ordinary_valueboundary−∞-\infty ranks last, ties by indexmasked logits from L8.7
nan_and_bad_k_are_einval_and_write_nothingboundaryNaN, k>nk > n, k<0k < 0, NULL pointers give TL_EINVAL and leave the outputs as they wereerrors instead of garbage
matches_a_full_sortproperty, differential300 random rows with many ties against qsort by ≺\prec, every kkheap shapes only random sizes reach
PitfallSymptomCaught by
replacing the root when the new value is greater or equala tie at the kk-th place keeps the last tied indextie_at_the_cut_goes_to_the_lower_index (mutant s01)
comparing values only inside the heapequal values come out in heap order, not index orderequal_values_are_listed_by_ascending_index (mutant s02)
no NaN checkthe result depends on where the NaN sitsnan_and_bad_k_are_einval_and_write_nothing (mutant s03)
no k > n checkthe heap writes past idx and val (ASan)nan_and_bad_k_are_einval_and_write_nothing (mutant s04)
forgetting the final sortthe right set in heap order; top-p (step 7) walks it in the wrong orderhand_example (mutant s05)
a max-heap (best at the root)the root comparison drops the wrong entrieshand_example, matches_a_full_sort (mutant s06)
treating k=0k = 0 as an errora request with top_k = 0 (off) failsk_zero_and_k_equal_n (mutant s07)
writing before validatinga rejected call has already changed the outputsnan_and_bad_k_are_einval_and_write_nothing (mutant s08)
rejecting with !isfinite instead of isnanevery masked request failsminus_infinity_is_an_ordinary_value (mutant s09)
DirectionModuleHow it uses this
Backrt.02allocation and error support for the optional standalone C module
BackL8.1process_logits is the Python reference; shared fixtures provide parity inputs
Backds.06the first heap chapter: sift up, sift down, the array layout (reading)
ForwardL10.1Rust implements its sampler independently and checks outputs against shared fixture files

This module is optional and standalone; no core module depends on it.

Your pieceProduction equivalentWhat it addsWhere to look
size-kk heap over one rowPyTorch torch.topk (CPU)a partial sort (std::partial_sort / nth_element) chosen by k/nk/n, many rows in parallelaten/src/ATen/native/TopKImpl.h
CPU top-kkvLLM and FlashInfer GPU samplerstop-kk and top-pp fused with sampling, by rejection instead of sortingFlashInfer sampling.cuh
one row per callllama.cpp llama_sampler_top_kpartial sort over the candidate array, reused by every sampler in the chainsrc/llama-sampling.cpp