GCC Code Coverage Report


Directory: src/
File: src/algo/kmeans.c
Date: 2026-09-30 11:11:31
Exec Total Coverage
Lines: 406 450 90.2%
Functions: 31 33 93.9%
Branches: 242 438 55.3%

Line Branch Exec Source
1 /*
2 * Copyright (c) 2026 Tiger Data, Inc.
3 * Licensed under the PostgreSQL License. See LICENSE for details.
4 *
5 * kmeans.c - K-means clustering orchestration
6 *
7 * Common infrastructure shared by all k-means variants:
8 * - k-means++ initialization (D²-weighted sampling)
9 * - Centroid update step (mean of assigned vectors)
10 * - Empty cluster handling (largest-cluster splitting)
11 * - Convergence checking (max centroid shift)
12 * - State management and public API
13 *
14 * Algorithm-specific assignment steps live in separate files:
15 * - kmeans_lloyd.c: brute-force (CBLAS sgemm or builtin batch)
16 * - kmeans_hamerly.c: Hamerly's single-bound acceleration
17 * - kmeans_elkan.c: Elkan's per-centroid bounds
18 */
19
20 #include "vs_config.h"
21
22 #include <float.h>
23 #include <math.h>
24 #include <stdint.h>
25 #include <stdlib.h>
26 #include <string.h>
27
28 #include "algo/kmeans_elkan.h"
29 #include "algo/kmeans_hamerly.h"
30 #include "algo/kmeans_internal.h"
31 #include "algo/kmeans_lloyd.h"
32 #include "algo/vecops.h"
33 #include "core/log.h"
34 #include "core/memory.h"
35
36 /*
37 * Xoshiro256** PRNG (same as matrix.c - duplicated to keep kmeans.c
38 * self-contained without exposing PRNG in a shared header).
39 */
40 typedef struct
41 {
42 uint64_t s[4];
43 } Xoshiro256State;
44
45 static inline uint64_t
46 22751 xo_rotl(uint64_t x, int k)
47 {
48 22751 return (x << k) | (x >> (64 - k));
49 }
50
51 static uint64_t
52 16821 xo_next(Xoshiro256State *state)
53 {
54 16821 uint64_t *s = state->s;
55 16821 uint64_t result = xo_rotl(s[1] * 5, 7) * 9;
56 16821 uint64_t t = s[1] << 17;
57
58 16821 s[2] ^= s[0];
59 16821 s[3] ^= s[1];
60 16821 s[1] ^= s[2];
61 16821 s[0] ^= s[3];
62 16821 s[2] ^= t;
63 16821 s[3] = xo_rotl(s[3], 45);
64
65 16821 return result;
66 }
67
68 static void
69 2662 xo_seed(Xoshiro256State *state, uint64_t seed)
70 {
71
2/2
✓ Branch 0 taken 10648 times.
✓ Branch 1 taken 2662 times.
13310 for (int i = 0; i < 4; i++)
72 {
73 10648 seed += 0x9e3779b97f4a7c15ULL;
74 10648 uint64_t z = seed;
75 10648 z = (z ^ (z >> 30)) * 0xbf58476d1ce4e5b9ULL;
76 10648 z = (z ^ (z >> 27)) * 0x94d049bb133111ebULL;
77 10648 state->s[i] = z ^ (z >> 31);
78 }
79 2662 }
80
81 /* Random double in [0, 1) */
82 static double
83 16821 xo_uniform(Xoshiro256State *state)
84 {
85 27712 uint64_t x = xo_next(state) >> 11;
86 16821 return (double)x / (double)(1ULL << 53);
87 }
88
89 /* CBLAS runtime toggle (same pattern as matrix.c) */
90 static bool g_use_cblas = true;
91
92 void
93 6 vs_kmeans_set_use_cblas(bool use_cblas)
94 {
95 6 g_use_cblas = use_cblas;
96 6 }
97
98 bool
99 2 vs_kmeans_get_use_cblas(void)
100 {
101 #ifdef VS_HAVE_CBLAS
102 return g_use_cblas;
103 #else
104 2 return false;
105 #endif
106 }
107
108 const char *
109 2 vs_kmeans_impl_name(void)
110 {
111 #ifdef VS_HAVE_CBLAS
112 return g_use_cblas ? "cblas" : "builtin";
113 #else
114 2 return "builtin";
115 #endif
116 }
117
118 static const char *algo_names[] = {
119 [KMEANS_ALGO_AUTO] = "auto",
120 [KMEANS_ALGO_LLOYD] = "lloyd",
121 [KMEANS_ALGO_HAMERLY] = "hamerly",
122 [KMEANS_ALGO_ELKAN] = "elkan",
123 [KMEANS_ALGO_CBLAS] = "lloyd(cblas)",
124 };
125
126 const char *
127 12 vs_kmeans_algo_name(KMeansAlgorithm algo)
128 {
129
2/2
✓ Branch 0 taken 10 times.
✓ Branch 1 taken 2 times.
12 if ((unsigned)algo <= KMEANS_ALGO_CBLAS)
130 10 return algo_names[algo];
131 2 return "unknown";
132 }
133
134 bool
135 2 vs_cblas_is_single_threaded(void)
136 {
137 2 const char *v = getenv("OMP_NUM_THREADS");
138
1/6
✗ Branch 0 not taken.
✓ Branch 1 taken 2 times.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✗ Branch 4 not taken.
✗ Branch 5 not taken.
2 return v != NULL && v[0] == '1' && v[1] == '\0';
139 }
140
141 /*
142 * Pin BLAS to single-threaded operation via the vendor's runtime API.
143 *
144 * Only symbols for the BLAS vendor detected at configure time are
145 * referenced. Declaring all vendors' externs and relying on weak
146 * symbols doesn't survive macOS's linker (Mach-O has no portable
147 * "undefined weak reference resolves to NULL" that holds up under
148 * LTO).
149 *
150 * Supported vendors: OpenBLAS (openblas_set_num_threads),
151 * BLIS/AOCL (bli_thread_set_num_threads). Others (Accelerate,
152 * MKL, ARM PL) must be pinned via env vars.
153 */
154 #ifdef VS_BLAS_OPENBLAS
155 extern void openblas_set_num_threads(int);
156 #endif
157 #ifdef VS_BLAS_BLIS
158 extern void bli_thread_set_num_threads(long);
159 #endif
160
161 void
162 255 vs_cblas_pin_single_thread(void)
163 {
164 #ifdef VS_BLAS_OPENBLAS
165 openblas_set_num_threads(1);
166 #endif
167 #ifdef VS_BLAS_BLIS
168 bli_thread_set_num_threads(1);
169 #endif
170 255 }
171
172 /*
173 * Precompute ||x||^2 for all input vectors.
174 */
175 __attribute__((always_inline)) static inline void
176 precompute_norms_x_impl(KMeansState *st, const Vec32TypeOps *ops)
177 {
178 1104 size_t esz = ops->element_size;
179
4/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✓ Branch 2 taken 2000 times.
✓ Branch 3 taken 12 times.
✓ Branch 4 taken 778198 times.
✓ Branch 5 taken 2552 times.
782762 for (uint32_t i = 0; i < st->nvecs; i++)
180
4/7
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 3 not taken.
✓ Branch 4 taken 2000 times.
✓ Branch 5 taken 298813 times.
✓ Branch 6 taken 395149 times.
✓ Branch 7 taken 84236 times.
1307553 st->norms_x[i] = ops->norm_sq(km_get_vector(st, i, esz), st->dim);
181 1104 }
182
183 /*
184 * Compute distance from a typed vector to a float32 centroid.
185 */
186 __attribute__((always_inline)) static inline float
187 5920914 vector_centroid_distance_impl(
188 DistanceMetric metric,
189 const void *vec,
190 const float *centroid,
191 Dimension dim,
192 const Vec32TypeOps *ops)
193 {
194 8297048 switch (metric)
195 {
196 8009754 case DISTANCE_L2:
197 8009754 return ops->l2_squared(vec, centroid, dim);
198 101600 case DISTANCE_INNER_PRODUCT:
199 101600 return -ops->dot_product(vec, centroid, dim);
200 185694 case DISTANCE_COSINE:
201 185694 return 1.0f - ops->dot_product(vec, centroid, dim);
202 }
203 ✗ return FLT_MAX;
204 }
205
206 /*
207 * k-means++ initialization.
208 */
209 __attribute__((always_inline)) static inline void
210 1512 kmeans_init_plusplus_impl(
211 KMeansState *st, uint64_t seed, const Vec32TypeOps *ops)
212 {
213 1512 Xoshiro256State rng;
214 2662 xo_seed(&rng, seed);
215
216 2662 Dimension dim = st->dim;
217 2662 uint32_t nvecs = st->nvecs;
218 2662 uint32_t nlist = st->nlist;
219 2662 size_t esz = ops->element_size;
220 2662 float *dists = vs_alloc(nvecs * sizeof(float));
221
222 /* 1. First centroid: random vector */
223 2662 uint32_t idx = (uint32_t)(xo_uniform(&rng) * nvecs);
224
2/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✓ Branch 3 taken 12 times.
✗ Branch 4 not taken.
✓ Branch 5 taken 2650 times.
2662 if (idx >= nvecs)
225 ✗ idx = nvecs - 1;
226
4/8
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✓ Branch 4 taken 1271 times.
✓ Branch 5 taken 253 times.
✓ Branch 6 taken 916 times.
✓ Branch 7 taken 222 times.
3812 ops->to_float_one(km_get_vector(st, idx, esz), st->centroids, dim);
227
228 /* Initialize distances to first centroid */
229
6/8
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✓ Branch 2 taken 2000 times.
✓ Branch 3 taken 12 times.
✓ Branch 4 taken 243902 times.
✓ Branch 5 taken 1138 times.
✓ Branch 7 taken 567676 times.
✓ Branch 8 taken 1512 times.
816240 for (uint32_t i = 0; i < nvecs; i++)
230
4/12
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✓ Branch 4 taken 2000 times.
✗ Branch 5 not taken.
✗ Branch 6 not taken.
✗ Branch 7 not taken.
✓ Branch 8 taken 778198 times.
✓ Branch 9 taken 15200 times.
✓ Branch 10 taken 18180 times.
✗ Branch 11 not taken.
1627156 dists[i] = vector_centroid_distance_impl(
231 st->metric,
232 km_get_vector(st, i, esz),
233
3/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✓ Branch 3 taken 2000 times.
✓ Branch 4 taken 462339 times.
✓ Branch 5 taken 349239 times.
813578 st->centroids,
234 dim,
235 ops);
236
237 /* 2. Pick remaining centroids */
238
4/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✓ Branch 2 taken 32 times.
✓ Branch 3 taken 12 times.
✓ Branch 4 taken 14127 times.
✓ Branch 5 taken 2650 times.
16821 for (uint32_t k = 1; k < nlist; k++)
239 {
240 /* Compute cumulative distribution */
241 4780 double total = 0.0;
242
4/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✓ Branch 2 taken 5600 times.
✓ Branch 3 taken 32 times.
✓ Branch 4 taken 7477870 times.
✓ Branch 5 taken 14127 times.
7497629 for (uint32_t i = 0; i < nvecs; i++)
243 7483470 total += (double)dists[i];
244
245 /* Sample from distribution */
246 14159 double r = xo_uniform(&rng) * total;
247 14159 double cum = 0.0;
248 14159 uint32_t chosen = 0;
249
2/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✓ Branch 2 taken 4164 times.
✗ Branch 3 not taken.
✓ Branch 4 taken 4809776 times.
✗ Branch 5 not taken.
4813940 for (uint32_t i = 0; i < nvecs; i++)
250 {
251 4813940 cum += (double)dists[i];
252
4/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✓ Branch 2 taken 32 times.
✓ Branch 3 taken 4132 times.
✓ Branch 4 taken 3553364 times.
✓ Branch 5 taken 1256412 times.
4813940 if (cum >= r)
253 {
254 4780 chosen = i;
255 4780 break;
256 }
257 }
258
259 /* Copy chosen vector as new centroid (convert to f32) */
260 23538 ops->to_float_one(
261 km_get_vector(st, chosen, esz),
262
3/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✓ Branch 3 taken 32 times.
✓ Branch 4 taken 10125 times.
✓ Branch 5 taken 4002 times.
14159 st->centroids + (size_t)k * dim,
263 dim);
264
265 /* Update min distances */
266
6/8
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✓ Branch 2 taken 5600 times.
✓ Branch 3 taken 32 times.
✓ Branch 4 taken 2124632 times.
✓ Branch 5 taken 4748 times.
✓ Branch 7 taken 5353238 times.
✓ Branch 8 taken 9379 times.
7497629 for (uint32_t i = 0; i < nvecs; i++)
267 {
268
4/12
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✓ Branch 4 taken 5600 times.
✗ Branch 5 not taken.
✗ Branch 6 not taken.
✗ Branch 7 not taken.
✓ Branch 8 taken 7223956 times.
✓ Branch 9 taken 86400 times.
✓ Branch 10 taken 167514 times.
✗ Branch 11 not taken.
7483470 float d = vector_centroid_distance_impl(
269 st->metric,
270 km_get_vector(st, i, esz),
271
3/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✓ Branch 3 taken 5600 times.
✓ Branch 4 taken 3353869 times.
✓ Branch 5 taken 4124001 times.
7483470 st->centroids + (size_t)k * dim,
272 dim,
273 ops);
274
4/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✓ Branch 2 taken 1920 times.
✓ Branch 3 taken 3680 times.
✓ Branch 4 taken 1242705 times.
✓ Branch 5 taken 6235165 times.
7483470 if (d < dists[i])
275 1244625 dists[i] = d;
276 }
277 }
278
279 2662 vs_free(dists);
280 2662 }
281
282 /*
283 * Update step: recompute centroids as mean of assigned vectors.
284 */
285 __attribute__((always_inline)) static inline void
286 ✗ kmeans_update_centroids_impl(KMeansState *st, const Vec32TypeOps *ops)
287 {
288 208 uint32_t nlist = st->nlist;
289 208 Dimension dim = st->dim;
290 208 size_t esz = ops->element_size;
291
292 /* Zero accumulators */
293 208 memset(st->new_centroids, 0, (size_t)nlist * dim * sizeof(float));
294 208 memset(st->cluster_sizes, 0, nlist * sizeof(uint32_t));
295
296 /* Accumulate */
297
4/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✓ Branch 2 taken 13600 times.
✓ Branch 3 taken 72 times.
✓ Branch 4 taken 22800 times.
✓ Branch 5 taken 136 times.
36608 for (uint32_t i = 0; i < st->nvecs; i++)
298 {
299 36400 ClusterId c = st->assignments[i];
300
2/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✓ Branch 3 taken 13600 times.
✗ Branch 4 not taken.
✓ Branch 5 taken 22800 times.
36400 st->cluster_sizes[c]++;
301
0/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✗ Branch 4 not taken.
✗ Branch 5 not taken.
36400 const void *vec = km_get_vector(st, i, esz);
302 36400 float *cent = st->new_centroids + (size_t)c * dim;
303 36400 ops->sum_to_float(vec, cent, dim);
304 }
305
306 /* Divide by cluster size */
307
4/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✓ Branch 2 taken 328 times.
✓ Branch 3 taken 72 times.
✓ Branch 4 taken 560 times.
✓ Branch 5 taken 136 times.
1096 for (uint32_t j = 0; j < nlist; j++)
308 {
309
2/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✓ Branch 3 taken 328 times.
✗ Branch 4 not taken.
✓ Branch 5 taken 560 times.
888 if (st->cluster_sizes[j] == 0)
310 ✗ continue;
311
312 888 float inv_size = 1.0f / (float)st->cluster_sizes[j];
313 888 vec32_scale(
314 888 st->new_centroids + (size_t)j * dim,
315 inv_size,
316 888 st->new_centroids + (size_t)j * dim,
317 dim);
318 }
319
320 /* For cosine: normalize centroids to unit length */
321
2/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✓ Branch 3 taken 72 times.
✗ Branch 4 not taken.
✓ Branch 5 taken 136 times.
208 if (st->metric == DISTANCE_COSINE)
322 {
323 ✗ for (uint32_t j = 0; j < nlist; j++)
324 {
325 ✗ if (st->cluster_sizes[j] == 0)
326 ✗ continue;
327 ✗ float *cent = st->new_centroids + (size_t)j * dim;
328 ✗ float norm = vs_l2_norm(cent, dim);
329 ✗ if (norm > 1e-10f)
330 ✗ vec32_scale(cent, 1.0f / norm, cent, dim);
331 }
332 }
333
334 /* Swap new centroids into place */
335 208 float *tmp = st->centroids;
336 208 st->centroids = st->new_centroids;
337 208 st->new_centroids = tmp;
338 208 }
339
340 /*
341 * Handle empty clusters by splitting the largest cluster.
342 *
343 * FAISS approach: replace empty centroid with a perturbation of the
344 * largest cluster's centroid.
345 */
346 static void
347 2814 kmeans_handle_empty_clusters(KMeansState *st)
348 {
349 2814 uint32_t dim = st->dim;
350 2814 uint32_t nlist = st->nlist;
351
352
2/2
✓ Branch 0 taken 17505 times.
✓ Branch 1 taken 2814 times.
20319 for (uint32_t j = 0; j < nlist; j++)
353 {
354
2/2
✓ Branch 0 taken 17375 times.
✓ Branch 1 taken 130 times.
17505 if (st->cluster_sizes[j] > 0)
355 17375 continue;
356
357 /* Find largest cluster */
358 130 uint32_t largest = 0;
359 130 uint32_t largest_size = st->cluster_sizes[0];
360
2/2
✓ Branch 0 taken 2076 times.
✓ Branch 1 taken 130 times.
2206 for (uint32_t k = 1; k < nlist; k++)
361 {
362
2/2
✓ Branch 0 taken 106 times.
✓ Branch 1 taken 1970 times.
2076 if (st->cluster_sizes[k] > largest_size)
363 {
364 106 largest = k;
365 106 largest_size = st->cluster_sizes[k];
366 }
367 }
368
369
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 130 times.
130 if (largest_size <= 1)
370 ✗ continue; /* Can't split a cluster of size 1 */
371
372 /* Perturb: empty = largest * (1 + eps), largest *= (1 - eps) */
373 130 float *c_empty = st->centroids + (size_t)j * dim;
374 130 float *c_largest = st->centroids + (size_t)largest * dim;
375
376
2/2
✓ Branch 0 taken 2390 times.
✓ Branch 1 taken 130 times.
2520 for (uint32_t d = 0; d < dim; d++)
377 {
378 2390 float val = c_largest[d];
379 2390 c_empty[d] = val * (1.0f + 1e-4f);
380 2390 c_largest[d] = val * (1.0f - 1e-4f);
381 }
382
383 /* Split size estimate (roughly half each) */
384 130 uint32_t half = largest_size / 2;
385 130 st->cluster_sizes[j] = half;
386 130 st->cluster_sizes[largest] -= half;
387 }
388 2814 }
389
390 /*
391 * Max centroid movement between two centroid arrays.
392 * Returns the maximum squared L2 shift.
393 */
394 float
395 19533 kmeans_max_centroid_shift_between(
396 const float *a, const float *b, uint32_t nlist, Dimension dim)
397 {
398 19533 float max_shift = 0.0f;
399
400
2/2
✓ Branch 0 taken 140812 times.
✓ Branch 1 taken 19533 times.
160345 for (uint32_t j = 0; j < nlist; j++)
401 {
402 240419 float shift = vs_l2_distance_squared(
403 140812 a + (size_t)j * dim, b + (size_t)j * dim, dim);
404
2/2
✓ Branch 0 taken 36930 times.
✓ Branch 1 taken 103882 times.
140812 if (shift > max_shift)
405 36930 max_shift = shift;
406 }
407
408 19533 return max_shift;
409 }
410
411 VS_TARGET_CLONES void
412 3333 kmeans_assign_accumulate(
413 const float *vectors,
414 const uint32_t *indices,
415 uint32_t start,
416 uint32_t end,
417 const float *centroids,
418 const float *norms_c,
419 uint32_t k,
420 Dimension dim,
421 DistanceMetric metric,
422 const uint32_t *filter,
423 uint32_t filter_val,
424 float *out_sums,
425 uint32_t *out_cnts,
426 float *out_cost)
427 {
428 3333 float cost = 0.0f;
429
430
2/2
✓ Branch 0 taken 3161981 times.
✓ Branch 1 taken 3333 times.
3165314 for (uint32_t i = start; i < end; i++)
431 {
432
1/4
✗ Branch 0 not taken.
✓ Branch 1 taken 3161981 times.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
3161981 if (filter != NULL && filter[i] != filter_val)
433 ✗ continue;
434
435
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 3161981 times.
3161981 uint32_t idx = indices ? indices[i] : i;
436 3161981 const float *vec = vectors + (size_t)idx * dim;
437 3161981 float best_d = __FLT_MAX__;
438 3161981 uint32_t best_c = 0;
439
440
2/2
✓ Branch 0 taken 31586908 times.
✓ Branch 1 taken 3161981 times.
34748889 for (uint32_t c = 0; c < k; c++)
441 {
442 31586908 const float *cent = centroids + (size_t)c * dim;
443 15425024 float d;
444
445
3/3
✓ Branch 0 taken 30622428 times.
✓ Branch 1 taken 360000 times.
✓ Branch 2 taken 604480 times.
31586908 switch (metric)
446 {
447 30622428 default: /* L2 */
448 {
449 30622428 float nx = vs_l2_norm_squared(vec, dim);
450 30622428 float dot = vs_dot_product(vec, cent, dim);
451 30622428 d = nx + norms_c[c] - 2.0f * dot;
452
2/2
✓ Branch 0 taken 105 times.
✓ Branch 1 taken 30622323 times.
30622428 if (d < 0.0f)
453 105 d = 0.0f;
454 15740284 break;
455 }
456 360000 case DISTANCE_INNER_PRODUCT:
457 360000 d = -vs_dot_product(vec, cent, dim);
458 360000 break;
459 604480 case DISTANCE_COSINE:
460 604480 d = 1.0f - vs_dot_product(vec, cent, dim);
461 604480 break;
462 }
463
464
2/2
✓ Branch 0 taken 9333959 times.
✓ Branch 1 taken 22252949 times.
31586908 if (d < best_d)
465 {
466 9333959 best_d = d;
467 9333959 best_c = c;
468 }
469 }
470
471 3161981 cost += best_d;
472 3161981 out_cnts[best_c]++;
473 3161981 float *sum = out_sums + (size_t)best_c * dim;
474
2/2
✓ Branch 0 taken 113202685 times.
✓ Branch 1 taken 3161981 times.
116364666 for (uint32_t d = 0; d < dim; d++)
475 113202685 sum[d] += vec[d];
476 }
477
478 3333 *out_cost += cost;
479 3333 }
480
481 VS_TARGET_CLONES void
482 318 kmeans_assign(
483 const float *vectors,
484 uint32_t start,
485 uint32_t end,
486 const float *centroids,
487 const float *norms_c,
488 uint32_t k,
489 Dimension dim,
490 DistanceMetric metric,
491 uint32_t *out_assignments)
492 {
493
2/2
✓ Branch 0 taken 215527 times.
✓ Branch 1 taken 318 times.
215845 for (uint32_t i = start; i < end; i++)
494 {
495 215527 const float *vec = vectors + (size_t)i * dim;
496 215527 float best_d = __FLT_MAX__;
497 215527 uint32_t best_c = 0;
498
499
2/2
✓ Branch 0 taken 1703126 times.
✓ Branch 1 taken 215527 times.
1918653 for (uint32_t c = 0; c < k; c++)
500 {
501 1703126 const float *cent = centroids + (size_t)c * dim;
502 838108 float d;
503
504
3/3
✓ Branch 0 taken 1653860 times.
✓ Branch 1 taken 18000 times.
✓ Branch 2 taken 31266 times.
1703126 switch (metric)
505 {
506 1653860 default: /* L2 */
507 {
508 1653860 float nx = vs_l2_norm_squared(vec, dim);
509 1653860 float dot = vs_dot_product(vec, cent, dim);
510 1653860 d = nx + norms_c[c] - 2.0f * dot;
511
2/2
✓ Branch 0 taken 4 times.
✓ Branch 1 taken 1653856 times.
1653860 if (d < 0.0f)
512 4 d = 0.0f;
513 843218 break;
514 }
515 18000 case DISTANCE_INNER_PRODUCT:
516 18000 d = -vs_dot_product(vec, cent, dim);
517 18000 break;
518 31266 case DISTANCE_COSINE:
519 31266 d = 1.0f - vs_dot_product(vec, cent, dim);
520 31266 break;
521 }
522
523
2/2
✓ Branch 0 taken 540848 times.
✓ Branch 1 taken 1162278 times.
1703126 if (d < best_d)
524 {
525 540848 best_d = d;
526 540848 best_c = c;
527 }
528 }
529
530 215527 out_assignments[i] = best_c;
531 }
532 318 }
533
534 float
535 1174 kmeans_merge_centroids(
536 float *centroids,
537 float *norms_c,
538 const float *old_cents,
539 const float *const *worker_sums,
540 const uint32_t *const *worker_cnts,
541 const float *worker_costs,
542 uint32_t nworkers,
543 uint32_t nlist,
544 Dimension dim,
545 DistanceMetric metric,
546 float *out_total_cost)
547 {
548 /* Sum costs */
549 1174 float total_cost = 0.0f;
550
2/2
✓ Branch 0 taken 3333 times.
✓ Branch 1 taken 1174 times.
4507 for (uint32_t t = 0; t < nworkers; t++)
551 3333 total_cost += worker_costs[t];
552 1174 *out_total_cost = total_cost;
553
554 /* Merge per-worker accumulators */
555 1174 float *new_cents = centroids;
556 1174 memset(new_cents, 0, (size_t)nlist * dim * sizeof(float));
557
558 1174 uint32_t *sizes = (uint32_t *)alloca(nlist * sizeof(uint32_t));
559 1174 memset(sizes, 0, nlist * sizeof(uint32_t));
560
561
2/2
✓ Branch 0 taken 3333 times.
✓ Branch 1 taken 1174 times.
4507 for (uint32_t t = 0; t < nworkers; t++)
562 {
563 3333 const float *sums = worker_sums[t];
564 3333 const uint32_t *cnts = worker_cnts[t];
565
566
2/2
✓ Branch 0 taken 30368 times.
✓ Branch 1 taken 3333 times.
33701 for (uint32_t c = 0; c < nlist; c++)
567 {
568 30368 sizes[c] += cnts[c];
569 30368 float *dst = new_cents + (size_t)c * dim;
570 30368 const float *src = sums + (size_t)c * dim;
571
2/2
✓ Branch 0 taken 2873356 times.
✓ Branch 1 taken 30368 times.
2903724 for (uint32_t d = 0; d < dim; d++)
572 2873356 dst[d] += src[d];
573 }
574 }
575
576 /* Divide by cluster size to get mean */
577
2/2
✓ Branch 0 taken 10168 times.
✓ Branch 1 taken 1174 times.
11342 for (uint32_t c = 0; c < nlist; c++)
578 {
579
2/2
✓ Branch 0 taken 7 times.
✓ Branch 1 taken 10161 times.
10168 if (sizes[c] == 0)
580 7 continue;
581 10161 float inv = 1.0f / (float)sizes[c];
582 10161 float *cent = new_cents + (size_t)c * dim;
583
2/2
✓ Branch 0 taken 956499 times.
✓ Branch 1 taken 10161 times.
966660 for (uint32_t d = 0; d < dim; d++)
584 956499 cent[d] *= inv;
585 }
586
587 /* Normalize for cosine metric */
588
2/2
✓ Branch 0 taken 103 times.
✓ Branch 1 taken 1071 times.
1174 if (metric == DISTANCE_COSINE)
589 {
590
2/2
✓ Branch 0 taken 980 times.
✓ Branch 1 taken 103 times.
1083 for (uint32_t c = 0; c < nlist; c++)
591 {
592
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 980 times.
980 if (sizes[c] == 0)
593 ✗ continue;
594 980 float *cent = new_cents + (size_t)c * dim;
595 980 float norm = vs_l2_norm(cent, dim);
596
1/2
✓ Branch 0 taken 980 times.
✗ Branch 1 not taken.
980 if (norm > 1e-10f)
597 980 vec32_scale(cent, 1.0f / norm, cent, dim);
598 }
599 }
600
601 /* Convergence: max centroid shift (squared) */
602 1174 float shift_sq = kmeans_max_centroid_shift_between(
603 centroids, old_cents, nlist, dim);
604
605 /* Precompute centroid norms for next iteration */
606
4/4
✓ Branch 0 taken 1115 times.
✓ Branch 1 taken 59 times.
✓ Branch 2 taken 596 times.
✓ Branch 3 taken 84 times.
1174 if (norms_c != NULL && metric == DISTANCE_L2)
607 {
608
2/2
✓ Branch 0 taken 9068 times.
✓ Branch 1 taken 1031 times.
10099 for (uint32_t j = 0; j < nlist; j++)
609 9068 norms_c[j] = vs_l2_norm_squared(centroids + (size_t)j * dim, dim);
610 }
611
612 1174 return shift_sq;
613 }
614
615 /*
616 * Allocate working state for one k-means run.
617 *
618 * Creates an arena context and allocates everything (including the
619 * struct itself) within it. kmeans_state_destroy() bulk-frees all
620 * memory by deleting the arena.
621 */
622 static KMeansState *
623 2662 kmeans_state_create(
624 const void *vectors,
625 const uint32_t *indices,
626 VecType vec_type,
627 uint32_t nvecs,
628 Dimension dim,
629 uint32_t nlist,
630 DistanceMetric metric)
631 {
632 2662 VsMemCtx ctx = vs_memctx_create(NULL, "kmeans_state");
633 2662 VsMemCtx old_ctx = vs_memctx_switch(ctx);
634
635 2662 KMeansState *st = vs_alloc0(sizeof(KMeansState));
636
637 2662 st->memctx = ctx;
638 2662 st->vectors = vectors;
639 2662 st->indices = indices;
640 2662 st->vec_type = vec_type;
641 2662 st->nvecs = nvecs;
642 2662 st->nlist = nlist;
643 2662 st->dim = dim;
644 2662 st->metric = metric;
645
646 2662 st->centroids = vs_alloc((size_t)nlist * dim * sizeof(float));
647 2662 st->assignments = vs_alloc(nvecs * sizeof(ClusterId));
648 2662 st->cluster_sizes = vs_alloc(nlist * sizeof(uint32_t));
649 2662 st->new_centroids = vs_alloc((size_t)nlist * dim * sizeof(float));
650
651 2662 uint32_t block = KMEANS_BLOCK_SIZE;
652
1/2
✓ Branch 0 taken 1150 times.
✗ Branch 1 not taken.
2662 if (block > nvecs)
653 1150 block = nvecs;
654 2662 st->dist_block = vs_alloc((size_t)block * nlist * sizeof(float));
655
656 /*
657 * Allocate conversion buffer for non-f32 types, or for f32 with
658 * indices (Lloyd gather needs a contiguous block).
659 */
660
5/5
✓ Branch 0 taken 1512 times.
✓ Branch 1 taken 1138 times.
✓ Branch 2 taken 1271 times.
✓ Branch 3 taken 1169 times.
✓ Branch 4 taken 222 times.
2662 if (vs_vec_element_size(vec_type) != sizeof(float) || indices != NULL)
661 2187 st->vec_block = vs_alloc((size_t)block * dim * sizeof(float));
662
663
2/2
✓ Branch 0 taken 2564 times.
✓ Branch 1 taken 98 times.
2662 if (metric == DISTANCE_L2)
664 {
665 2564 st->norms_x = vs_alloc(nvecs * sizeof(float));
666 2564 st->norms_c = vs_alloc(nlist * sizeof(float));
667 }
668
669 2662 vs_memctx_switch(old_ctx);
670 2662 return st;
671 }
672
673 static void
674 2662 kmeans_state_destroy(KMeansState *st)
675 {
676
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 1150 times.
2662 if (st == NULL)
677 ✗ return;
678 2662 VsMemCtx ctx = (VsMemCtx)st->memctx;
679 2662 vs_memctx_delete(ctx); /* st is now invalid */
680 }
681
682 /*
683 * Algorithm vtable for the k-means iteration loop.
684 *
685 * Each variant provides: create/destroy for per-algorithm state,
686 * assign for the assignment step, and optionally update_bounds
687 * for bound-accelerated algorithms (Hamerly, Elkan).
688 */
689 typedef struct
690 {
691 void *(*create)(const KMeansState *st);
692 void (*destroy)(void *algo_state);
693 void (*assign)(KMeansState *st, void *algo_state);
694 void (*update_bounds)(
695 KMeansState *st, void *algo_state, const float *old_cents);
696 } KMeansAlgoOps;
697
698 /* Lloyd wrappers (no per-algorithm state) */
699
700 static void
701 12 lloyd_plain_assign(KMeansState *st, void *state)
702 {
703 ✗ (void)state;
704 12 lloyd_assign(st, false);
705 12 }
706
707 static void
708 ✗ lloyd_cblas_assign(KMeansState *st, void *state)
709 {
710 ✗ (void)state;
711 ✗ lloyd_assign(st, true);
712 ✗ }
713
714 static const KMeansAlgoOps lloyd_ops = {
715 .assign = lloyd_plain_assign,
716 };
717
718 static const KMeansAlgoOps lloyd_cblas_ops = {
719 .assign = lloyd_cblas_assign,
720 };
721
722 /* Hamerly wrappers */
723
724 static void *
725 26 hamerly_wrap_create(const KMeansState *st)
726 {
727 26 return hamerly_create(st->nvecs, st->nlist, st->dim);
728 }
729
730 static void
731 126 hamerly_wrap_assign(KMeansState *st, void *state)
732 {
733 126 hamerly_assign(st, state);
734 126 }
735
736 static void
737 100 hamerly_wrap_bounds(KMeansState *st, void *state, const float *old_cents)
738 {
739 100 hamerly_update_bounds(st, state, old_cents);
740 100 }
741
742 static void
743 26 hamerly_wrap_destroy(void *state)
744 {
745 26 hamerly_destroy(state);
746 26 }
747
748 static const KMeansAlgoOps hamerly_ops = {
749 .create = hamerly_wrap_create,
750 .destroy = hamerly_wrap_destroy,
751 .assign = hamerly_wrap_assign,
752 .update_bounds = hamerly_wrap_bounds,
753 };
754
755 /* Elkan wrappers */
756
757 static void *
758 26 elkan_wrap_create(const KMeansState *st)
759 {
760 26 return elkan_create(st->nvecs, st->nlist, st->dim);
761 }
762
763 static void
764 126 elkan_wrap_assign(KMeansState *st, void *state)
765 {
766 126 elkan_assign(st, state);
767 126 }
768
769 static void
770 100 elkan_wrap_bounds(KMeansState *st, void *state, const float *old_cents)
771 {
772 100 elkan_update_bounds(st, state, old_cents);
773 100 }
774
775 static void
776 26 elkan_wrap_destroy(void *state)
777 {
778 26 elkan_destroy(state);
779 26 }
780
781 static const KMeansAlgoOps elkan_ops = {
782 .create = elkan_wrap_create,
783 .destroy = elkan_wrap_destroy,
784 .assign = elkan_wrap_assign,
785 .update_bounds = elkan_wrap_bounds,
786 };
787
788 /*
789 * Run one complete k-means attempt (init + iterate to convergence).
790 *
791 * always_inline — the specialized wrappers below pass a static const
792 * Vec32TypeOps from the header, so the compiler inlines through
793 * every vtable function pointer. VS_TARGET_CLONES on the wrappers
794 * generates AVX2/AVX-512 variants of the entire inlined body.
795 */
796 __attribute__((always_inline)) static inline void
797 1512 kmeans_run_one_impl(
798 KMeansState *st,
799 const KMeansOptions *opts,
800 uint64_t seed,
801 const KMeansAlgoOps *algo,
802 const Vec32TypeOps *ops)
803 {
804 /* Precompute input vector norms for L2 */
805 2662 if (st->metric == DISTANCE_L2)
806 1512 precompute_norms_x_impl(st, ops);
807
808 2662 size_t cent_bytes = (size_t)st->nlist * st->dim * sizeof(float);
809
4/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✓ Branch 3 taken 8 times.
✓ Branch 4 taken 4 times.
✓ Branch 6 taken 44 times.
✓ Branch 7 taken 2606 times.
2662 void *algo_state = algo->create ? algo->create(st) : NULL;
810
4/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✓ Branch 3 taken 8 times.
✓ Branch 4 taken 4 times.
✓ Branch 6 taken 44 times.
✓ Branch 7 taken 2606 times.
2662 float *old_cents = algo->update_bounds ? vs_alloc(cent_bytes) : NULL;
811
812
2/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✓ Branch 3 taken 12 times.
✗ Branch 4 not taken.
✓ Branch 5 taken 2650 times.
2662 if (opts->initial_centroids != NULL)
813 {
814 ✗ memcpy(st->centroids,
815 ✗ opts->initial_centroids,
816 ✗ (size_t)st->nlist * st->dim * sizeof(float));
817 }
818 else
819 {
820 3024 kmeans_init_plusplus_impl(st, seed, ops);
821 }
822
823 /* Use fused iterate path for Lloyd with f32 vectors.
824 * Works with or without a thread pool — serial fallback
825 * calls work+reduce in a plain loop.
826 * Hamerly/Elkan use the generic loop below (needs bounds). */
827
3/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✓ Branch 3 taken 12 times.
✓ Branch 4 taken 1138 times.
✓ Branch 5 taken 1512 times.
3800 bool use_iterate = st->vec_type == VS_VEC_F32 &&
828
2/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✓ Branch 4 taken 1094 times.
✓ Branch 5 taken 1556 times.
2650 algo->update_bounds == NULL;
829
830
3/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✓ Branch 3 taken 12 times.
✓ Branch 4 taken 1094 times.
✓ Branch 5 taken 44 times.
2662 if (use_iterate)
831 {
832 2606 bool use_cblas = (algo == &lloyd_cblas_ops);
833 2606 lloyd_iterate(st, use_cblas, opts);
834 2606 kmeans_handle_empty_clusters(st);
835
1/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✗ Branch 4 not taken.
✓ Branch 5 taken 2606 times.
2606 if (algo->destroy)
836 ✗ algo->destroy(algo_state);
837
1/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✗ Branch 4 not taken.
✓ Branch 5 taken 2606 times.
2606 if (old_cents)
838 ✗ vs_free(old_cents);
839 1094 return;
840 }
841
842
2/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✓ Branch 2 taken 72 times.
✗ Branch 3 not taken.
✓ Branch 4 taken 136 times.
✗ Branch 5 not taken.
208 for (uint32_t iter = 0; iter < opts->max_iterations; iter++)
843 {
844 /* Save old centroids for convergence check */
845
0/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✗ Branch 4 not taken.
✗ Branch 5 not taken.
208 memcpy(st->new_centroids, st->centroids, cent_bytes);
846
847 /* Save separately for bounds update (Hamerly/Elkan) */
848
3/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✓ Branch 2 taken 64 times.
✓ Branch 3 taken 8 times.
✓ Branch 4 taken 136 times.
✗ Branch 5 not taken.
208 if (old_cents)
849 200 memcpy(old_cents, st->centroids, cent_bytes);
850
851 208 algo->assign(st, algo_state);
852 ✗ kmeans_update_centroids_impl(st, ops);
853 208 kmeans_handle_empty_clusters(st);
854
855
3/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✓ Branch 2 taken 64 times.
✓ Branch 3 taken 8 times.
✓ Branch 4 taken 136 times.
✗ Branch 5 not taken.
208 if (algo->update_bounds)
856 200 algo->update_bounds(st, algo_state, old_cents);
857
858 208 float shift_sq = kmeans_max_centroid_shift_between(
859 208 st->centroids, st->new_centroids, st->nlist, st->dim);
860
861
3/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✓ Branch 3 taken 72 times.
✓ Branch 4 taken 16 times.
✓ Branch 5 taken 120 times.
208 if (opts->verbose)
862
0/6
✗ Branch 1 not taken.
✗ Branch 2 not taken.
✗ Branch 6 not taken.
✗ Branch 7 not taken.
✗ Branch 11 not taken.
✗ Branch 12 not taken.
16 vs_log(" iter %u: cost=%.4f, max_shift=%.6f\n",
863 iter,
864 st->total_cost,
865 sqrtf(shift_sq));
866
867 208 st->total_cost = 0.0f;
868
869 208 float tol_sq = opts->tolerance * opts->tolerance;
870
4/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✓ Branch 2 taken 12 times.
✓ Branch 3 taken 60 times.
✓ Branch 4 taken 44 times.
✓ Branch 5 taken 92 times.
208 if (shift_sq < tol_sq)
871 {
872 56 algo->assign(st, algo_state);
873
3/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✓ Branch 2 taken 8 times.
✓ Branch 3 taken 4 times.
✓ Branch 4 taken 44 times.
✗ Branch 5 not taken.
56 if (algo->destroy)
874 52 algo->destroy(algo_state);
875
3/6
✗ Branch 0 not taken.
✗ Branch 1 not taken.
✓ Branch 2 taken 8 times.
✓ Branch 3 taken 4 times.
✓ Branch 4 taken 44 times.
✗ Branch 5 not taken.
56 if (old_cents)
876 52 vs_free(old_cents);
877 56 return;
878 }
879 }
880
881 /* Final assignment after max iterations */
882 ✗ algo->assign(st, algo_state);
883 ✗ if (algo->destroy)
884 ✗ algo->destroy(algo_state);
885 ✗ if (old_cents)
886 ✗ vs_free(old_cents);
887 }
888
889 /* Specialized wrappers — VS_TARGET_CLONES generates SIMD variants */
890
891 VS_TARGET_CLONES static void
892
2/2
✓ Branch 0 taken 1092 times.
✓ Branch 1 taken 46 times.
2650 kmeans_run_one_f32(
893 KMeansState *st,
894 const KMeansOptions *opts,
895 uint64_t seed,
896 const KMeansAlgoOps *algo)
897 {
898
2/2
✓ Branch 0 taken 1460 times.
✓ Branch 1 taken 52 times.
1512 kmeans_run_one_impl(st, opts, seed, algo, &vs_f32_type_ops);
899 2650 }
900
901 VS_TARGET_CLONES static void
902
1/2
✓ Branch 0 taken 12 times.
✗ Branch 1 not taken.
12 kmeans_run_one_f16(
903 KMeansState *st,
904 const KMeansOptions *opts,
905 uint64_t seed,
906 const KMeansAlgoOps *algo)
907 {
908 ✗ kmeans_run_one_impl(st, opts, seed, algo, &vs_f16_type_ops);
909 12 }
910
911 #if defined(VS_F16C_SUPPORT) && !defined(VS_SIMD_NONE)
912 VS_TARGET_F16C_AVX2 static void
913 ✗ kmeans_run_one_f16c(
914 KMeansState *st,
915 const KMeansOptions *opts,
916 uint64_t seed,
917 const KMeansAlgoOps *algo)
918 {
919 ✗ kmeans_run_one_impl(st, opts, seed, algo, &vs_f16c_type_ops);
920 ✗ }
921 #endif
922
923 /* Single dispatch point — selects the inline vtable once */
924 static void
925 2662 kmeans_run_one(
926 KMeansState *st,
927 const KMeansOptions *opts,
928 uint64_t seed,
929 const KMeansAlgoOps *algo)
930 {
931
2/3
✓ Branch 0 taken 2650 times.
✗ Branch 1 not taken.
✓ Branch 2 taken 12 times.
2662 switch (st->vec_type)
932 {
933 2650 case VS_VEC_F32:
934 2650 kmeans_run_one_f32(st, opts, seed, algo);
935 2650 break;
936 #if defined(VS_F16C_SUPPORT) && !defined(VS_SIMD_NONE)
937 ✗ case VS_VEC_F16C:
938 ✗ kmeans_run_one_f16c(st, opts, seed, algo);
939 ✗ break;
940 #endif
941 12 default:
942 12 kmeans_run_one_f16(st, opts, seed, algo);
943 12 break;
944 }
945 2662 }
946
947 /*
948 * Public API
949 */
950
951 KMeansResult *
952 2643 vs_kmeans(
953 const void *vectors,
954 const uint32_t *indices,
955 VecType vec_type,
956 uint32_t nvecs,
957 Dimension dim,
958 uint32_t nlist,
959 DistanceMetric metric,
960 const KMeansOptions *options)
961 {
962
8/8
✓ Branch 0 taken 2641 times.
✓ Branch 1 taken 2 times.
✓ Branch 2 taken 2639 times.
✓ Branch 3 taken 2 times.
✓ Branch 4 taken 1134 times.
✓ Branch 5 taken 2 times.
✓ Branch 6 taken 2 times.
✓ Branch 7 taken 1132 times.
2643 if (vectors == NULL || nvecs == 0 || dim == 0 || nlist == 0)
963 8 return NULL;
964
965 /* Clamp nlist to nvecs */
966
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 1132 times.
2635 if (nlist > nvecs)
967 ✗ nlist = nvecs;
968
969 2635 KMeansOptions opts = VS_KMEANS_OPTIONS_DEFAULT;
970
2/2
✓ Branch 0 taken 2633 times.
✓ Branch 1 taken 2 times.
2635 if (options != NULL)
971 2633 opts = *options;
972
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 2635 times.
2635 if (opts.nredo == 0)
973 ✗ opts.nredo = 1;
974
975 2635 KMeansResult *best = NULL;
976 2635 float best_cost = FLT_MAX;
977
978 /* Resolve AUTO to a concrete algorithm */
979 2635 KMeansAlgorithm algo = opts.algorithm;
980
2/2
✓ Branch 0 taken 556 times.
✓ Branch 1 taken 576 times.
2635 if (algo == KMEANS_ALGO_AUTO)
981 {
982 #ifdef VS_HAVE_CBLAS
983 algo = g_use_cblas ? KMEANS_ALGO_CBLAS : KMEANS_ALGO_LLOYD;
984 #else
985 556 algo = KMEANS_ALGO_LLOYD;
986 #endif
987 }
988
989 /* Hamerly/Elkan only support L2 — fall back to Lloyd for others */
990
6/6
✓ Branch 0 taken 1110 times.
✓ Branch 1 taken 1525 times.
✓ Branch 2 taken 22 times.
✓ Branch 3 taken 1088 times.
✓ Branch 4 taken 4 times.
✓ Branch 5 taken 40 times.
2635 if ((algo == KMEANS_ALGO_HAMERLY || algo == KMEANS_ALGO_ELKAN) &&
991 metric != DISTANCE_L2)
992 4 algo = KMEANS_ALGO_LLOYD;
993
994 /* Select algorithm vtable */
995 1503 const KMeansAlgoOps *algo_ops;
996
4/4
✓ Branch 0 taken 20 times.
✓ Branch 1 taken 20 times.
✓ Branch 2 taken 1503 times.
✓ Branch 3 taken 1092 times.
2635 switch (algo)
997 {
998 20 case KMEANS_ALGO_HAMERLY:
999 20 algo_ops = &hamerly_ops;
1000 20 break;
1001 20 case KMEANS_ALGO_ELKAN:
1002 20 algo_ops = &elkan_ops;
1003 20 break;
1004 ✗ case KMEANS_ALGO_CBLAS:
1005 ✗ algo_ops = &lloyd_cblas_ops;
1006 ✗ break;
1007 2595 case KMEANS_ALGO_LLOYD:
1008 default:
1009 2595 algo_ops = &lloyd_ops;
1010 2595 break;
1011 }
1012
1013 2635 size_t cent_sz = (size_t)nlist * dim * sizeof(float);
1014 2635 size_t assign_sz = nvecs * sizeof(ClusterId);
1015 2635 size_t clsize_sz = nlist * sizeof(uint32_t);
1016
1017
2/2
✓ Branch 0 taken 2662 times.
✓ Branch 1 taken 2635 times.
5297 for (uint32_t redo = 0; redo < opts.nredo; redo++)
1018 {
1019 2662 KMeansState *st = kmeans_state_create(
1020 vectors, indices, vec_type, nvecs, dim, nlist, metric);
1021
1022 2662 uint64_t seed = opts.seed + redo;
1023
1024 /* Run in arena context so per-iteration temps land there */
1025 2662 VsMemCtx run_ctx = vs_memctx_switch((VsMemCtx)st->memctx);
1026 2662 kmeans_run_one(st, &opts, seed, algo_ops);
1027
1/2
✗ Branch 0 not taken.
✓ Branch 1 taken 1512 times.
2662 vs_memctx_switch(run_ctx);
1028
1029
3/4
✓ Branch 0 taken 12 times.
✓ Branch 1 taken 2650 times.
✓ Branch 2 taken 12 times.
✗ Branch 3 not taken.
2662 if (opts.verbose && opts.nredo > 1)
1030
2/5
✓ Branch 0 taken 10 times.
✓ Branch 1 taken 2 times.
✗ Branch 2 not taken.
✗ Branch 3 not taken.
✗ Branch 4 not taken.
12 vs_log("redo %u/%u: cost=%.4f%s\n",
1031 redo + 1,
1032 opts.nredo,
1033 st->total_cost,
1034 st->total_cost < best_cost ? " (best)" : "");
1035
1036
2/2
✓ Branch 0 taken 2651 times.
✓ Branch 1 taken 11 times.
2662 if (st->total_cost < best_cost)
1037 {
1038 /* Save as best — copy out of arena into caller ctx */
1039
2/2
✓ Branch 0 taken 16 times.
✓ Branch 1 taken 2635 times.
2651 if (best != NULL)
1040 16 vs_kmeans_result_destroy(best);
1041
1042 2651 best = vs_alloc0(sizeof(KMeansResult));
1043 2651 best->nlist = nlist;
1044 2651 best->dim = dim;
1045 2651 best->total_cost = st->total_cost;
1046
1047 2651 best->centroids = vs_alloc(cent_sz);
1048 2651 memcpy(best->centroids, st->centroids, cent_sz);
1049
1050 2651 best->assignments = vs_alloc(assign_sz);
1051 2651 memcpy(best->assignments, st->assignments, assign_sz);
1052
1053 2651 best->cluster_sizes = vs_alloc(clsize_sz);
1054 2651 memcpy(best->cluster_sizes, st->cluster_sizes, clsize_sz);
1055
1056 2651 best_cost = st->total_cost;
1057 }
1058
1059 2662 kmeans_state_destroy(st);
1060 }
1061
1062 1132 return best;
1063 }
1064
1065 void
1066 2651 vs_kmeans_result_destroy(KMeansResult *result)
1067 {
1068
2/2
✓ Branch 0 taken 1507 times.
✓ Branch 1 taken 1144 times.
2651 if (result == NULL)
1069 ✗ return;
1070 2651 vs_free(result->centroids);
1071 2651 vs_free(result->assignments);
1072 2651 vs_free(result->cluster_sizes);
1073 2651 vs_free(result);
1074 }
1075