GCC Code Coverage Report


Directory: src/
File: src/algo/hkmeans.c
Date: 2026-09-30 11:11:31
Exec Total Coverage
Lines: 267 290 92.1%
Functions: 12 12 100.0%
Branches: 119 139 85.6%

Line Branch Exec Source
1 /*
2 * Copyright (c) 2026 Tiger Data, Inc.
3 * Licensed under the PostgreSQL License. See LICENSE for details.
4 *
5 * hkmeans.c - Hierarchical k-means tree builder
6 *
7 * BFS tree construction using vs_kmeans() with indexed access
8 * at each node — avoids copying sub-vector arrays.
9 *
10 * The result is packed into a single contiguous allocation so it
11 * can be memcpy'd into shared memory for parallel builds.
12 */
13
14 #include <math.h>
15 #include <string.h>
16
17 #include "algo/distance.h"
18 #include "algo/hkmeans.h"
19 #include "core/memory.h"
20
21 /* BFS work queue entry */
22 typedef struct HKWorkItem
23 {
24 uint32_t *vec_indices; /* original vector indices; NULL = identity */
25 uint32_t count;
26 uint32_t level;
27 uint32_t idx_in_level;
28 } HKWorkItem;
29
30 /*
31 * Compute the number of tree levels needed.
32 *
33 * nlevels = max(1, ceil(log(nlist) / log(fan_out)))
34 * When nlist <= fan_out, nlevels = 1 (flat).
35 */
36 static uint32_t
37 1148 compute_nlevels(uint32_t nlist, uint32_t fan_out)
38 {
39
2/2
✓ Branch 0 taken 730 times.
✓ Branch 1 taken 418 times.
1148 if (nlist <= fan_out)
40 296 return 1;
41
42 418 double levels = ceil(log((double)nlist) / log((double)fan_out));
43 418 uint32_t n = (uint32_t)levels;
44
2/2
✓ Branch 0 taken 242 times.
✓ Branch 1 taken 176 times.
418 return n < 1 ? 1 : n;
45 }
46
47 uint32_t
48 795 vs_hkmeans_nlevels(uint32_t nlist, uint32_t fan_out)
49 {
50
2/2
✓ Branch 0 taken 150 times.
✓ Branch 1 taken 168 times.
795 if (fan_out < 2)
51 150 fan_out = 2;
52 795 return compute_nlevels(nlist, fan_out);
53 }
54
55 /*
56 * power_u32 - Compute base^exp for small unsigned integers.
57 */
58 static uint32_t
59 258 power_u32(uint32_t base, uint32_t exp)
60 {
61 258 uint32_t result = 1;
62
2/2
✓ Branch 0 taken 172 times.
✓ Branch 1 taken 391 times.
563 for (uint32_t i = 0; i < exp; i++)
63 172 result *= base;
64 391 return result;
65 }
66
67 /*
68 * Upper bound on total nodes: sum of fan_out^l for l = 0..nlevels-1.
69 */
70 static uint32_t
71 283 max_total_nodes(uint32_t fan_out, uint32_t nlevels)
72 {
73 283 uint32_t total = 0;
74
2/2
✓ Branch 0 taken 391 times.
✓ Branch 1 taken 283 times.
674 for (uint32_t l = 0; l < nlevels; l++)
75 391 total += power_u32(fan_out, l);
76 283 return total;
77 }
78
79 /*
80 * Temporary node used during BFS construction. Holds pointers
81 * to centroid data before everything is packed into the final
82 * contiguous allocation.
83 */
84 typedef struct TmpNode
85 {
86 float *centroids;
87 uint32_t nchildren;
88 uint32_t level;
89 uint32_t first_child;
90 uint32_t first_leaf;
91 size_t cent_bytes;
92 bool is_leaf_parent;
93 } TmpNode;
94
95 HKMeansResult *
96 293 vs_hkmeans_f32(
97 const float *vectors,
98 uint32_t nvecs,
99 const uint32_t *indices,
100 Dimension dim,
101 uint32_t nlist,
102 uint32_t fan_out,
103 DistanceMetric metric,
104 const KMeansOptions *options)
105 {
106
10/10
✓ Branch 0 taken 291 times.
✓ Branch 1 taken 2 times.
✓ Branch 2 taken 289 times.
✓ Branch 3 taken 2 times.
✓ Branch 4 taken 174 times.
✓ Branch 5 taken 115 times.
✓ Branch 6 taken 172 times.
✓ Branch 7 taken 2 times.
✓ Branch 8 taken 2 times.
✓ Branch 9 taken 170 times.
293 if (vectors == NULL || nvecs == 0 || dim == 0 || nlist == 0 || fan_out < 2)
107 10 return NULL;
108
109 283 uint32_t nlevels = compute_nlevels(nlist, fan_out);
110 283 uint32_t max_nodes = max_total_nodes(fan_out, nlevels);
111
112 /* Work context for BFS temporaries */
113 283 VsMemCtx work_ctx = vs_memctx_create(NULL, "hkmeans_work");
114 283 VsMemCtx caller_ctx = vs_memctx_switch(work_ctx);
115
116 /* Temporary nodes (pointers, packed later) */
117 283 TmpNode *tmp_nodes = vs_alloc(max_nodes * sizeof(TmpNode));
118
2/2
✓ Branch 0 taken 110 times.
✓ Branch 1 taken 3 times.
283 memset(tmp_nodes, 0, max_nodes * sizeof(TmpNode));
119 283 uint32_t nnodes = 0;
120 283 uint32_t nleaves = 0;
121
122 /* If caller provided indices, copy them as the root's
123 * vec_indices so vs_kmeans uses indirect access. */
124 283 uint32_t *root_indices = NULL;
125
2/2
✓ Branch 0 taken 236 times.
✓ Branch 1 taken 47 times.
283 if (indices != NULL)
126 {
127 236 root_indices = vs_alloc(nvecs * sizeof(uint32_t));
128 236 memcpy(root_indices, indices, nvecs * sizeof(uint32_t));
129 }
130
131 /* BFS work queue */
132 283 HKWorkItem *queue = vs_alloc(max_nodes * sizeof(HKWorkItem));
133 283 uint32_t q_tail = 0;
134
135 283 queue[q_tail++] = (HKWorkItem){
136 .vec_indices = root_indices,
137 .count = nvecs,
138 .level = 0,
139 .idx_in_level = 0,
140 };
141
142 283 bool ok = true;
143 283 uint32_t *counts = vs_alloc(fan_out * sizeof(uint32_t));
144
145 /* Mutable copy: initial_centroids applies only to root */
146 283 KMeansOptions local_opts = VS_KMEANS_OPTIONS_DEFAULT;
147
1/2
✓ Branch 0 taken 283 times.
✗ Branch 1 not taken.
283 if (options != NULL)
148 283 local_opts = *options;
149
150
3/4
✓ Branch 0 taken 871 times.
✓ Branch 1 taken 283 times.
✓ Branch 2 taken 686 times.
✗ Branch 3 not taken.
1154 for (uint32_t qi = 0; qi < q_tail && ok; qi++)
151 {
152 871 HKWorkItem item = queue[qi];
153 871 bool is_leaf_parent = (item.level == nlevels - 1);
154
155 /*
156 * Determine K for this node.
157 *
158 * Single-level (flat): K = nlist (clamped to vector count)
159 * Multi-level root/internal: K = fan_out (clamped)
160 */
161 185 uint32_t k;
162
2/2
✓ Branch 0 taken 197 times.
✓ Branch 1 taken 674 times.
871 if (nlevels == 1)
163 197 k = nlist < item.count ? nlist : item.count;
164 else
165 674 k = fan_out < item.count ? fan_out : item.count;
166
167 /* Use indexed k-means — no vector copy needed */
168 871 KMeansResult *km = vs_kmeans(
169 vectors,
170 686 item.vec_indices,
171 VS_VEC_F32,
172 item.count,
173 dim,
174 k,
175 metric,
176 &local_opts);
177
178 /* Initial centroids only for root node */
179 871 local_opts.initial_centroids = NULL;
180
181
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 871 times.
871 if (km == NULL)
182 {
183 ✗ ok = false;
184 ✗ break;
185 }
186
187 871 uint32_t node_idx = nnodes++;
188 871 TmpNode *tn = &tmp_nodes[node_idx];
189
190 871 size_t cent_sz = (size_t)km->nlist * dim * sizeof(float);
191 871 tn->centroids = vs_alloc(cent_sz);
192 871 tn->cent_bytes = cent_sz;
193
2/2
✓ Branch 0 taken 165 times.
✓ Branch 1 taken 20 times.
871 memcpy(tn->centroids, km->centroids, cent_sz);
194 871 tn->nchildren = km->nlist;
195 871 tn->level = item.level;
196 871 tn->first_child = HKMEANS_NO_CHILD;
197 871 tn->first_leaf = 0;
198 871 tn->is_leaf_parent = is_leaf_parent;
199
200
2/2
✓ Branch 0 taken 647 times.
✓ Branch 1 taken 224 times.
871 if (is_leaf_parent)
201 647 nleaves += km->nlist;
202 else
203 {
204 /*
205 * Count vectors per cluster from assignments.
206 *
207 * We cannot use km->cluster_sizes because k-means
208 * does a final reassignment after the last update
209 * step, so cluster_sizes may be stale.
210 */
211 224 memset(counts, 0, km->nlist * sizeof(uint32_t));
212
2/2
✓ Branch 0 taken 52490 times.
✓ Branch 1 taken 224 times.
52714 for (uint32_t v = 0; v < item.count; v++)
213 52490 counts[km->assignments[v]]++;
214
215 /*
216 * Keep only non-empty clusters as children, compacting their
217 * centroids to match. k-means can leave a cluster empty on
218 * degenerate/collapsing data; tree descent indexes a node's
219 * children as first_child + c over its centroids, so the kept
220 * centroids and the enqueued child nodes must stay 1:1 and
221 * contiguous. (Empty internal clusters previously left
222 * nchildren > children-created, so descent could read past the
223 * nodes array.) item.count > 0 here, so at least one cluster is
224 * non-empty and kept >= 1.
225 */
226 224 uint32_t kept = 0;
227 224 tn->first_child = q_tail;
228
2/2
✓ Branch 0 taken 594 times.
✓ Branch 1 taken 224 times.
818 for (uint32_t c = 0; c < km->nlist; c++)
229 {
230 594 uint32_t sub_n = counts[c];
231
2/2
✓ Branch 0 taken 6 times.
✓ Branch 1 taken 588 times.
594 if (sub_n == 0)
232 6 continue;
233
234
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 588 times.
588 if (kept != c)
235 ✗ memcpy(tn->centroids + (size_t)kept * dim,
236 ✗ km->centroids + (size_t)c * dim,
237 ✗ (size_t)dim * sizeof(float));
238
239 588 uint32_t *sub_indices = vs_alloc(sub_n * sizeof(uint32_t));
240 588 uint32_t idx = 0;
241
3/3
✓ Branch 0 taken 175296 times.
✓ Branch 1 taken 8680 times.
✓ Branch 2 taken 72 times.
184048 for (uint32_t v = 0; v < item.count; v++)
242 {
243
2/2
✓ Branch 0 taken 52490 times.
✓ Branch 1 taken 130970 times.
183460 if (km->assignments[v] == c)
244 {
245 136914 uint32_t orig = item.vec_indices ? item.vec_indices[v]
246
2/2
✓ Branch 0 taken 31934 times.
✓ Branch 1 taken 20556 times.
52490 : v;
247 52490 sub_indices[idx++] = orig;
248 }
249 }
250
251 588 queue[q_tail++] = (HKWorkItem){
252 .vec_indices = sub_indices,
253 .count = sub_n,
254 588 .level = item.level + 1,
255 588 .idx_in_level = item.idx_in_level * fan_out + kept,
256 };
257 588 kept++;
258 }
259 224 tn->nchildren = kept;
260 224 tn->cent_bytes = (size_t)kept * dim * sizeof(float);
261 }
262
263 871 vs_kmeans_result_destroy(km);
264 }
265
266
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 170 times.
283 if (!ok)
267 {
268 ✗ vs_memctx_switch(caller_ctx);
269 ✗ vs_memctx_delete(work_ctx);
270 ✗ return NULL;
271 }
272
273 /* Compute first_leaf for leaf-parent nodes */
274 170 uint32_t leaf_off = 0;
275
2/2
✓ Branch 0 taken 871 times.
✓ Branch 1 taken 283 times.
1154 for (uint32_t i = 0; i < nnodes; i++)
276 {
277
2/2
✓ Branch 0 taken 224 times.
✓ Branch 1 taken 647 times.
871 if (!tmp_nodes[i].is_leaf_parent)
278 224 continue;
279 647 tmp_nodes[i].first_leaf = leaf_off;
280 647 leaf_off += tmp_nodes[i].nchildren;
281 }
282
283 /*
284 * Pack the whole tree into one contiguous allocation, laid out as
285 * [header][nodes][leaf centroids][internal centroids] and addressed
286 * by byte offsets rather than pointers. This is what lets a built
287 * tree be memcpy'd into shared memory and used by another process
288 * (the PG parallel build) with no pointer fix-up or deserialization.
289 */
290 283 size_t hdr_sz = sizeof(HKMeansResult);
291 283 size_t nodes_sz = (size_t)nnodes * sizeof(HKMeansNode);
292 283 size_t leaf_sz = (size_t)nleaves * dim * sizeof(float);
293 283 size_t intern_sz = 0;
294
2/2
✓ Branch 0 taken 871 times.
✓ Branch 1 taken 283 times.
1154 for (uint32_t i = 0; i < nnodes; i++)
295 {
296
2/2
✓ Branch 0 taken 224 times.
✓ Branch 1 taken 647 times.
871 if (!tmp_nodes[i].is_leaf_parent)
297 224 intern_sz += tmp_nodes[i].cent_bytes;
298 }
299 283 size_t total = hdr_sz + nodes_sz + leaf_sz + intern_sz;
300
301 283 vs_memctx_switch(caller_ctx);
302 283 HKMeansResult *result = vs_alloc(total);
303 283 memset(result, 0, total);
304
305 283 result->nodes_offset = (uint32_t)hdr_sz;
306 283 result->leaf_offset = (uint32_t)(hdr_sz + nodes_sz);
307 283 result->total_size = (uint32_t)total;
308 283 result->nnodes = nnodes;
309 283 result->nlevels = nlevels;
310 283 result->nleaves = nleaves;
311 283 result->fan_out = fan_out;
312 283 result->dim = dim;
313
314 283 HKMeansNode *nodes = hk_nodes(result);
315 283 float *leaf_cents = hk_leaf_centroids(result);
316 283 char *intern_dst = (char *)result + hdr_sz + nodes_sz + leaf_sz;
317
318 /* Copy leaf centroids into the leaf area */
319
2/2
✓ Branch 0 taken 871 times.
✓ Branch 1 taken 283 times.
1154 for (uint32_t i = 0; i < nnodes; i++)
320 {
321 871 TmpNode *tn = &tmp_nodes[i];
322
2/2
✓ Branch 0 taken 224 times.
✓ Branch 1 taken 647 times.
871 if (!tn->is_leaf_parent)
323 224 continue;
324 647 memcpy(leaf_cents + (size_t)tn->first_leaf * dim,
325 647 tn->centroids,
326 tn->cent_bytes);
327 }
328
329 /* Copy internal centroids and build final nodes */
330
2/2
✓ Branch 0 taken 871 times.
✓ Branch 1 taken 283 times.
1154 for (uint32_t i = 0; i < nnodes; i++)
331 {
332 871 TmpNode *tn = &tmp_nodes[i];
333
334 871 nodes[i].nchildren = tn->nchildren;
335 871 nodes[i].level = tn->level;
336 871 nodes[i].first_child = tn->first_child;
337 871 nodes[i].first_leaf = tn->first_leaf;
338
339
2/2
✓ Branch 0 taken 647 times.
✓ Branch 1 taken 224 times.
871 if (tn->is_leaf_parent)
340 {
341 /* Point into the leaf centroids area */
342 647 nodes[i].centroid_offset = result->leaf_offset +
343 647 (uint32_t)((size_t)tn->first_leaf *
344 dim * sizeof(float));
345 }
346 else
347 {
348 /* Copy into the internal area */
349 224 memcpy(intern_dst, tn->centroids, tn->cent_bytes);
350 224 nodes[i].centroid_offset = (uint32_t)((size_t)(intern_dst -
351 (char *)result));
352 224 intern_dst += tn->cent_bytes;
353 }
354 }
355
356 /* Bulk-free all BFS temporaries */
357 283 vs_memctx_delete(work_ctx);
358
359 283 return result;
360 }
361
362 size_t
363 70 vs_hkmeans_max_blob_size_capped(
364 uint32_t nlist, uint32_t fan_out, Dimension dim, uint64_t max_leaves)
365 {
366
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 50 times.
70 if (fan_out < 2)
367 ✗ fan_out = 2;
368
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 50 times.
70 if (max_leaves < 1)
369 ✗ max_leaves = 1;
370
371 70 uint32_t nlevels = compute_nlevels(nlist, fan_out);
372
373 /* All worst-case counts are fan_out powers; at large nlist with a small
374 * fan_out they exceed 32 bits, so the whole bound is computed in 64-bit
375 * (the caller compares it against the blob format's 32-bit offset limit
376 * and fails the build rather than wrapping into an undersized slot).
377 *
378 * max_leaves caps every level's node count: a node exists only where at
379 * least one training vector landed, so no level can hold more nodes
380 * than the tree has vectors -- the depth stays the full nlevels (few
381 * vectors under a deep target degenerate into chains), but each level's
382 * width is min(fan_out^l, max_leaves). UINT64_MAX = the analytic
383 * worst case. */
384 70 uint64_t width = 1; /* fan_out^l, capped at max_leaves */
385 70 uint64_t nleaves = 0;
386 70 uint64_t nnodes = 0; /* sum of capped widths, l = 0..nlevels-1 */
387 70 uint64_t intern_nodes = 0; /* same sum, one level shorter */
388
3/3
✓ Branch 0 taken 120 times.
✓ Branch 1 taken 75 times.
✓ Branch 2 taken 20 times.
215 for (uint32_t l = 0; l < nlevels; l++)
389 {
390 145 nnodes += width;
391
4/4
✓ Branch 0 taken 110 times.
✓ Branch 1 taken 35 times.
✓ Branch 2 taken 75 times.
✓ Branch 3 taken 35 times.
145 if (nlevels >= 2 && l < nlevels - 1)
392 75 intern_nodes += width;
393
2/2
✓ Branch 0 taken 41 times.
✓ Branch 1 taken 104 times.
145 if (width >= max_leaves / fan_out)
394 28 width = max_leaves;
395 else
396 105 width *= fan_out;
397 }
398 70 nleaves = width;
399
400 70 return sizeof(HKMeansResult) + nnodes * sizeof(HKMeansNode) +
401 120 nleaves * dim * sizeof(float) +
402 70 intern_nodes * fan_out * dim * sizeof(float);
403 }
404
405 size_t
406 10 vs_hkmeans_max_blob_size(uint32_t nlist, uint32_t fan_out, Dimension dim)
407 {
408 10 return vs_hkmeans_max_blob_size_capped(nlist, fan_out, dim, UINT64_MAX);
409 }
410
411 HKMeansResult *
412 78 vs_hkmeans_build_flat(
413 const float *centroids,
414 uint32_t nleaves,
415 uint32_t fan_out,
416 Dimension dim)
417 {
418
3/6
✓ Branch 0 taken 78 times.
✗ Branch 1 not taken.
✓ Branch 2 taken 78 times.
✗ Branch 3 not taken.
✗ Branch 4 not taken.
✓ Branch 5 taken 54 times.
78 if (centroids == NULL || nleaves == 0 || dim == 0)
419 ✗ return NULL;
420
421 78 size_t hdr_sz = sizeof(HKMeansResult);
422 78 size_t nodes_sz = sizeof(HKMeansNode); /* single leaf-parent root */
423 78 size_t leaf_sz = (size_t)nleaves * dim * sizeof(float);
424 78 size_t total = hdr_sz + nodes_sz + leaf_sz;
425
426 78 HKMeansResult *result = vs_alloc(total);
427 78 memset(result, 0, total);
428
429 78 result->nodes_offset = (uint32_t)hdr_sz;
430 78 result->leaf_offset = (uint32_t)(hdr_sz + nodes_sz);
431 78 result->total_size = (uint32_t)total;
432 78 result->nnodes = 1;
433 78 result->nlevels = 1;
434 78 result->nleaves = nleaves;
435 78 result->fan_out = fan_out;
436 78 result->dim = dim;
437
438 78 HKMeansNode *root = hk_nodes(result);
439 78 root->level = 0;
440 78 root->nchildren = nleaves;
441 78 root->first_child = HKMEANS_NO_CHILD; /* leaf-parent */
442 78 root->first_leaf = 0;
443 78 root->centroid_offset = result->leaf_offset;
444 78 memcpy(hk_leaf_centroids(result), centroids, leaf_sz);
445
446 78 return result;
447 }
448
449 uint32_t
450 11448 vs_hkmeans_assign(
451 const HKMeansResult *tree,
452 const float *vec,
453 DistanceMetric metric,
454 Distance *out_distance)
455 {
456 11448 Dimension dim = tree->dim;
457 11448 const HKMeansNode *nodes = hk_nodes(tree);
458 11448 uint32_t node_idx = 0;
459
460
1/2
✓ Branch 0 taken 22886 times.
✗ Branch 1 not taken.
22886 for (uint32_t level = 0; level < tree->nlevels; level++)
461 {
462 22886 const HKMeansNode *node = &nodes[node_idx];
463 22886 const float *cents = hk_node_centroids(tree, node);
464 22886 Distance best_dist = INFINITY;
465 22886 uint32_t best_c = 0;
466
467
2/2
✓ Branch 0 taken 92754 times.
✓ Branch 1 taken 22886 times.
115640 for (uint32_t c = 0; c < node->nchildren; c++)
468 {
469 92754 const float *centroid = cents + (size_t)c * dim;
470 92754 Vec32Ref qref = {.data = vec, .dim = dim};
471 92754 Vec32Ref cref = {.data = centroid, .dim = dim};
472 92754 Distance d = vs_distance(qref, cref, metric);
473
2/2
✓ Branch 0 taken 48028 times.
✓ Branch 1 taken 44726 times.
92754 if (d < best_dist)
474 {
475 48028 best_dist = d;
476 48028 best_c = c;
477 }
478 }
479
480 /*
481 * A node with no internal children is a leaf-parent: its children are
482 * leaf centroids, reached via first_leaf. This is the authoritative
483 * test -- using level == nlevels - 1 instead breaks on non-uniform
484 * depth trees, where a branch that bottoms out early leaves a
485 * leaf-parent above the max level; first_child (NO_CHILD = UINT32_MAX)
486 * would then be added to best_c and index past the nodes array.
487 */
488
2/2
✓ Branch 0 taken 11448 times.
✓ Branch 1 taken 11438 times.
22886 if (node->first_child == HKMEANS_NO_CHILD)
489 {
490
2/2
✓ Branch 0 taken 11438 times.
✓ Branch 1 taken 10 times.
11448 if (out_distance != NULL)
491 11438 *out_distance = best_dist;
492 11448 return node->first_leaf + best_c;
493 }
494
495 11438 node_idx = node->first_child + best_c;
496 }
497
498 ✗ if (out_distance != NULL)
499 ✗ *out_distance = INFINITY;
500 ✗ return 0;
501 }
502
503 /*
504 * Insert (id, d) into an unsorted bounded set of the `cap` smallest
505 * entries. Keeps at most `cap` items; evicts the current max when full.
506 */
507 static inline void
508 1384 topk_insert(
509 uint32_t *ids,
510 Distance *dists,
511 uint32_t *n,
512 uint32_t cap,
513 uint32_t id,
514 Distance d)
515 {
516
2/2
✓ Branch 0 taken 546 times.
✓ Branch 1 taken 838 times.
1384 if (*n < cap)
517 {
518 546 ids[*n] = id;
519 546 dists[*n] = d;
520 546 (*n)++;
521 546 return;
522 }
523
524 /* Full: replace the worst entry if this one is better. */
525 838 uint32_t worst_i = 0;
526 838 Distance worst_d = dists[0];
527
2/2
✓ Branch 0 taken 2020 times.
✓ Branch 1 taken 838 times.
2858 for (uint32_t i = 1; i < cap; i++)
528
2/2
✓ Branch 0 taken 750 times.
✓ Branch 1 taken 1270 times.
2020 if (dists[i] > worst_d)
529 {
530 750 worst_d = dists[i];
531 750 worst_i = i;
532 }
533
2/2
✓ Branch 0 taken 394 times.
✓ Branch 1 taken 444 times.
838 if (d < worst_d)
534 {
535 394 ids[worst_i] = id;
536 394 dists[worst_i] = d;
537 }
538 }
539
540 /* Insertion sort by ascending distance (n is small, <= VS_HK_MAX_TOPK). */
541 static inline void
542 92 topk_sort(uint32_t *ids, Distance *dists, uint32_t n)
543 {
544
2/2
✓ Branch 0 taken 200 times.
✓ Branch 1 taken 92 times.
292 for (uint32_t i = 1; i < n; i++)
545 {
546 200 Distance d = dists[i];
547 200 uint32_t id = ids[i];
548 200 uint32_t j = i;
549
4/4
✓ Branch 0 taken 312 times.
✓ Branch 1 taken 78 times.
✓ Branch 2 taken 190 times.
✓ Branch 3 taken 122 times.
390 while (j > 0 && dists[j - 1] > d)
550 {
551 190 dists[j] = dists[j - 1];
552 190 ids[j] = ids[j - 1];
553 190 j--;
554 }
555 200 dists[j] = d;
556 200 ids[j] = id;
557 }
558 92 }
559
560 uint32_t
561 92 vs_hkmeans_assign_topk(
562 const HKMeansResult *tree,
563 const float *vec,
564 DistanceMetric metric,
565 uint32_t k,
566 uint32_t beam_width,
567 uint32_t *out_leaves,
568 Distance *out_dists)
569 {
570
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 92 times.
92 if (k == 0)
571 ✗ return 0;
572
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 92 times.
92 if (k > VS_HK_MAX_TOPK)
573 ✗ k = VS_HK_MAX_TOPK;
574
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 92 times.
92 if (beam_width < 1)
575 ✗ beam_width = 1;
576
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 92 times.
92 if (beam_width > VS_HK_MAX_TOPK)
577 ✗ beam_width = VS_HK_MAX_TOPK;
578
579 92 const Dimension dim = tree->dim;
580 92 const HKMeansNode *nodes = hk_nodes(tree);
581
582 /* Current beam: indices into nodes[] for the next level to expand. */
583 ✗ uint32_t beam[VS_HK_MAX_TOPK];
584 92 uint32_t beam_n = 1;
585 92 beam[0] = 0; /* root */
586
587 /*
588 * Terminal leaves found so far. A leaf-parent (first_child == NO_CHILD)
589 * can appear at any level on a non-uniform-depth tree, so its leaf
590 * children are collected here directly rather than assumed to all sit at
591 * nlevels - 1.
592 */
593 ✗ uint32_t res_id[VS_HK_MAX_TOPK];
594 ✗ Distance res_d[VS_HK_MAX_TOPK];
595 92 uint32_t res_n = 0;
596
597
3/4
✓ Branch 0 taken 184 times.
✓ Branch 1 taken 92 times.
✓ Branch 2 taken 184 times.
✗ Branch 3 not taken.
276 for (uint32_t level = 0; level < tree->nlevels && beam_n > 0; level++)
598 {
599 /* Internal child nodes to expand at the next level. */
600 ✗ uint32_t next[VS_HK_MAX_TOPK];
601 ✗ Distance next_d[VS_HK_MAX_TOPK];
602 184 uint32_t next_n = 0;
603
604
2/2
✓ Branch 0 taken 346 times.
✓ Branch 1 taken 184 times.
530 for (uint32_t b = 0; b < beam_n; b++)
605 {
606 346 const HKMeansNode *node = &nodes[beam[b]];
607 346 const float *cents = hk_node_centroids(tree, node);
608 346 bool node_leaf = (node->first_child == HKMEANS_NO_CHILD);
609
2/2
✓ Branch 0 taken 1384 times.
✓ Branch 1 taken 346 times.
1730 for (uint32_t c = 0; c < node->nchildren; c++)
610 {
611 1384 Vec32Ref qref = {.data = vec, .dim = dim};
612 1384 Vec32Ref cref = {.data = cents + (size_t)c * dim, .dim = dim};
613 1384 Distance d = vs_distance(qref, cref, metric);
614
2/2
✓ Branch 0 taken 1016 times.
✓ Branch 1 taken 368 times.
1384 if (node_leaf)
615 1016 topk_insert(
616 1016 res_id, res_d, &res_n, k, node->first_leaf + c, d);
617 else
618 368 topk_insert(
619 next,
620 next_d,
621 &next_n,
622 beam_width,
623 368 node->first_child + c,
624 d);
625 }
626 }
627
628 184 memcpy(beam, next, next_n * sizeof(uint32_t));
629 184 beam_n = next_n;
630 }
631
632 92 topk_sort(res_id, res_d, res_n);
633
2/3
✓ Branch 0 taken 292 times.
✓ Branch 1 taken 92 times.
✗ Branch 2 not taken.
384 for (uint32_t i = 0; i < res_n; i++)
634 {
635 292 out_leaves[i] = res_id[i];
636
1/2
✓ Branch 0 taken 292 times.
✗ Branch 1 not taken.
292 if (out_dists != NULL)
637 292 out_dists[i] = res_d[i];
638 }
639 92 return res_n;
640 }
641