GCC Code Coverage Report


Directory: src/
File: src/index/query_scan.c
Date: 2026-09-30 11:11:31
Exec Total Coverage
Lines: 262 277 94.6%
Functions: 16 16 100.0%
Branches: 100 125 80.0%

Line Branch Exec Source
1 /*
2 * Copyright (c) 2026 Tiger Data, Inc.
3 * Licensed under the PostgreSQL License. See LICENSE for details.
4 *
5 * query_scan.c - Shared query execution for ANN search
6 *
7 * Beam search over centroids, posting list scan with RaBitQ
8 * scoring, and candidate extraction. All buffers are pre-allocated
9 * in init; the per-query path does zero allocations (except rare
10 * topk candidate buffer growth).
11 */
12
13 #include "vs_config.h"
14
15 #include <float.h>
16 #include <math.h>
17 #include <stdlib.h>
18 #include <string.h>
19 #include <time.h>
20
21 #include "algo/topk.h"
22 #include "algo/vecops.h"
23 #include "core/injection.h"
24 #include "core/log.h"
25 #include "core/memory.h"
26 #include "core/platform.h"
27 #include "index/centroid_search.h"
28 #include "index/posting_page.h"
29 #include "index/posting_scan.h"
30 #include "index/query_scan.h"
31 #include "quant/rabitq.h"
32
33 /* ----------------------------------------------------------------
34 * Init / cleanup
35 * ---------------------------------------------------------------- */
36
37 void
38 18105 prism_query_state_init(
39 PrismQueryState *qs,
40 PrismIndexBase *index,
41 uint32_t max_k,
42 uint32_t max_nprobe)
43 {
44 18105 memset(qs, 0, sizeof(*qs));
45 18105 qs->index = index;
46 18105 qs->max_k = max_k;
47 18105 qs->max_nprobe = max_nprobe;
48
49 18105 prism_index_ensure_rabitq(index);
50
51 /*
52 * Everything below is allocated here, including the contexts the top-k
53 * and the centroid scratch create for themselves, so that the whole
54 * state can be released at once. Parented to the caller's current
55 * context, so a caller that never calls cleanup still loses it with
56 * whatever scope it allocated the state in.
57 */
58 18105 qs->memctx = vs_memctx_create(NULL, "vs query state");
59 18105 VsMemCtx caller_ctx = vs_memctx_switch(qs->memctx);
60
61 18105 Dimension dim = index->dim;
62 18105 uint32_t packed_bytes = VS_RABITQ_BYTES(dim);
63
64 /* Query buffers */
65 18105 qs->query_buf = vs_alloc(dim * sizeof(float));
66 18105 qs->pt_query = vs_alloc_aligned(dim * sizeof(float), 64);
67 18105 qs->beam_transformed = vs_alloc_aligned(dim * sizeof(float), 64);
68 18105 qs->cluster_transformed = vs_alloc_aligned(dim * sizeof(float), 64);
69 18105 qs->beam_query_bits = vs_alloc_aligned(packed_bytes, 64);
70 18105 qs->cluster_query_bits = vs_alloc_aligned(packed_bytes, 64);
71
72 /* Wire up query state buffers */
73 18105 qs->beam_qs.transformed = qs->beam_transformed;
74 18105 qs->beam_qs.query_bits = qs->beam_query_bits;
75 18105 qs->cluster_qs.transformed = qs->cluster_transformed;
76 18105 qs->cluster_qs.query_bits = qs->cluster_query_bits;
77
78 18105 vs_rabitq_init_query_constants(&qs->beam_qs, dim);
79 18105 vs_rabitq_init_query_constants(&qs->cluster_qs, dim);
80
81 /* Beam search results + per-scan scratch.
82 * Pre-allocating the scratch here means prism_centroid_beam_search
83 * skips 8 vs_alloc calls and 2 memory-context creations on
84 * every query (the largest remaining source of per-query
85 * allocator traffic after the dedup-gens fix). Sized to the
86 * worst case beam_width == max_nprobe. */
87 18105 qs->beam_results = vs_alloc(max_nprobe * sizeof(PrismCentroidResult));
88 18105 qs->centroid_scratch = prism_centroid_scratch_create(dim, max_nprobe);
89
90 /* Probe-order scratch (exact centroid re-rank of the expanded
91 * probe set; see prism_query_set_probe_expand). */
92 18105 qs->probe_dists = vs_alloc(max_nprobe * sizeof(float));
93 18105 qs->probe_order = vs_alloc(max_nprobe * sizeof(uint32_t));
94
95 /* Top-K */
96 18105 vs_topk_init(&qs->topk, max_k);
97
98 /* Candidate extraction buffer */
99 18105 qs->cand_cap = max_k * PRISM_QUERY_CAND_PER_K;
100 18105 qs->candidates = vs_alloc(qs->cand_cap * sizeof(VsTopKEntry));
101
102 /* Result ordering */
103 18105 qs->result_order = vs_alloc(qs->cand_cap * sizeof(uint32_t));
104 18105 qs->result_dists = vs_alloc(qs->cand_cap * sizeof(Distance));
105
106 /* Posting scan iterator */
107 18105 uint32_t max_entries = prism_posting_max_entries(dim);
108 18105 prism_posting_scan_init(
109 &qs->pscan,
110 index->posting_storage,
111 index->page_base,
112 18105 index->params,
113 dim,
114 max_entries);
115
116 18105 vs_memctx_switch(caller_ctx);
117 18105 }
118
119 void
120 18103 prism_query_state_cleanup(PrismQueryState *qs)
121 {
122
2/2
✓ Branch 0 taken 17967 times.
✓ Branch 1 taken 136 times.
18103 if (qs == NULL)
123 ✗ return;
124
125 /*
126 * The pinned page first: it is the one resource the state holds that is
127 * not memory, and deleting the context below would strand it.
128 */
129 18103 prism_posting_scan_cleanup(&qs->pscan);
130
131 /* Sub-contexts before the context that parents them. */
132 18103 vs_topk_cleanup(&qs->topk);
133 18103 prism_centroid_scratch_free(qs->centroid_scratch);
134 18103 qs->centroid_scratch = NULL;
135
136
1/2
✓ Branch 0 taken 18103 times.
✗ Branch 1 not taken.
18103 if (qs->memctx != NULL)
137 {
138 18103 vs_memctx_delete(qs->memctx);
139 18103 qs->memctx = NULL;
140 }
141 }
142
143 /* ----------------------------------------------------------------
144 * Per-query execution
145 * ---------------------------------------------------------------- */
146
147 static const float *
148 624875 prepare_query(PrismQueryState *qs, const float *query)
149 {
150
2/2
✓ Branch 0 taken 184093 times.
✓ Branch 1 taken 440782 times.
624875 if (qs->index->metric != DISTANCE_COSINE)
151 172184 return query;
152
153 13791 Dimension dim = qs->index->dim;
154 13791 memcpy(qs->query_buf, query, dim * sizeof(float));
155 13791 float norm = vs_l2_norm(qs->query_buf, dim);
156
2/2
✓ Branch 0 taken 13665 times.
✓ Branch 1 taken 126 times.
13791 if (norm > 0.0f)
157 13665 vec32_scale(qs->query_buf, 1.0f / norm, qs->query_buf, dim);
158 13791 return qs->query_buf;
159 }
160
161 static uint32_t
162 624875 search_centroids(
163 PrismQueryState *qs,
164 const float *qvec,
165 uint32_t nprobe,
166 VsDistanceMode mode,
167 PrismCentroidSearchStats *beam_stats)
168 {
169 624875 const PrismIndexBase *idx = qs->index;
170 624875 Dimension dim = idx->dim;
171
172 624875 RaBitQQueryState *rqs = NULL;
173
2/2
✓ Branch 0 taken 434286 times.
✓ Branch 1 taken 190589 times.
624875 if (idx->centroid_format == PRISM_CENTROID_FMT_RABITQ ||
174
2/2
✓ Branch 0 taken 6000 times.
✓ Branch 1 taken 12000 times.
18000 idx->centroid_format == PRISM_CENTROID_FMT_FASTSCAN)
175 {
176 578352 vs_rabitq_init_query_state(
177 578352 &qs->beam_qs, qs->pt_query, idx->pt_global_mean, dim, mode);
178 578352 rqs = &qs->beam_qs;
179 }
180
181 1075684 uint32_t beam_w = prism_query_beam_width(
182 624875 nprobe, idx->nlist, idx->fan_out, idx->centroid_beam_scale);
183
184 624875 PrismCentroidSearchState search = {
185 .qstate = rqs,
186 .query = qvec,
187 624875 .storage = idx->centroid_storage,
188 .beam_width = beam_w,
189 .nprobe = nprobe,
190 .dim = dim,
191 624875 .metric = idx->metric,
192 624875 .error_scale = idx->centroid_error_scale,
193 624875 .scratch = qs->centroid_scratch,
194 /* NULL on every query path; only the build route sets it. */
195 624875 .exact_internal = idx->exact_internal,
196 };
197
198 1249750 return prism_centroid_beam_search(
199 &search,
200 624875 idx->first_centroid,
201 624875 idx->nlevels,
202 qs->beam_results,
203 NULL,
204 beam_stats);
205 }
206
207 /* Probe-order refinement factor (prism.probe_expand); see query_scan.h.
208 * Enabled by default: expansion gains saturate around a factor of 2,
209 * so 2.0 captures ~all the recall benefit of exact probe ordering.
210 * 1.0 means no expansion (identity). */
211 static double g_probe_expand = 2.0;
212
213 /* Cap on extra routed candidates. At large nprobe a deep scan already
214 * covers cluster membership, so ordering refinement adds little while
215 * the phase-A cost keeps growing linearly; capping the expansion keeps
216 * the overhead bounded (measured to retain nearly all of the recall
217 * gain at high nprobe). */
218 #define PRISM_PROBE_EXPAND_MAX_EXTRA 256
219
220 void
221 257 prism_query_set_probe_expand(double expand)
222 {
223 257 g_probe_expand = expand;
224 257 }
225
226 /*
227 * Centroid slots the beam keeps at each intermediate level.
228 *
229 * Starts at a fraction of nprobe and is then raised by three floors, each
230 * of which exists because dropping below it loses leaves outright rather
231 * than merely ranking them lower:
232 *
233 * - Routing floor. Below PRISM_CENTROID_BEAM_FLOOR a scaled-down beam saves
234 * almost nothing and mis-routes, so the beam covers every probed list.
235 *
236 * - Coverage floor. A kept set of beam_w parents exposes at most
237 * beam_w * fan_out leaves, so returning nprobe of them needs
238 * ceil(nprobe / fan_out) parents.
239 *
240 * - Probe-everything floor. The coverage floor assumes every child
241 * carries fan_out leaves; an unbalanced tree holds fewer, so once
242 * nprobe covers every leaf the beam keeps whole levels.
243 *
244 * Called by prism_query_execute and by the cost model, which prices a page
245 * read per kept slot per level.
246 */
247 uint32_t
248 625257 prism_query_beam_width(
249 uint32_t nprobe, uint32_t nlist, uint32_t fan_out, double beam_scale)
250 {
251 625257 uint32_t beam_w = (uint32_t)((double)nprobe * beam_scale);
252
253
2/2
✓ Branch 0 taken 430 times.
✓ Branch 1 taken 173636 times.
625257 if (beam_w < 1)
254 430 beam_w = 1;
255
256 625257 uint32_t floor_w = nprobe < PRISM_CENTROID_BEAM_FLOOR
257 ? nprobe
258 : PRISM_CENTROID_BEAM_FLOOR;
259
260
2/2
✓ Branch 0 taken 3274 times.
✓ Branch 1 taken 170792 times.
625257 if (beam_w < floor_w)
261 3274 beam_w = floor_w;
262
263
1/2
✓ Branch 0 taken 625257 times.
✗ Branch 1 not taken.
625257 if (fan_out > 0)
264 {
265 625257 floor_w = (nprobe + fan_out - 1) / fan_out;
266
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 174066 times.
625257 if (beam_w < floor_w)
267 451191 beam_w = floor_w;
268 }
269
270
5/6
✓ Branch 0 taken 276331 times.
✓ Branch 1 taken 348926 times.
✓ Branch 2 taken 127592 times.
✓ Branch 3 taken 46474 times.
✗ Branch 4 not taken.
✓ Branch 5 taken 127592 times.
625257 if (nlist > 0 && nprobe >= nlist && beam_w < nprobe)
271 451191 beam_w = nprobe;
272
273 625257 return beam_w;
274 }
275
276 /*
277 * Leaf clusters the centroid beam routes to when the scan will read nprobe
278 * of them, capped at cap (the most the caller can route to: the scan's beam
279 * capacity, or the number of posting lists).
280 *
281 * Exact centroid formats need no expansion -- their probe order is already
282 * correct -- so those route exactly nprobe. Compressed formats route more
283 * and let phase A re-rank the wider set on exact distances, keeping the
284 * best nprobe of them.
285 *
286 * Called by prism_query_execute and by the cost model, which prices one head
287 * page read per routed cluster.
288 */
289 uint32_t
290 4368 prism_query_routed_clusters(
291 uint32_t nprobe, uint32_t cap, PrismCentroidFormat centroid_format)
292 {
293
4/6
✓ Branch 0 taken 4368 times.
✗ Branch 1 not taken.
✓ Branch 2 taken 4328 times.
✓ Branch 3 taken 40 times.
✗ Branch 4 not taken.
✓ Branch 5 taken 3704 times.
4368 if (g_probe_expand <= 1.0 || centroid_format == PRISM_CENTROID_FMT_FLOAT ||
294 centroid_format == PRISM_CENTROID_FMT_HALF)
295 ✗ return nprobe;
296
297 4328 double expanded = (double)nprobe * g_probe_expand;
298 4328 uint32_t n_route = (uint32_t)(expanded + 0.5);
299
300
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 3704 times.
4328 if (n_route > nprobe + PRISM_PROBE_EXPAND_MAX_EXTRA)
301 ✗ n_route = nprobe + PRISM_PROBE_EXPAND_MAX_EXTRA;
302
2/2
✓ Branch 0 taken 2696 times.
✓ Branch 1 taken 1008 times.
4328 if (n_route > cap)
303 2696 n_route = cap;
304
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 3704 times.
4328 if (n_route < nprobe)
305 664 n_route = nprobe;
306
307 3704 return n_route;
308 }
309
310 /* qsort comparator for probe_order indices by probe_dists (context via
311 * a file-static base pointer; the scan path is single-threaded per
312 * backend). */
313 static const float *g_probe_sort_dists;
314
315 static int
316 19789 cmp_probe_order(const void *a, const void *b)
317 {
318 19789 float da = g_probe_sort_dists[*(const uint32_t *)a];
319 19789 float db = g_probe_sort_dists[*(const uint32_t *)b];
320
2/2
✓ Branch 0 taken 10924 times.
✓ Branch 1 taken 8865 times.
19789 if (da < db)
321 8875 return -1;
322
2/2
✓ Branch 0 taken 8807 times.
✓ Branch 1 taken 2 times.
8809 if (da > db)
323 8807 return 1;
324 ✗ return 0;
325 }
326
327 static void
328 3986 scan_clusters(
329 PrismQueryState *qs,
330 const PrismCentroidResult *beam_results,
331 uint32_t n_results,
332 uint32_t scan_limit,
333 VsDistanceMode mode,
334 VsTopK *topk,
335 PrismQueryStats *stats)
336 {
337 3986 const PrismIndexBase *idx = qs->index;
338 3986 Dimension dim = idx->dim;
339
340 3986 qs->pscan.storage = idx->posting_storage;
341
342 /*
343 * Warm the cache before the serial per-cluster reads below. Both the
344 * Phase-A re-rank and the scan itself fetch each probed cluster's head
345 * page one at a time; on a cold buffer cache that is a string of
346 * synchronous random reads. Issue async prefetches for every head up
347 * front so the reads overlap. Best-effort and a no-op where the storage
348 * has no prefetch (standalone) or the page is already resident.
349 */
350
2/2
✓ Branch 0 taken 40152 times.
✓ Branch 1 taken 3986 times.
44138 for (uint32_t j = 0; j < n_results; j++)
351 40152 vs_storage_prefetch(qs->pscan.storage, beam_results[j].posting_head);
352
353 /*
354 * Phase A (only when the probe set was expanded): re-rank the routed
355 * clusters by EXACT query-centroid distance. The beam's RaBitQ
356 * distances are 1-bit estimates whose noise scrambles the probe
357 * order; each cluster's first posting page stores the full-precision
358 * rotated centroid, so one page read + one O(dim) distance per
359 * candidate recovers the true order. Only the best `scan_limit`
360 * clusters are then scanned.
361 */
362 3986 const uint32_t *order = NULL;
363 3986 uint32_t n_scan = n_results;
364
365
2/2
✓ Branch 0 taken 1049 times.
✓ Branch 1 taken 2937 times.
3986 if (n_results > scan_limit)
366 {
367
2/2
✓ Branch 0 taken 9092 times.
✓ Branch 1 taken 1049 times.
10141 for (uint32_t j = 0; j < n_results; j++)
368 {
369 9092 qs->probe_order[j] = j;
370 9092 qs->probe_dists[j] = FLT_MAX;
371
372 9092 BlockNumber ph = beam_results[j].posting_head;
373
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 9092 times.
9092 if (ph == InvalidBlockNumber)
374 ✗ continue;
375
376 9092 prism_posting_scan_begin_cluster(&qs->pscan, &qs->cluster_qs, ph);
377 9092 const float *pt_cent = prism_posting_scan_pt_centroid(&qs->pscan);
378
1/2
✓ Branch 0 taken 9092 times.
✗ Branch 1 not taken.
9092 if (pt_cent != NULL)
379 9092 qs->probe_dists[j] =
380 9092 vs_l2_distance_squared(qs->pt_query, pt_cent, dim);
381 9092 prism_posting_scan_end_cluster(&qs->pscan);
382 }
383
384 1049 g_probe_sort_dists = qs->probe_dists;
385 1049 qsort(qs->probe_order, n_results, sizeof(uint32_t), cmp_probe_order);
386
387 1049 order = qs->probe_order;
388 1049 n_scan = scan_limit;
389 }
390
391 3986 uint32_t total_pages = 0;
392 3986 uint32_t total_skipped = 0;
393 3986 uint32_t total_entries = 0;
394 3986 uint32_t scanned = 0;
395
396
2/2
✓ Branch 0 taken 35452 times.
✓ Branch 1 taken 3986 times.
39438 for (uint32_t r = 0; r < n_scan; r++)
397 {
398
2/2
✓ Branch 0 taken 4392 times.
✓ Branch 1 taken 31060 times.
35452 uint32_t j = order ? order[r] : r;
399 35452 BlockNumber ph = beam_results[j].posting_head;
400
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 35452 times.
35452 if (ph == InvalidBlockNumber)
401 ✗ continue;
402
403 /* Diagnostic: stamp candidates inserted while scanning this cluster
404 * with its scan rank r (the position in the exact-re-ranked probe
405 * order, not the beam index j), so the deepest contributing rank
406 * measures how many probed clusters the query actually needed. */
407 35452 topk->cur_src = r;
408
409 35452 prism_posting_scan_begin_cluster(&qs->pscan, &qs->cluster_qs, ph);
410
411 35452 const float *pt_cent = prism_posting_scan_pt_centroid(&qs->pscan);
412
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 35452 times.
35452 if (pt_cent == NULL)
413 {
414 ✗ prism_posting_scan_end_cluster(&qs->pscan);
415 ✗ continue;
416 }
417
418 35452 vs_rabitq_init_query_state(
419 35452 &qs->cluster_qs, qs->pt_query, pt_cent, dim, mode);
420
421
3/4
✓ Branch 0 taken 9136 times.
✓ Branch 1 taken 26316 times.
✓ Branch 2 taken 9136 times.
✗ Branch 3 not taken.
35452 if (idx->fastscan && qs->pscan.fs_lut != NULL)
422 9136 prism_posting_scan_cluster_fastscan(&qs->pscan, topk);
423 else
424 26316 prism_posting_scan_cluster(&qs->pscan, topk);
425 35452 total_pages += qs->pscan.pages_read;
426 35452 total_skipped += qs->pscan.pages_skipped;
427 35452 total_entries += qs->pscan.entries_scanned;
428 35452 prism_posting_scan_end_cluster(&qs->pscan);
429 35452 scanned++;
430 }
431
432 3986 qs->pscan.storage = NULL;
433
434
2/2
✓ Branch 0 taken 282 times.
✓ Branch 1 taken 3704 times.
3986 if (stats != NULL)
435 {
436 282 stats->clusters_scanned = scanned;
437 282 stats->posting_pages_read = total_pages;
438 282 stats->posting_pages_skipped = total_skipped;
439 282 stats->posting_entries_scanned = total_entries;
440 }
441 3986 }
442
443 static uint32_t
444 3986 extract_candidates(PrismQueryState *qs, uint32_t cap)
445 {
446
2/2
✓ Branch 0 taken 99 times.
✓ Branch 1 taken 3887 times.
3986 if (qs->topk.cand_count > qs->cand_cap)
447 {
448 /* Grow geometrically so a run of gradually larger queries does
449 * not realloc+copy on every step to the high-water mark. */
450 99 uint32_t want = qs->cand_cap * 2;
451
2/2
✓ Branch 0 taken 18 times.
✓ Branch 1 taken 31 times.
99 if (want < qs->topk.cand_count)
452 18 want = qs->topk.cand_count;
453 99 qs->cand_cap = want;
454 149 qs->candidates =
455 99 vs_realloc(qs->candidates, qs->cand_cap * sizeof(VsTopKEntry));
456 149 qs->result_order =
457 99 vs_realloc(qs->result_order, qs->cand_cap * sizeof(uint32_t));
458 99 qs->result_dists =
459 99 vs_realloc(qs->result_dists, qs->cand_cap * sizeof(Distance));
460 }
461
462 282 uint32_t ncands;
463 3986 vs_topk_extract_sorted_capped(&qs->topk, qs->candidates, &ncands, cap);
464 3986 qs->ncandidates = ncands;
465 3986 return ncands;
466 }
467
468 /* Monotonic nanosecond clock for per-phase query instrumentation. */
469 static inline uint64_t
470 1265694 prism_query_now_ns(void)
471 {
472 902746 struct timespec ts;
473 1265694 clock_gettime(CLOCK_MONOTONIC, &ts);
474 1265694 return (uint64_t)ts.tv_sec * VS_NS_PER_SEC + (uint64_t)ts.tv_nsec;
475 }
476
477 uint32_t
478 624875 prism_query_route(
479 PrismQueryState *qs,
480 const float *query,
481 uint32_t nprobe,
482 VsDistanceMode mode,
483 PrismCentroidSearchStats *beam_stats)
484 {
485
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 174066 times.
624875 if (nprobe > qs->max_nprobe)
486 ✗ nprobe = qs->max_nprobe;
487
488 624875 PrismCentroidSearchStats local = {0};
489
2/2
✓ Branch 0 taken 170644 times.
✓ Branch 1 taken 454231 times.
624875 PrismCentroidSearchStats *bs = beam_stats ? beam_stats : &local;
490
491 624875 const float *qvec = prepare_query(qs, query);
492 624875 uint64_t t_rot = prism_query_now_ns();
493 624875 vs_rabitq_rotate(qs->index->params, qvec, qs->pt_query);
494 624875 bs->rotation_ns = prism_query_now_ns() - t_rot;
495
496 624875 return search_centroids(qs, qvec, nprobe, mode, bs);
497 }
498
499 /* Cap on the rerank candidate pool (prism.rerank_pool). Candidates are
500 * sorted by approximate distance, so capping keeps the most promising
501 * ones and bounds the exact-distance heap fetches. 0 (default) resolves
502 * to an automatic cap of max(3 * k * nprobe^0.15, candidate-buffer count
503 * / 8): the buffer population directly measures estimate noise, so the
504 * nprobe-scaled floor (fit against rekall's cohere-1m recall-vs-pool
505 * sweep: the pool needed to keep the rerank-induced recall deficit under
506 * 0.1% relative to an unbounded pool, at nprobe in {10,20,40,80,160},
507 * power-law-fits to 30 * nprobe^0.15 for k=10) grows further when noisy
508 * estimates flood the buffer and ranking into just that floor would
509 * silently cap recall far below what the probed clusters contain. A
510 * flat 16 * k floor (this formula's predecessor) was measured
511 * recall-neutral, but oversized at every nprobe on that sweep -- e.g.
512 * 2-4x more pool than needed below nprobe=80, wasting rerank work
513 * without buying recall. -1 disables the cap entirely; positive values
514 * are absolute. The effective cap is never below k, so a cap can never
515 * truncate the result set. */
516 #define PRISM_RERANK_POOL_AUTO_COEFF 3.0
517 #define PRISM_RERANK_POOL_AUTO_EXP 0.15
518
519 static int32_t g_rerank_pool = 0;
520
521 void
522 266 prism_query_set_rerank_pool(int32_t n)
523 {
524 266 g_rerank_pool = n;
525 266 }
526
527 /*
528 * Calculates the size of the rerank pool from k, the number of neighbours
529 * the query will return, and nprobe, the number of posting lists it will
530 * scan -- as far as that can be known without running the scan.
531 *
532 * A return of 0 means the pool is uncapped: it is the value
533 * vs_topk_extract_sorted_capped reads as "keep every survivor", so 0 is
534 * the widest possible pool and not the narrowest. Callers that price the
535 * pool have to special-case it.
536 *
537 * Shared with the cost model, which has to price the fetches the scan will
538 * actually make. The scan adds one term this cannot: a floor at an eighth of
539 * the candidate buffer, which measures estimate noise and so does not exist
540 * until the clusters have been scanned.
541 */
542 uint32_t
543 4367 prism_query_rerank_pool_estimate(uint32_t k, uint32_t nprobe)
544 {
545
2/2
✓ Branch 0 taken 661 times.
✓ Branch 1 taken 3706 times.
4367 if (g_rerank_pool < 0)
546 ✗ return 0; /* uncapped: every survivor is reranked */
547
548
2/2
✓ Branch 0 taken 5 times.
✓ Branch 1 taken 4360 times.
4365 if (g_rerank_pool > 0)
549 {
550 5 uint32_t pool = (uint32_t)g_rerank_pool;
551
552 5 return pool < k ? k : pool;
553 }
554
555 4360 double auto_floor = PRISM_RERANK_POOL_AUTO_COEFF * (double)k *
556 4360 pow((double)nprobe, PRISM_RERANK_POOL_AUTO_EXP);
557 4360 uint32_t pool = (uint32_t)(auto_floor + 0.5);
558
559 4360 return pool < k ? k : pool;
560 }
561
562 uint32_t
563 3986 prism_query_execute(
564 PrismQueryState *qs,
565 const float *query,
566 uint32_t k,
567 uint32_t nprobe,
568 VsDistanceMode mode,
569 bool rerank,
570 PrismQueryStats *stats)
571 {
572
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 3704 times.
3986 if (k > qs->max_k)
573 ✗ k = qs->max_k;
574
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 3704 times.
3986 if (nprobe > qs->max_nprobe)
575 ✗ nprobe = qs->max_nprobe;
576
577 /* reset_to_k also (re)sizes the heap when k grows between queries —
578 * resetting first and assigning k afterwards left the heap sized for
579 * the previous k. */
580 3986 vs_topk_reset_to_k(&qs->topk, k);
581
582 /* Probe expansion: route extra leaf candidates so phase A of
583 * scan_clusters can pick the best `nprobe` by exact centroid
584 * distance. n_route == nprobe (expand == 1, no expansion) keeps
585 * the classic single-phase behavior. Skipped entirely when the
586 * centroid pages are exact (float/half): the beam distances are
587 * already exact, so there is no ordering noise to correct. */
588 4268 uint32_t n_route = prism_query_routed_clusters(
589 3986 nprobe, qs->max_nprobe, qs->index->centroid_format);
590
591 3986 uint64_t t0 = prism_query_now_ns();
592
593 3986 PrismCentroidSearchStats beam_stats = {0};
594 282 uint32_t ncentroids =
595 3986 prism_query_route(qs, query, n_route, mode, &beam_stats);
596
597 /* prism_query_route already normalized the query into qs->query_buf (for
598 * cosine) via prepare_query; reuse it for the rerank below instead of
599 * re-normalizing (a redundant O(dim) memcpy+norm+scale per query). For
600 * non-cosine metrics prepare_query is a no-op and returns the raw query.
601 */
602 7972 const float *qvec = (qs->index->metric == DISTANCE_COSINE) ? qs->query_buf
603
2/2
✓ Branch 0 taken 14 times.
✓ Branch 1 taken 3972 times.
3986 : query;
604
605 3986 uint64_t t1 = prism_query_now_ns();
606
607 /*
608 * The probed heads are chosen and their centroid pages released, but none
609 * of the lists has been opened yet. A concurrent split can replace and
610 * retire a head in this window, which is why the old chain stays readable
611 * until no snapshot can reach it; an isolation test pauses here to hold a
612 * head across exactly that.
613 */
614 282 VS_INJECTION_POINT("prism-scan-routed");
615
616 3986 scan_clusters(
617 3986 qs, qs->beam_results, ncentroids, nprobe, mode, &qs->topk, stats);
618
619 /* Rerank-pool cap: the first `pool` candidates by approximate
620 * distance are the most promising; see prism_query_set_rerank_pool.
621 * Resolved before extraction so the extract can select the capped
622 * prefix instead of fully sorting an unbounded survivor set.
623 *
624 * The automatic cap also scales with the candidate-buffer
625 * population, which directly measures estimate noise: accurate
626 * estimates keep the threshold at the nprobe-scaled floor above,
627 * while noisy estimates (low dimension, wide norm spread) flood
628 * the buffer -- and then ranking into just that floor is
629 * meaningless, silently capping recall well below what the probed
630 * clusters contain. 1/8th of the buffer restores the recall
631 * ceiling at a rerank cost proportionate to the observed noise. */
632 3986 uint32_t pool = prism_query_rerank_pool_estimate(k, nprobe);
633
634 /*
635 * The noise term needs the candidate population, which exists only now
636 * that the clusters have been scanned -- so it cannot be part of the
637 * shared estimate the planner uses.
638 */
639
3/4
✓ Branch 0 taken 3986 times.
✗ Branch 1 not taken.
✓ Branch 2 taken 3168 times.
✓ Branch 3 taken 536 times.
3986 if (g_rerank_pool == 0 && pool < qs->topk.cand_count / 8)
640 3450 pool = qs->topk.cand_count / 8;
641
3/4
✓ Branch 0 taken 3704 times.
✓ Branch 1 taken 282 times.
✗ Branch 2 not taken.
✓ Branch 3 taken 3704 times.
3986 if (pool > 0 && pool < k)
642 ✗ pool = k;
643
644 3986 uint32_t ncands = extract_candidates(qs, pool);
645
646 3986 uint64_t t2 = prism_query_now_ns();
647
648 /* Rerank with exact distances if enabled and storage supports it */
649 3986 VsStorage *ps = qs->index->posting_storage;
650
5/8
✓ Branch 0 taken 3982 times.
✓ Branch 1 taken 4 times.
✓ Branch 2 taken 3982 times.
✗ Branch 3 not taken.
✓ Branch 4 taken 3982 times.
✗ Branch 5 not taken.
✓ Branch 6 taken 3704 times.
✗ Branch 7 not taken.
3986 if (rerank && ncands > 0 && ps != NULL && ps->ops->rerank != NULL)
651 {
652 3982 qs->nresults = vs_storage_rerank(
653 ps,
654 qvec,
655 3982 qs->index->dim,
656 3982 qs->candidates,
657 ncands,
658 k,
659 qs->result_order,
660 qs->result_dists);
661 }
662 else
663 {
664 4 qs->nresults = ncands < k ? ncands : k;
665
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 4 times.
4 for (uint32_t i = 0; i < qs->nresults; i++)
666 {
667 ✗ qs->result_order[i] = i;
668 ✗ qs->result_dists[i] = qs->candidates[i].distance;
669 }
670 }
671
672 #ifndef NDEBUG
673
2/2
✓ Branch 0 taken 12933 times.
✓ Branch 1 taken 3986 times.
16919 for (uint32_t i = 0; i < qs->nresults; i++)
674 {
675 12933 uint64_t id_i = qs->candidates[qs->result_order[i]].id;
676
2/2
✓ Branch 0 taken 1118409 times.
✓ Branch 1 taken 12933 times.
1131342 for (uint32_t j = i + 1; j < qs->nresults; j++)
677
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 1118409 times.
1118409 if (id_i == qs->candidates[qs->result_order[j]].id)
678
0/2
✗ Branch 1 not taken.
✗ Branch 2 not taken.
1112789 vs_warn(VS_EXTENSION_NAME ": duplicate result at positions "
679 "%u and %u",
680 i,
681 j);
682 }
683 #endif
684
685 3986 uint64_t t3 = prism_query_now_ns();
686
687 /* Routing-quality diagnostic: deepest probe rank contributing a final
688 * top-k result. Low values (relative to nprobe) => over-probing; values
689 * near nprobe => neighbors genuinely routed deep (mis-routing). */
690 3986 uint32_t max_rank = 0;
691
3/3
✓ Branch 0 taken 4836 times.
✓ Branch 1 taken 11801 times.
✓ Branch 2 taken 282 times.
16919 for (uint32_t i = 0; i < qs->nresults; i++)
692 {
693 12933 uint32_t r = qs->candidates[qs->result_order[i]].src;
694
2/2
✓ Branch 0 taken 1069 times.
✓ Branch 1 taken 3767 times.
12933 if (r > max_rank)
695 1069 max_rank = r;
696 }
697
698
2/2
✓ Branch 0 taken 282 times.
✓ Branch 1 taken 3704 times.
3986 if (stats != NULL)
699 {
700 /* clusters_scanned is set by scan_clusters (actual count). */
701 282 stats->max_contrib_rank = max_rank;
702 282 stats->centroid_pages_read = beam_stats.pages_read;
703 282 stats->centroid_ns = t1 - t0;
704 282 stats->posting_ns = t2 - t1;
705 282 stats->rerank_ns = t3 - t2;
706 282 stats->rotation_ns = beam_stats.rotation_ns;
707 282 stats->centroid_lut_ns = beam_stats.lut_ns;
708 282 stats->centroid_pageread_ns = beam_stats.pageread_ns;
709 282 stats->centroid_score_ns = beam_stats.score_ns;
710 }
711
712 3986 return qs->nresults;
713 }
714
715 uint32_t
716 351 prism_auto_nprobe(uint32_t nlist)
717 {
718 /* See the header: ~0.5 * sqrt(nlist), floored at 10, capped at
719 * 2048, never above nlist. */
720 351 uint32_t nprobe = (uint32_t)ceil(0.5 * sqrt((double)nlist));
721
2/2
✓ Branch 0 taken 6 times.
✓ Branch 1 taken 8 times.
351 if (nprobe < 10)
722 6 nprobe = 10;
723
2/2
✓ Branch 0 taken 2 times.
✓ Branch 1 taken 12 times.
351 if (nprobe > 2048)
724 2 nprobe = 2048;
725
4/4
✓ Branch 0 taken 192 times.
✓ Branch 1 taken 159 times.
✓ Branch 2 taken 2 times.
✓ Branch 3 taken 12 times.
351 if (nlist > 0 && nprobe > nlist)
726 180 nprobe = nlist;
727 351 return nprobe;
728 }
729