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