GCC Code Coverage Report


Directory: src/
File: src/index/centroid_search.c
Date: 2026-09-30 11:11:31
Exec Total Coverage
Lines: 344 402 85.6%
Functions: 8 8 100.0%
Branches: 138 187 73.8%

Line Branch Exec Source
1 /*
2 * Copyright (c) 2026 Tiger Data, Inc.
3 * Licensed under the PostgreSQL License. See LICENSE for details.
4 *
5 * centroid_search.c - Beam search over centroid tree
6 *
7 * Implements level-by-level descent through centroid pages using
8 * format-aware distance computation. RaBitQ pages use approximate
9 * distance with error bounds; float and half pages use exact L2
10 * distance (error = 0). Candidate selection uses VsTopK for
11 * error-bound-aware pruning.
12 */
13
14 #include "vs_config.h"
15
16 #include <math.h>
17 #include <string.h>
18 #include <time.h>
19
20 #include "algo/topk.h"
21 #include "algo/vecops.h"
22 #include "core/log.h"
23 #include "core/memory.h"
24 #include "core/platform.h"
25 #include "index/centroid_search.h"
26 #include "types/vec16.h"
27
28 /* ----------------------------------------------------------------
29 * Internal candidate for beam search
30 * ---------------------------------------------------------------- */
31 typedef struct Candidate
32 {
33 BlockNumber child_blkno; /* next level's page (or posting head) */
34 ItemPointerData origin; /* (page, entry) this came from */
35 Distance distance;
36 Distance error; /* symmetric error (0 for exact) */
37 } Candidate;
38
39 /* ----------------------------------------------------------------
40 * Per-scan scratch (public allocation; see centroid_search.h)
41 *
42 * Holds the candidate-buffer pair and score-page scratch that beam
43 * search previously palloc'd per call. Allocated once at scan setup
44 * and reused across all queries on the scan.
45 * ---------------------------------------------------------------- */
46 struct PrismCentroidScratch
47 {
48 /* Owning context, captured at scratch_create: every buffer below
49 * lives here until scratch_free. Beam search runs under whatever
50 * context the caller routes from — during index builds that is a
51 * per-row scratch context reset after every tuple — so any lazy
52 * (re)allocation of a scratch buffer must go to this context, never
53 * to the current one. */
54 VsMemCtx memctx;
55 uint32_t cand_cap; /* size of buf_a / buf_b */
56 uint32_t max_per_page; /* size of the f_add..symmetric_scratch arrays */
57 Candidate *buf_a;
58 Candidate *buf_b;
59 float *f_add;
60 float *f_rescale;
61 Distance *distances;
62 Distance *lower_bounds;
63 float *multi_scratch;
64 uint32_t *symmetric_scratch;
65 /* Reusable top-K + extraction buffer for select_topk_bounded.
66 * Avoids creating a fresh memctx + ub_heap + ub_ids + candidates
67 * + entries-buf on every beam-search level (was 2 sets of 4 allocs
68 * + 2 memctx creates per query). The VsTopK is initialised once
69 * at scratch_create with the worst-case k; select_topk_bounded
70 * calls vs_topk_reset_to_k() to adjust between levels. */
71 VsTopK level_topk;
72 VsTopKEntry *entries_buf;
73 uint32_t entries_cap;
74 /* Fastscan LUT used by the FASTSCAN centroid format. The LUT only
75 * depends on the query (qstate->transformed), which is constant
76 * for the entire centroid descent — so we build it once per query
77 * and reuse for every FASTSCAN page. fs_lut_valid is cleared at
78 * the start of every beam-search call. */
79 uint8_t *fs_lut;
80 uint32_t fs_lut_bytes;
81 float fs_lut_delta;
82 float fs_lut_bias;
83 bool fs_lut_valid;
84 /* Fine-grained timing accumulators (ns), reset per beam search. */
85 uint64_t t_lut_ns;
86 uint64_t t_pageread_ns;
87 uint64_t t_score_ns;
88 };
89
90 /* Monotonic nanosecond clock for fine-grained centroid instrumentation. */
91 static inline uint64_t
92 36124342 cs_now_ns(void)
93 {
94 28527566 struct timespec ts;
95 36124342 clock_gettime(CLOCK_MONOTONIC, &ts);
96 36124342 return (uint64_t)ts.tv_sec * VS_NS_PER_SEC + (uint64_t)ts.tv_nsec;
97 }
98
99 PrismCentroidScratch *
100 18119 prism_centroid_scratch_create(Dimension dim, uint32_t max_beam_width)
101 {
102 18119 uint32_t max_per_page = prism_centroid_max_entries(dim);
103
104 /* A kept node's children can span two pages when fan_out exceeds
105 * the per-page entry capacity (e.g. fan_out 74 vs 72 entries at
106 * dim 768), so a level can expose up to beam * 2 pages of
107 * candidates; sizing by pages * capacity keeps the buffer an upper
108 * bound and stops the level scan from silently truncating the
109 * farthest kept parent's tail children. */
110 18119 uint32_t cand_cap = max_beam_width * max_per_page * 2;
111
2/2
✓ Branch 0 taken 2 times.
✓ Branch 1 taken 148 times.
18119 if (cand_cap < max_per_page * 4)
112 2 cand_cap = max_per_page * 4;
113
114 /* The scratch owns a dedicated child context: every buffer —
115 * including later growth, which can run under a caller's per-row
116 * reset context — lives and dies with it, and cleanup is a single
117 * context delete. */
118 17969 VsMemCtx ctx =
119 18119 vs_memctx_create(vs_memctx_current(), "vs centroid scratch");
120 18119 VsMemCtx old_ctx = vs_memctx_switch(ctx);
121 18119 PrismCentroidScratch *s = vs_alloc(sizeof(PrismCentroidScratch));
122
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 18119 times.
18119 if (s == NULL)
123 {
124 ✗ vs_memctx_switch(old_ctx);
125 ✗ vs_memctx_delete(ctx);
126 ✗ return NULL;
127 }
128
129 18119 s->memctx = vs_memctx_current();
130 18119 s->cand_cap = cand_cap;
131 18119 s->max_per_page = max_per_page;
132 18119 s->buf_a = vs_alloc(cand_cap * sizeof(Candidate));
133 18119 s->buf_b = vs_alloc(cand_cap * sizeof(Candidate));
134 18119 s->f_add = vs_alloc(max_per_page * sizeof(float));
135 18119 s->f_rescale = vs_alloc(max_per_page * sizeof(float));
136 18119 s->distances = vs_alloc(max_per_page * sizeof(Distance));
137 18119 s->lower_bounds = vs_alloc(max_per_page * sizeof(Distance));
138 18119 s->multi_scratch = vs_alloc(max_per_page * sizeof(float));
139 18119 s->symmetric_scratch = vs_alloc(max_per_page * sizeof(uint32_t));
140
141 /* Reusable top-K and extract buffer (resized to actual k per
142 * select_topk_bounded call). Initial k=max_beam_width is just
143 * a starting size — the reset path repalloc's within the
144 * topk's memctx for different k. */
145 18119 vs_topk_init(&s->level_topk, max_beam_width);
146 18119 s->entries_cap = max_beam_width * 4;
147
2/2
✓ Branch 0 taken 17608 times.
✓ Branch 1 taken 511 times.
18119 if (s->entries_cap < 64)
148 17608 s->entries_cap = 64;
149 18119 s->entries_buf = vs_alloc(s->entries_cap * sizeof(VsTopKEntry));
150
151 /* Fastscan LUT (worst-case hacc size for this dim). Allocated
152 * once and reused for every centroid page scored in the
153 * FASTSCAN format. The LUT depends on the query so it's rebuilt
154 * per page; the buffer is reusable. */
155 18119 s->fs_lut_bytes = VS_FASTSCAN_LUT_HACC_BYTES(dim);
156 18119 s->fs_lut = vs_alloc(s->fs_lut_bytes);
157 18119 vs_memctx_switch(old_ctx);
158 18119 return s;
159 }
160
161 void
162 18117 prism_centroid_scratch_free(PrismCentroidScratch *s)
163 {
164
2/2
✓ Branch 0 taken 17967 times.
✓ Branch 1 taken 150 times.
18117 if (s == NULL)
165 ✗ return;
166 18117 vs_topk_cleanup(&s->level_topk);
167 /* Everything the scratch owns — struct included — lives in its
168 * context; one delete frees it all. */
169 18117 vs_memctx_delete(s->memctx);
170 }
171
172 /*
173 * Build-time exact-scoring hook: resolve a page's exact-centroid slots,
174 * or NULL when the hook is unset, the page lies outside the collected
175 * region, or the page carries no internal entries (the leaf level).
176 */
177 static const float *
178 5217888 exact_internal_slots(
179 const PrismCentroidSearchState *state,
180 BlockNumber page_blkno,
181 Dimension dim)
182 {
183 5217888 const PrismExactInternalCentroids *ex = state->exact_internal;
184
185
4/4
✓ Branch 0 taken 5049216 times.
✓ Branch 1 taken 168672 times.
✓ Branch 2 taken 3929403 times.
✓ Branch 3 taken 1119813 times.
5217888 if (ex == NULL || page_blkno < ex->base ||
186
2/2
✓ Branch 0 taken 2956203 times.
✓ Branch 1 taken 973200 times.
3929403 page_blkno - ex->base >= ex->npages)
187 97262 return NULL;
188
189 3929403 uint32_t off = ex->page_off[page_blkno - ex->base];
190
2/2
✓ Branch 0 taken 1158578 times.
✓ Branch 1 taken 2770825 times.
3929403 if (off == PRISM_EXACT_INTERNAL_NONE)
191 921200 return NULL;
192 289378 return ex->cents + (size_t)off * dim;
193 }
194
195 /*
196 * Score all centroids on a single page, appending to candidates.
197 * Returns the new candidate count.
198 *
199 * Dispatches based on page data format:
200 * RABITQ → batch multi-candidate scoring via cs
201 * FLOAT → vs_l2_distance_squared (exact, error=0)
202 * HALF → vs_f16_l2_squared (exact, error=0)
203 */
204 static uint32_t
205 5644020 score_page(
206 const PrismCentroidSearchState *state,
207 Page page,
208 BlockNumber page_blkno,
209 Dimension dim,
210 Candidate *cands,
211 uint32_t cand_count,
212 uint32_t cand_cap,
213 PrismCentroidScratch *cs)
214 {
215 5644020 PrismCentroidPageOpaque *opaque = PRISM_CENTROID_OPAQUE(page);
216 5644020 uint16_t count = opaque->entry_count;
217 5644020 PrismCentroidFormat fmt = prism_centroid_page_format(page);
218
219 /* The on-disk entry_count drives every loop below -- the RABITQ
220 * branch fills the max_per_page-sized scratch arrays with `count`
221 * entries, and the other branches read `count` entries off the page.
222 * A count past the format's real capacity (corruption, a truncated
223 * write) would overrun the scratch or read past the page, so reject
224 * it loudly rather than act on it. The RABITQ capacity equals the
225 * scratch size the state was built with (prism_centroid_max_entries). */
226 5644020 uint32_t max_entries = prism_centroid_max_entries_fmt(dim, fmt);
227
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 5644020 times.
5644020 if (count > max_entries)
228 ✗ vs_error(
229 VS_EXTENSION_NAME
230 ": centroid page %u has an invalid entry count (%u > %u); "
231 "the index may be corrupted -- REINDEX it",
232 page_blkno,
233 (unsigned)count,
234 max_entries);
235
236
2/2
✓ Branch 0 taken 4447540 times.
✓ Branch 1 taken 1196480 times.
5644020 if (count == 0)
237 ✗ return cand_count;
238
239 /* Build-time hook: when the build collected this page's exact float
240 * centroids (internal levels only; see PrismExactInternalCentroids), score
241 * them with exact L2 instead of the RaBitQ estimate. The estimated
242 * formats approximate the original-space squared L2 (the rotation is
243 * norm-preserving), so this is the same quantity with the estimate
244 * noise removed (error = 0). Exact formats (FLOAT/HALF) never take
245 * this path — they are exact already, including their metric
246 * handling. */
247
4/4
✓ Branch 0 taken 4345444 times.
✓ Branch 1 taken 1298576 times.
✓ Branch 2 taken 72000 times.
✓ Branch 3 taken 126018 times.
5644020 if (fmt == PRISM_CENTROID_FMT_RABITQ || fmt == PRISM_CENTROID_FMT_FASTSCAN)
248 {
249 5217888 const float *exact = exact_internal_slots(state, page_blkno, dim);
250
251
2/2
✓ Branch 0 taken 289378 times.
✓ Branch 1 taken 4928510 times.
5217888 if (exact != NULL)
252 {
253 52000 char *content = (char *)PageGetContents(page);
254
255
3/4
✓ Branch 0 taken 1140904 times.
✓ Branch 1 taken 289378 times.
✓ Branch 2 taken 172000 times.
✗ Branch 3 not taken.
1430282 for (uint16_t i = 0; i < count && cand_count < cand_cap; i++)
256 {
257 968904 BlockNumber child;
258
259
2/2
✓ Branch 0 taken 1006904 times.
✓ Branch 1 taken 134000 times.
1140904 if (fmt == PRISM_CENTROID_FMT_FASTSCAN)
260 1006904 child = prism_centroid_fastscan_group_child(
261 content,
262 i / VS_FASTSCAN_GROUP,
263 1006904 dim)[i % VS_FASTSCAN_GROUP];
264 else
265 134000 child = prism_centroid_meta(page, i)->child_blkno;
266
267 1140904 cands[cand_count].child_blkno = child;
268 1140904 ItemPointerSet(&cands[cand_count].origin, page_blkno, i);
269 2281808 cands[cand_count].distance = vs_l2_distance_squared(
270 1140904 state->query, exact + (size_t)i * dim, dim);
271 1140904 cands[cand_count].error = 0.0f;
272 1140904 cand_count++;
273 }
274 289378 return cand_count;
275 }
276 }
277
278
4/5
✓ Branch 0 taken 1258309 times.
✓ Branch 1 taken 408683 times.
✓ Branch 2 taken 3670201 times.
✓ Branch 3 taken 17449 times.
✗ Branch 4 not taken.
5354642 switch (fmt)
279 {
280 1258309 case PRISM_CENTROID_FMT_RABITQ:
281 {
282 1258309 uint32_t data_size = VS_RABITQ_DATA_SIZE(dim);
283
284 /* Gather f_add/f_rescale in reverse order
285 * (page data grows backward: entry 0 at highest address) */
286
2/2
✓ Branch 0 taken 14809475 times.
✓ Branch 1 taken 1258309 times.
16067784 for (uint16_t i = 0; i < count; i++)
287 {
288 3287115 const RaBitQData *d =
289 14809475 prism_centroid_data(page, count - 1 - i, dim);
290 14809475 cs->f_add[i] = d->f_add;
291 14809475 cs->f_rescale[i] = d->f_rescale;
292 }
293
294 /* bits_base = last entry's bits (lowest address) */
295 1552156 const uint8_t *bits_base =
296 1258309 prism_centroid_data(page, count - 1, dim)->bits;
297
298 /* Batch distance + error bound computation. With error_scale = 0
299 * (the default) the pruning bounds are multiplied by zero anyway,
300 * so skip computing them entirely: the batch functions take a
301 * NULL lower_bounds and omit the per-entry error derivation (a
302 * divide + sqrt per centroid that also blocks vectorization of
303 * the distance-apply loop). */
304 1258309 bool want_bounds = (state->error_scale != 0.0f);
305
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 1258309 times.
1258309 Distance *lb = want_bounds ? cs->lower_bounds : NULL;
306
307
2/2
✓ Branch 0 taken 6 times.
✓ Branch 1 taken 1258303 times.
1258309 if (state->qstate->mode == VS_DISTANCE_MODE_SYMMETRIC)
308 6 vs_rabitq_distance_batch_symmetric_with_bound(
309 6 state->qstate,
310 6 cs->f_add,
311 6 cs->f_rescale,
312 bits_base,
313 data_size,
314 count,
315 dim,
316 cs->distances,
317 lb,
318 cs->symmetric_scratch);
319 else
320 1258303 vs_rabitq_distance_batch_multi_with_bound(
321 964456 state->qstate,
322 1258303 cs->f_add,
323 1258303 cs->f_rescale,
324 bits_base,
325 data_size,
326 count,
327 dim,
328 cs->distances,
329 lb,
330 cs->multi_scratch);
331
332 /* Build candidates (result j → page entry count-1-j) */
333
3/4
✓ Branch 0 taken 14809475 times.
✓ Branch 1 taken 1258309 times.
✓ Branch 2 taken 11522360 times.
✗ Branch 3 not taken.
16067784 for (uint16_t j = 0; j < count && cand_count < cand_cap; j++)
334 {
335 14809475 uint16_t page_idx = count - 1 - j;
336 3287115 const PrismCentroidEntryMeta *meta =
337 14809475 prism_centroid_meta(page, page_idx);
338
339 14809475 cands[cand_count].child_blkno = meta->child_blkno;
340 14809475 ItemPointerSet(&cands[cand_count].origin, page_blkno, page_idx);
341 14809475 cands[cand_count].distance = cs->distances[j];
342 18096590 cands[cand_count].error = want_bounds
343 ✗ ? state->error_scale *
344 ✗ (cs->distances[j] -
345 ✗ cs->lower_bounds[j])
346
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 14809475 times.
14809475 : 0.0f;
347 14809475 cand_count++;
348 }
349 964462 break;
350 }
351 408683 case PRISM_CENTROID_FMT_FLOAT:
352 {
353 /* Hoist query norm out of the inner loop: it depends only on
354 * the query, not the centroid, but was previously recomputed
355 * for every entry (one full norm² per centroid scored). */
356 817366 float norm_q = (state->metric == DISTANCE_COSINE)
357 2002 ? vs_l2_norm_squared(state->query, dim)
358
2/2
✓ Branch 0 taken 2002 times.
✓ Branch 1 taken 406681 times.
408683 : 0.0f;
359
360
3/4
✓ Branch 0 taken 3616481 times.
✓ Branch 1 taken 408683 times.
✓ Branch 2 taken 360058 times.
✗ Branch 3 not taken.
4025164 for (uint16_t i = 0; i < count && cand_count < cand_cap; i++)
361 {
362 3616481 const PrismCentroidEntryMeta *meta = prism_centroid_meta(page, i);
363 3616481 const float *fvec = prism_centroid_float_data(page, i, dim);
364 3256423 Distance dist;
365
366
3/3
✓ Branch 0 taken 145001 times.
✓ Branch 1 taken 62062 times.
✓ Branch 2 taken 3409418 times.
3616481 switch (state->metric)
367 {
368 145001 case DISTANCE_INNER_PRODUCT:
369 145001 dist = -vs_dot_product(state->query, fvec, dim);
370 145001 break;
371 62062 case DISTANCE_COSINE:
372 {
373 62062 float dot = vs_dot_product(state->query, fvec, dim);
374 62062 float norm_v = vs_l2_norm_squared(fvec, dim);
375 62062 float denom = sqrtf(norm_q * norm_v);
376
2/2
✓ Branch 0 taken 62000 times.
✓ Branch 1 taken 62 times.
62062 dist = (denom > 0.0f) ? 1.0f - dot / denom : 1.0f;
377 ✗ break;
378 }
379 3409418 default: /* L2 */
380 3409418 dist = vs_l2_distance_squared(state->query, fvec, dim);
381 3409418 break;
382 }
383
384 3616481 cands[cand_count].child_blkno = meta->child_blkno;
385 3616481 ItemPointerSet(&cands[cand_count].origin, page_blkno, i);
386 3616481 cands[cand_count].distance = dist;
387 3616481 cands[cand_count].error = 0.0f;
388 3616481 cand_count++;
389 }
390 126012 break;
391 }
392 3670201 case PRISM_CENTROID_FMT_FASTSCAN:
393 {
394 /* Fastscan centroid pages: same RaBitQ codes as the RABITQ
395 * format but rearranged into 32-vector groups so we can
396 * score them with vs_fastscan_accumulate (~150 M vec/s on
397 * Graviton 4) instead of the per-vector kernel used by the
398 * RABITQ branch above.
399 *
400 * Layout per group section:
401 * BlockNumber child_blkno[32]
402 * float f_add[32]
403 * float f_rescale[32]
404 * float f_error[32]
405 * uint8_t codes[nsq_pairs*32]
406 *
407 * The LUT depends on the query (and via qstate->transformed,
408 * implicitly on the cluster's centroid for posting scans;
409 * for centroid descent the relevant query state is the
410 * pre-cluster qstate->transformed which is just P^T*query
411 * minus the global_mean rotation already absorbed by
412 * pt_global_mean). It is rebuilt once per centroid page
413 * scored. Amortising the LUT build across all 32 entries in
414 * a group is the whole reason this is faster than the
415 * per-vector kernel. */
416 /* Build the LUT once per query — qstate->transformed is
417 * constant across the whole centroid descent, so we cache
418 * the LUT bytes + lut_delta + lut_bias in PrismCentroidScratch
419 * and reuse them on every subsequent FASTSCAN page. */
420
2/2
✓ Branch 0 taken 385255 times.
✓ Branch 1 taken 3284946 times.
3670201 if (!cs->fs_lut_valid)
421 {
422 385255 uint64_t t_lut = cs_now_ns();
423 385255 vs_fastscan_build_lut_hacc(
424 385255 state->qstate->transformed,
425 dim,
426 cs->fs_lut,
427 &cs->fs_lut_delta,
428 &cs->fs_lut_bias);
429 385255 cs->fs_lut_valid = true;
430 385255 cs->t_lut_ns += cs_now_ns() - t_lut;
431 }
432 3670201 float lut_delta = cs->fs_lut_delta;
433 3670201 float lut_bias = cs->fs_lut_bias;
434
435 3670201 float g_add = state->qstate->g_add;
436 3670201 float sum_t = state->qstate->sum_transformed;
437 3670201 float inv_sqrt_d = state->qstate->inv_sqrt_d;
438 3670201 float g_error = state->qstate->g_error;
439 3670201 float err_mult = state->qstate->error_multiplier;
440
441 3670201 char *content = (char *)PageGetContents(page);
442 3670201 uint32_t entry_count = count;
443 3670201 uint32_t ngroups = (entry_count + VS_FASTSCAN_GROUP - 1) /
444 VS_FASTSCAN_GROUP;
445
446 3616201 int32_t accum[VS_FASTSCAN_GROUP];
447
448
2/2
✓ Branch 0 taken 3670201 times.
✓ Branch 1 taken 3670201 times.
7340402 for (uint32_t g = 0; g < ngroups; g++)
449 {
450 3670201 uint32_t g_start = g * VS_FASTSCAN_GROUP;
451 3670201 uint32_t g_count = entry_count - g_start;
452
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 54000 times.
3670201 if (g_count > VS_FASTSCAN_GROUP)
453 ✗ g_count = VS_FASTSCAN_GROUP;
454
455 3616201 const BlockNumber *child = (const BlockNumber *)
456 3670201 prism_centroid_fastscan_group_child(content, g, dim);
457 3616201 const float *f_add_arr =
458 3670201 prism_centroid_fastscan_group_f_add(content, g, dim);
459 3616201 const float *f_rescale_arr =
460 3670201 prism_centroid_fastscan_group_f_rescale(content, g, dim);
461 3616201 const float *f_error_arr =
462 3670201 prism_centroid_fastscan_group_f_error(content, g, dim);
463 3616201 const uint8_t *codes =
464 3670201 prism_centroid_fastscan_group_codes(content, g, dim);
465
466 3670201 vs_fastscan_accumulate_hacc(codes, cs->fs_lut, accum, dim);
467
468
3/4
✓ Branch 0 taken 162000 times.
✓ Branch 1 taken 34295325 times.
✓ Branch 2 taken 3778201 times.
✗ Branch 3 not taken.
38073526 for (uint32_t v = 0; v < g_count && cand_count < cand_cap; v++)
469 {
470 /* De-quantise the LUT accumulator the same way the
471 * posting fastscan path does (see prune_group_neon
472 * in posting_scan.c). */
473 34403325 float binary_ip = (float)accum[v] * lut_delta + lut_bias;
474 34403325 float final_dot = (2.0f * binary_ip - sum_t) * inv_sqrt_d;
475
476 34403325 Distance est = f_add_arr[v] + g_add -
477 34403325 2.0f * f_rescale_arr[v] * final_dot;
478 /* Matches rabitq_lower_bound(): err_margin =
479 * multiplier * f_error * g_error, plus a small
480 * floating-point margin proportional to |est|.
481 * error_scale = 0 (the default) zeroes the margin, so
482 * skip the arithmetic in that case. */
483 34403325 Distance err = 0.0f;
484
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 34403325 times.
34403325 if (state->error_scale != 0.0f)
485 ✗ err = state->error_scale *
486 ✗ (err_mult * f_error_arr[v] * g_error +
487 ✗ 1e-5f * fabsf(est));
488
489 34403325 uint32_t page_idx = g_start + v;
490 34403325 cands[cand_count].child_blkno = child[v];
491 34403325 ItemPointerSet(
492 162000 &cands[cand_count].origin, page_blkno, page_idx);
493 34403325 cands[cand_count].distance = est;
494 34403325 cands[cand_count].error = err;
495 34403325 cand_count++;
496 }
497 }
498 3670201 break;
499 }
500 17449 case PRISM_CENTROID_FMT_HALF:
501 {
502 34898 float norm_q = (state->metric == DISTANCE_COSINE)
503 ✗ ? vs_l2_norm_squared(state->query, dim)
504
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 17449 times.
17449 : 0.0f;
505
506
3/4
✓ Branch 0 taken 140135 times.
✓ Branch 1 taken 17449 times.
✓ Branch 2 taken 40 times.
✗ Branch 3 not taken.
157584 for (uint16_t i = 0; i < count && cand_count < cand_cap; i++)
507 {
508 140135 const PrismCentroidEntryMeta *meta = prism_centroid_meta(page, i);
509 140135 const half *hvec = prism_centroid_half_data(page, i, dim);
510 140095 Distance dist;
511
512
2/3
✓ Branch 0 taken 1001 times.
✗ Branch 1 not taken.
✓ Branch 2 taken 139134 times.
140135 switch (state->metric)
513 {
514 1001 case DISTANCE_INNER_PRODUCT:
515 1001 dist = -vs_f16_dot_product(hvec, state->query, dim);
516 1001 break;
517 ✗ case DISTANCE_COSINE:
518 {
519 ✗ float dot = vs_f16_dot_product(hvec, state->query, dim);
520 ✗ float norm_v = vs_f16_norm_sq(hvec, dim);
521 ✗ float denom = sqrtf(norm_q * norm_v);
522 ✗ dist = (denom > 0.0f) ? 1.0f - dot / denom : 1.0f;
523 ✗ break;
524 }
525 139134 default: /* L2 */
526 139134 dist = vs_f16_l2_squared(hvec, state->query, dim);
527 139134 break;
528 }
529
530 140135 cands[cand_count].child_blkno = meta->child_blkno;
531 140135 ItemPointerSet(&cands[cand_count].origin, page_blkno, i);
532 140135 cands[cand_count].distance = dist;
533 140135 cands[cand_count].error = 0.0f;
534 140135 cand_count++;
535 }
536 6 break;
537 }
538 }
539
540 1144480 return cand_count;
541 }
542
543 /* ----------------------------------------------------------------
544 * Top-K selection via VsTopK (error-bound-aware)
545 *
546 * Selects the best candidates using VsTopK pruning. Candidates
547 * with lower_bound >= threshold are pruned. With overlapping
548 * error intervals, may return more than k entries.
549 *
550 * Results are placed sorted in out[0..return_count).
551 * out must not alias cands.
552 * ---------------------------------------------------------------- */
553 static uint32_t
554 895617 select_topk_bounded(
555 PrismCentroidScratch *scratch,
556 Candidate *cands,
557 uint32_t count,
558 uint32_t k,
559 Candidate *out,
560 uint32_t out_cap)
561 {
562
2/2
✓ Branch 0 taken 692497 times.
✓ Branch 1 taken 203120 times.
895617 if (count == 0)
563 ✗ return 0;
564
565 /* Reuse the per-scan topk and extract buffer instead of allocating
566 * new ones every level; vs_topk_reset_to_k is O(1) unless k grows
567 * past the allocated capacity. */
568 895617 VsTopK *topk = &scratch->level_topk;
569 895617 vs_topk_reset_to_k(topk, k);
570
571 /* Centroid candidates have unique ids (the buf index), so we can
572 * skip the O(k) per-insert dedup scan. */
573
3/3
✓ Branch 0 taken 12216458 times.
✓ Branch 1 taken 42096982 times.
✓ Branch 2 taken 692497 times.
55005937 for (uint32_t i = 0; i < count; i++)
574 54110320 vs_topk_insert_unique(topk, cands[i].distance, cands[i].error, i);
575
576 /* entries_buf must hold topk->cand_count survivors; grow if needed.
577 * Grow in the scratch's owning context: this runs under whatever
578 * context the route was issued from — during index builds a per-row
579 * scratch context that is reset after every tuple — while the buffer
580 * must survive for the scratch's whole lifetime. */
581
2/2
✓ Branch 0 taken 7 times.
✓ Branch 1 taken 895610 times.
895617 if (topk->cand_count > scratch->entries_cap)
582 {
583 7 uint32_t new_cap = scratch->entries_cap * 2;
584
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 7 times.
7 while (new_cap < topk->cand_count)
585 ✗ new_cap *= 2;
586 7 vs_free(scratch->entries_buf);
587 7 scratch->entries_buf = vs_memctx_alloc(
588 scratch->memctx, new_cap * sizeof(VsTopKEntry));
589 7 scratch->entries_cap = new_cap;
590 }
591
592 692497 uint32_t nresults;
593 895617 vs_topk_extract_sorted_unique(topk, scratch->entries_buf, &nresults);
594
595
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 895617 times.
895617 if (nresults > out_cap)
596 ✗ nresults = out_cap;
597
598
2/2
✓ Branch 0 taken 15434044 times.
✓ Branch 1 taken 895617 times.
16329661 for (uint32_t i = 0; i < nresults; i++)
599 {
600 15434044 uint32_t idx = (uint32_t)scratch->entries_buf[i].id;
601 15434044 out[i] = cands[idx];
602 15434044 out[i].distance = scratch->entries_buf[i].distance;
603 15434044 out[i].error = scratch->entries_buf[i].error;
604 }
605
606 203120 return nresults;
607 }
608
609 /*
610 * Emit all of a centroid page's children as candidates without scoring.
611 *
612 * Used when a level's selection cannot reject anything (keep >= entry
613 * count): the distances would exist only to rank candidates for a
614 * selection that keeps them all, and the next level re-scores its own
615 * children from scratch, so computing them is pure waste. Candidates
616 * are emitted with zero distance/error.
617 */
618 static uint32_t
619 372428 emit_page_children(
620 Page page,
621 BlockNumber page_blkno,
622 Dimension dim,
623 Candidate *cands,
624 uint32_t cand_count,
625 uint32_t cand_cap)
626 {
627 372428 PrismCentroidPageOpaque *opaque = PRISM_CENTROID_OPAQUE(page);
628 372428 uint16_t count = opaque->entry_count;
629 372428 PrismCentroidFormat fmt = prism_centroid_page_format(page);
630
631
2/2
✓ Branch 0 taken 227206 times.
✓ Branch 1 taken 145222 times.
372428 if (fmt == PRISM_CENTROID_FMT_FASTSCAN)
632 {
633 227206 char *content = (char *)PageGetContents(page);
634 227206 uint32_t ngroups = (count + VS_FASTSCAN_GROUP - 1) / VS_FASTSCAN_GROUP;
635
636
2/2
✓ Branch 0 taken 227206 times.
✓ Branch 1 taken 227206 times.
454412 for (uint32_t g = 0; g < ngroups; g++)
637 {
638 227206 uint32_t g_start = g * VS_FASTSCAN_GROUP;
639 227206 uint32_t g_count = count - g_start;
640
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 6000 times.
227206 if (g_count > VS_FASTSCAN_GROUP)
641 ✗ g_count = VS_FASTSCAN_GROUP;
642
643 221206 const BlockNumber *child = (const BlockNumber *)
644 227206 prism_centroid_fastscan_group_child(content, g, dim);
645
646
3/4
✓ Branch 0 taken 2083088 times.
✓ Branch 1 taken 227206 times.
✓ Branch 2 taken 18000 times.
✗ Branch 3 not taken.
2310294 for (uint32_t v = 0; v < g_count && cand_count < cand_cap; v++)
647 {
648 2083088 cands[cand_count].child_blkno = child[v];
649 2083088 ItemPointerSet(
650 2083088 &cands[cand_count].origin, page_blkno, g_start + v);
651 2083088 cands[cand_count].distance = 0.0f;
652 2083088 cands[cand_count].error = 0.0f;
653 2083088 cand_count++;
654 }
655 }
656 227206 return cand_count;
657 }
658
659
3/4
✓ Branch 0 taken 1379163 times.
✓ Branch 1 taken 145222 times.
✓ Branch 2 taken 835530 times.
✗ Branch 3 not taken.
1524385 for (uint16_t i = 0; i < count && cand_count < cand_cap; i++)
660 {
661 1379163 const PrismCentroidEntryMeta *meta = prism_centroid_meta(page, i);
662
663 1379163 cands[cand_count].child_blkno = meta->child_blkno;
664 1379163 ItemPointerSet(&cands[cand_count].origin, page_blkno, i);
665 1379163 cands[cand_count].distance = 0.0f;
666 1379163 cands[cand_count].error = 0.0f;
667 1379163 cand_count++;
668 }
669 95474 return cand_count;
670 }
671
672 uint32_t
673 624913 prism_centroid_beam_search(
674 const PrismCentroidSearchState *state,
675 BlockNumber first_centroid_blkno,
676 uint8_t nlevels,
677 PrismCentroidResult *results,
678 float *centroid_vecs,
679 PrismCentroidSearchStats *stats)
680 {
681
8/8
✓ Branch 0 taken 624911 times.
✓ Branch 1 taken 2 times.
✓ Branch 2 taken 174100 times.
✓ Branch 3 taken 2 times.
✓ Branch 4 taken 174098 times.
✓ Branch 5 taken 2 times.
✓ Branch 6 taken 2 times.
✓ Branch 7 taken 174096 times.
624913 if (state == NULL || results == NULL || nlevels == 0 ||
682
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 450809 times.
450809 first_centroid_blkno == InvalidBlockNumber)
683 8 return 0;
684
685 624905 Dimension dim = state->dim;
686 624905 uint32_t beam_width = state->beam_width;
687 624905 uint32_t nprobe = state->nprobe;
688
689 /* Caller-owned scratch: pre-allocated buffers for candidate
690 * arrays and per-page batch scoring. Avoids 8 palloc + 2 memctx
691 * creations per query — a meaningful chunk of the allocator
692 * traffic on the hot path.
693 *
694 * If state->scratch is NULL we fall back to a one-shot
695 * allocation. The fallback exists for tests and ad-hoc callers;
696 * production query paths (PrismQueryState / PrismQueryCtx)
697 * pre-allocate and pass it in. */
698 624905 PrismCentroidScratch *scratch = state->scratch;
699 624905 PrismCentroidScratch *owned_scratch = NULL;
700
2/2
✓ Branch 0 taken 6 times.
✓ Branch 1 taken 624899 times.
624905 if (scratch == NULL)
701 {
702 6 owned_scratch = prism_centroid_scratch_create(dim, beam_width);
703
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 6 times.
6 if (owned_scratch == NULL)
704 ✗ return 0;
705 6 scratch = owned_scratch;
706 }
707 624905 uint32_t cand_cap = scratch->cand_cap;
708 624905 Candidate *buf_a = scratch->buf_a;
709 624905 Candidate *buf_b = scratch->buf_b;
710
711 /* Invalidate the per-query fastscan LUT cache. Built lazily on
712 * first FASTSCAN page encountered, then reused for all subsequent
713 * pages in this query. */
714 624905 scratch->fs_lut_valid = false;
715
716 /* Reset fine-grained timing accumulators for this query. */
717 624905 scratch->t_lut_ns = 0;
718 624905 scratch->t_pageread_ns = 0;
719 624905 scratch->t_score_ns = 0;
720
721 /* centroid_vecs: will be used later for copying centroid vectors */
722 450809 (void)centroid_vecs;
723
724 /*
725 * buf_a accumulates raw candidates from score_page.
726 * buf_b receives the topk-selected survivors.
727 * After selection, buf_b becomes the live set for expansion.
728 */
729
730 624905 uint32_t centroid_pages_read = 0;
731
732 /* beam_width is the intermediate-level keep; the leaf level always
733 * returns nprobe (see keep below), and beam_width*fan_out >= nprobe
734 * covers the top-nprobe leaves, so beam_width may be < nprobe. */
735
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 174096 times.
624905 if (beam_width < 1)
736 ✗ beam_width = 1;
737
738
2/2
✓ Branch 0 taken 424572 times.
✓ Branch 1 taken 200333 times.
624905 uint32_t level0_keep = (nlevels == 1) ? nprobe : beam_width;
739
740 /* Level 0: read root centroid page(s), score ALL centroids.
741 *
742 * Fast path: when the root is a single page whose entry count does
743 * not exceed the level-0 keep, selection cannot reject anything —
744 * scoring would rank candidates for a no-op selection, and level 1
745 * re-scores its own children anyway. Emit the children directly
746 * and skip both the scoring and the top-K pass. */
747 624905 uint32_t raw_count = 0;
748 624905 bool level0_full = false;
749 624905 BlockNumber blkno = first_centroid_blkno;
750
751
2/2
✓ Branch 0 taken 625411 times.
✓ Branch 1 taken 624905 times.
1250316 while (blkno != InvalidBlockNumber)
752 {
753 625411 uint64_t t_r = cs_now_ns();
754 625411 Page page = vs_storage_read_page(state->storage, blkno);
755 625411 scratch->t_pageread_ns += cs_now_ns() - t_r;
756 625411 PrismCentroidPageOpaque *opaque = PRISM_CENTROID_OPAQUE(page);
757 625411 centroid_pages_read++;
758
759
4/4
✓ Branch 0 taken 624905 times.
✓ Branch 1 taken 506 times.
✓ Branch 2 taken 624403 times.
✓ Branch 3 taken 502 times.
625411 if (raw_count == 0 && opaque->next_blkno == InvalidBlockNumber &&
760
4/4
✓ Branch 0 taken 540342 times.
✓ Branch 1 taken 84061 times.
✓ Branch 2 taken 372428 times.
✓ Branch 3 taken 167914 times.
624403 opaque->entry_count <= level0_keep && nlevels > 1)
761 {
762 270954 raw_count =
763 372428 emit_page_children(page, blkno, dim, buf_a, 0, cand_cap);
764 372428 level0_full = true;
765 }
766 else
767 {
768 252983 uint64_t t_s = cs_now_ns();
769 252983 raw_count = score_page(
770 state,
771 page,
772 blkno,
773 dim,
774 buf_a,
775 raw_count,
776 cand_cap,
777 scratch);
778 252983 scratch->t_score_ns += cs_now_ns() - t_s;
779 }
780
781 625411 BlockNumber next_blkno = opaque->next_blkno;
782 625411 t_r = cs_now_ns();
783 625411 vs_storage_release_page(state->storage, blkno);
784 625411 scratch->t_pageread_ns += cs_now_ns() - t_r;
785 625411 blkno = next_blkno;
786 }
787
4/4
✓ Branch 0 taken 353945 times.
✓ Branch 1 taken 270960 times.
✓ Branch 2 taken 72618 times.
✓ Branch 3 taken 101472 times.
624905 if (stats && !level0_full)
788 252473 stats->dist_calcs += raw_count;
789
790 /* Select top-K from level 0 into buf_b (skipped when the fast path
791 * already kept everything).
792 *
793 * Error-bound-aware selection via VsTopK: keeps the beam_width
794 * candidates with smallest upper bounds, plus any additional
795 * candidates whose lower bound overlaps the threshold. For
796 * exact formats (error=0) this returns exactly beam_width. */
797 450809 Candidate *live;
798 450809 Candidate *expand_buf;
799 450809 uint32_t cand_count;
800
801
2/2
✓ Branch 0 taken 281329 times.
✓ Branch 1 taken 343576 times.
624905 if (level0_full)
802 {
803 101474 live = buf_a;
804 101474 expand_buf = buf_b;
805 101474 cand_count = raw_count;
806 }
807 else
808 {
809 252477 cand_count = select_topk_bounded(
810 scratch, buf_a, raw_count, level0_keep, buf_b, cand_cap);
811 252477 live = buf_b;
812 252477 expand_buf = buf_a;
813 }
814
815 /* Intermediate levels: expand winners via child_blkno */
816
2/2
✓ Branch 0 taken 643140 times.
✓ Branch 1 taken 624905 times.
1268045 for (uint8_t level = 1; level < nlevels; level++)
817 {
818 130498 uint32_t next_count = 0;
819
820
2/2
✓ Branch 0 taken 5391037 times.
✓ Branch 1 taken 643140 times.
6034177 for (uint32_t i = 0; i < cand_count; i++)
821 {
822 5391037 BlockNumber child_blkno = live[i].child_blkno;
823
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 5391037 times.
5391037 if (child_blkno == InvalidBlockNumber)
824 ✗ continue;
825
826 1123858 BlockNumber cb = child_blkno;
827
2/2
✓ Branch 0 taken 5391037 times.
✓ Branch 1 taken 5391037 times.
10782074 while (cb != InvalidBlockNumber)
828 {
829 5391037 uint64_t t_r = cs_now_ns();
830 5391037 Page page = vs_storage_read_page(state->storage, cb);
831 5391037 scratch->t_pageread_ns += cs_now_ns() - t_r;
832 5391037 centroid_pages_read++;
833 5391037 uint64_t t_s = cs_now_ns();
834 5391037 next_count = score_page(
835 state,
836 page,
837 cb,
838 dim,
839 expand_buf,
840 next_count,
841 cand_cap,
842 scratch);
843 5391037 scratch->t_score_ns += cs_now_ns() - t_s;
844 5391037 PrismCentroidPageOpaque *opaque = PRISM_CENTROID_OPAQUE(page);
845 5391037 BlockNumber nb = opaque->next_blkno;
846 5391037 t_r = cs_now_ns();
847 5391037 vs_storage_release_page(state->storage, cb);
848 5391037 scratch->t_pageread_ns += cs_now_ns() - t_r;
849 5391037 cb = nb;
850 }
851 }
852
2/2
✓ Branch 0 taken 643134 times.
✓ Branch 1 taken 6 times.
643140 if (stats)
853 643134 stats->dist_calcs += next_count;
854
855 /* Select winners into live; expand_buf is the raw input. */
856
2/2
✓ Branch 0 taken 262166 times.
✓ Branch 1 taken 380974 times.
643140 uint32_t keep = (level == nlevels - 1) ? nprobe : beam_width;
857 643140 cand_count = select_topk_bounded(
858 scratch, expand_buf, next_count, keep, live, cand_cap);
859 }
860
861 /* Build results in caller-owned memory (cap at nprobe) */
862 624905 uint32_t result_count = cand_count < nprobe ? cand_count : nprobe;
863
2/2
✓ Branch 0 taken 13505249 times.
✓ Branch 1 taken 624905 times.
14130154 for (uint32_t i = 0; i < result_count; i++)
864 {
865 13505249 results[i].posting_head = live[i].child_blkno;
866 13505249 results[i].distance = live[i].distance;
867 13505249 results[i].error = live[i].error;
868 }
869
870 /* Extract centroid vectors for winning clusters.
871 * Only meaningful for FLOAT/HALF centroid pages — RaBitQ is a
872 * lossy binary encoding, so the full-precision centroid cannot
873 * be recovered. RaBitQ callers must obtain pt_centroid from its
874 * dedicated location (e.g., the first posting page). Skip the
875 * loop entirely rather than re-reading pages for nothing. */
876
1/6
✗ Branch 0 not taken.
✓ Branch 1 taken 624905 times.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✗ Branch 4 not taken.
✗ Branch 5 not taken.
624905 if (centroid_vecs != NULL && result_count > 0 && state->qstate == NULL)
877 {
878 ✗ BlockNumber prev_blk = InvalidBlockNumber;
879 ✗ Page prev_page = NULL;
880
881 ✗ for (uint32_t i = 0; i < result_count; i++)
882 {
883 ✗ BlockNumber blk = ItemPointerGetBlockNumber(&live[i].origin);
884 ✗ OffsetNumber idx = ItemPointerGetOffsetNumber(&live[i].origin);
885
886 ✗ Page page;
887 ✗ if (blk == prev_blk)
888 {
889 ✗ page = prev_page;
890 }
891 else
892 {
893 ✗ if (prev_page != NULL)
894 ✗ vs_storage_release_page(state->storage, prev_blk);
895 ✗ page = vs_storage_read_page(state->storage, blk);
896 ✗ centroid_pages_read++;
897 ✗ prev_blk = blk;
898 ✗ prev_page = page;
899 }
900
901 ✗ float *dst = centroid_vecs + (size_t)i * dim;
902 ✗ PrismCentroidFormat fmt = prism_centroid_page_format(page);
903
904 ✗ switch (fmt)
905 {
906 ✗ case PRISM_CENTROID_FMT_FLOAT:
907 {
908 ✗ const float *src = prism_centroid_float_data(page, idx, dim);
909 ✗ memcpy(dst, src, dim * sizeof(float));
910 ✗ break;
911 }
912 ✗ case PRISM_CENTROID_FMT_HALF:
913 {
914 ✗ const half *src = prism_centroid_half_data(page, idx, dim);
915 ✗ for (Dimension d = 0; d < dim; d++)
916 ✗ dst[d] = vs_half_to_float(src[d]);
917 ✗ break;
918 }
919 ✗ case PRISM_CENTROID_FMT_RABITQ:
920 case PRISM_CENTROID_FMT_FASTSCAN:
921 /* Unreachable — both are lossy binary encodings and
922 * the outer guard skips this branch when the page
923 * format isn't FLOAT/HALF. */
924 ✗ break;
925 }
926 }
927
928 ✗ if (prev_page != NULL)
929 ✗ vs_storage_release_page(state->storage, prev_blk);
930 }
931
932
2/2
✓ Branch 0 taken 624899 times.
✓ Branch 1 taken 6 times.
624905 if (stats)
933 {
934 624899 stats->pages_read = centroid_pages_read;
935 624899 stats->lut_ns = scratch->t_lut_ns;
936 624899 stats->pageread_ns = scratch->t_pageread_ns;
937 624899 stats->score_ns = scratch->t_score_ns;
938 }
939
940
2/2
✓ Branch 0 taken 6 times.
✓ Branch 1 taken 624899 times.
624905 if (owned_scratch != NULL)
941 6 prism_centroid_scratch_free(owned_scratch);
942
943 174096 return result_count;
944 }
945