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