| 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 |