| 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_lloyd.c - Lloyd's k-means assignment step | ||
| 6 | * | ||
| 7 | * Brute-force assignment: compute the full N×K distance matrix each | ||
| 8 | * iteration. Two backends share the same block-processing structure: | ||
| 9 | * | ||
| 10 | * CBLAS path: | ||
| 11 | * Uses sgemm for the [block × dim] × [dim × nlist] dot product | ||
| 12 | * matrix. BLAS cache-tiling makes this 10-50x faster than naive | ||
| 13 | * loops for large K. | ||
| 14 | * | ||
| 15 | * Builtin path: | ||
| 16 | * Batch dot-product kernel with VS_TARGET_CLONES for AVX-512/AVX2 | ||
| 17 | * auto-vectorization. FMA generation requires -ffp-contract=fast | ||
| 18 | * (set in meson.build). | ||
| 19 | * | ||
| 20 | * Both use the decomposition: | ||
| 21 | * ||x - c||² = ||x||² + ||c||² - 2⟨x,c⟩ | ||
| 22 | * The ⟨x,c⟩ term is computed in batch, norms are precomputed. | ||
| 23 | */ | ||
| 24 | |||
| 25 | #include "vs_config.h" | ||
| 26 | |||
| 27 | #include <math.h> | ||
| 28 | #include <stddef.h> | ||
| 29 | #include <stdint.h> | ||
| 30 | #include <string.h> | ||
| 31 | |||
| 32 | #ifdef VS_HAVE_CBLAS | ||
| 33 | /* See matrix.c for the rationale on the Apple branch. */ | ||
| 34 | #ifdef __APPLE__ | ||
| 35 | #include <vecLib/cblas_new.h> | ||
| 36 | #else | ||
| 37 | #include <cblas.h> | ||
| 38 | #endif | ||
| 39 | #endif | ||
| 40 | |||
| 41 | #include "algo/kmeans_lloyd.h" | ||
| 42 | #include "algo/vecops.h" | ||
| 43 | #include "core/log.h" | ||
| 44 | #include "core/memory.h" | ||
| 45 | #include "types/vec16.h" | ||
| 46 | |||
| 47 | /* | ||
| 48 | * Precompute ||c||² for all centroids (L2 only). | ||
| 49 | */ | ||
| 50 | static void | ||
| 51 | 17243 | precompute_norms_c(KMeansState *st) | |
| 52 | { | ||
| 53 |
2/2✓ Branch 0 taken 126479 times.
✓ Branch 1 taken 17243 times.
|
143722 | for (uint32_t j = 0; j < st->nlist; j++) |
| 54 | 126479 | st->norms_c[j] = vs_l2_norm_squared( | |
| 55 | 126479 | st->centroids + (size_t)j * st->dim, st->dim); | |
| 56 | 17243 | } | |
| 57 | |||
| 58 | /* | ||
| 59 | * Assignment step: CBLAS path for one block. | ||
| 60 | * | ||
| 61 | * Computes the N_block x K distance matrix using sgemm, then finds | ||
| 62 | * the nearest centroid per vector (argmin per row). | ||
| 63 | * | ||
| 64 | * For L2: dist[i][j] = ||x_i||² + ||c_j||² - 2⟨x_i, c_j⟩ | ||
| 65 | * For IP: dist[i][j] = -⟨x_i, c_j⟩ | ||
| 66 | * For cos: dist[i][j] = 1 - ⟨x_i, c_j⟩ | ||
| 67 | */ | ||
| 68 | #ifdef VS_HAVE_CBLAS | ||
| 69 | __attribute__((always_inline)) static inline void | ||
| 70 | lloyd_assign_block_cblas_impl( | ||
| 71 | KMeansState *st, | ||
| 72 | uint32_t block_start, | ||
| 73 | uint32_t block_count, | ||
| 74 | const Vec32TypeOps *ops) | ||
| 75 | { | ||
| 76 | uint32_t nlist = st->nlist; | ||
| 77 | uint32_t dim = st->dim; | ||
| 78 | float *dist = st->dist_block; | ||
| 79 | |||
| 80 | /* Get float32 view of this block */ | ||
| 81 | const float *block_vecs; | ||
| 82 | if (st->indices != NULL) | ||
| 83 | { | ||
| 84 | /* Gather indexed vectors into vec_block */ | ||
| 85 | size_t esz = ops->element_size; | ||
| 86 | for (uint32_t i = 0; i < block_count; i++) | ||
| 87 | { | ||
| 88 | uint32_t idx = st->indices[block_start + i]; | ||
| 89 | const void *src = (const char *)st->vectors + | ||
| 90 | (size_t)idx * dim * esz; | ||
| 91 | /* Unreachable | ||
| 92 | * unless the standalone arena is exhausted, which no | ||
| 93 | * caller handles. */ | ||
| 94 | /* NOLINTNEXTLINE(clang-analyzer-core.NullPointerArithm) */ | ||
| 95 | ops->to_float_one(src, st->vec_block + (size_t)i * dim, dim); | ||
| 96 | } | ||
| 97 | block_vecs = st->vec_block; | ||
| 98 | } | ||
| 99 | else | ||
| 100 | { | ||
| 101 | /* Zero-copy for f32, converts for f16 */ | ||
| 102 | const void *raw = (const char *)st->vectors + | ||
| 103 | (size_t)block_start * dim * ops->element_size; | ||
| 104 | block_vecs = ops->to_float_block(raw, st->vec_block, block_count, dim); | ||
| 105 | } | ||
| 106 | |||
| 107 | /* Compute dot products via sgemm */ | ||
| 108 | float alpha = -2.0f; | ||
| 109 | if (st->metric == DISTANCE_INNER_PRODUCT || st->metric == DISTANCE_COSINE) | ||
| 110 | alpha = -1.0f; | ||
| 111 | |||
| 112 | cblas_sgemm( | ||
| 113 | CblasRowMajor, | ||
| 114 | CblasNoTrans, | ||
| 115 | CblasTrans, | ||
| 116 | (int)block_count, /* M: rows */ | ||
| 117 | (int)nlist, /* N: cols */ | ||
| 118 | (int)dim, /* K: inner dim */ | ||
| 119 | alpha, | ||
| 120 | block_vecs, | ||
| 121 | (int)dim, | ||
| 122 | st->centroids, | ||
| 123 | (int)dim, | ||
| 124 | 0.0f, | ||
| 125 | dist, | ||
| 126 | (int)nlist); | ||
| 127 | |||
| 128 | /* Add norms / offset */ | ||
| 129 | switch (st->metric) | ||
| 130 | { | ||
| 131 | case DISTANCE_L2: | ||
| 132 | for (uint32_t i = 0; i < block_count; i++) | ||
| 133 | { | ||
| 134 | float nx = st->norms_x[block_start + i]; | ||
| 135 | for (uint32_t j = 0; j < nlist; j++) | ||
| 136 | dist[(size_t)i * nlist + j] += nx + st->norms_c[j]; | ||
| 137 | } | ||
| 138 | break; | ||
| 139 | case DISTANCE_COSINE: | ||
| 140 | for (uint32_t i = 0; i < block_count; i++) | ||
| 141 | for (uint32_t j = 0; j < nlist; j++) | ||
| 142 | dist[(size_t)i * nlist + j] += 1.0f; | ||
| 143 | break; | ||
| 144 | case DISTANCE_INNER_PRODUCT: | ||
| 145 | /* dist = -⟨x,c⟩, already correct */ | ||
| 146 | break; | ||
| 147 | } | ||
| 148 | |||
| 149 | /* Find argmin per row */ | ||
| 150 | for (uint32_t i = 0; i < block_count; i++) | ||
| 151 | { | ||
| 152 | uint32_t vec_idx = block_start + i; | ||
| 153 | float min_dist = dist[(size_t)i * nlist]; | ||
| 154 | uint32_t min_j = 0; | ||
| 155 | |||
| 156 | for (uint32_t j = 1; j < nlist; j++) | ||
| 157 | { | ||
| 158 | float d = dist[(size_t)i * nlist + j]; | ||
| 159 | if (d < min_dist) | ||
| 160 | { | ||
| 161 | min_dist = d; | ||
| 162 | min_j = j; | ||
| 163 | } | ||
| 164 | } | ||
| 165 | |||
| 166 | st->assignments[vec_idx] = min_j; | ||
| 167 | st->total_cost += min_dist; | ||
| 168 | } | ||
| 169 | } | ||
| 170 | |||
| 171 | static void | ||
| 172 | lloyd_assign_block_cblas_f32( | ||
| 173 | KMeansState *st, uint32_t block_start, uint32_t block_count) | ||
| 174 | { | ||
| 175 | lloyd_assign_block_cblas_impl( | ||
| 176 | st, block_start, block_count, &vs_f32_type_ops); | ||
| 177 | } | ||
| 178 | |||
| 179 | static void | ||
| 180 | lloyd_assign_block_cblas_f16( | ||
| 181 | KMeansState *st, uint32_t block_start, uint32_t block_count) | ||
| 182 | { | ||
| 183 | lloyd_assign_block_cblas_impl( | ||
| 184 | st, block_start, block_count, &vs_f16_type_ops); | ||
| 185 | } | ||
| 186 | |||
| 187 | #if defined(VS_F16C_SUPPORT) && !defined(VS_SIMD_NONE) | ||
| 188 | static void | ||
| 189 | lloyd_assign_block_cblas_f16c( | ||
| 190 | KMeansState *st, uint32_t block_start, uint32_t block_count) | ||
| 191 | { | ||
| 192 | lloyd_assign_block_cblas_impl( | ||
| 193 | st, block_start, block_count, &vs_f16c_type_ops); | ||
| 194 | } | ||
| 195 | #endif | ||
| 196 | #endif /* VS_HAVE_CBLAS */ | ||
| 197 | |||
| 198 | /* | ||
| 199 | * Batch dot-product matrix: dots[i*nlist + j] = dot(vecs[i], cents[j]) | ||
| 200 | * | ||
| 201 | * VS_TARGET_CLONES generates AVX-512, AVX2, and default versions. | ||
| 202 | * The innermost loop over dim is auto-vectorized by the compiler, | ||
| 203 | * eliminating per-vector function pointer dispatch overhead. | ||
| 204 | * | ||
| 205 | * FMA generation requires -ffp-contract=fast (set in meson.build for | ||
| 206 | * release builds). Without it, GCC with -std=c2x generates separate | ||
| 207 | * vmulps + horizontal scalar adds instead of vfmadd231ps accumulate. | ||
| 208 | */ | ||
| 209 | VS_TARGET_CLONES static void | ||
| 210 | 21446 | lloyd_compute_dot_products( | |
| 211 | const float *vecs, | ||
| 212 | const float *centroids, | ||
| 213 | float *dots, | ||
| 214 | uint32_t block_count, | ||
| 215 | uint32_t nlist, | ||
| 216 | uint32_t dim) | ||
| 217 | { | ||
| 218 |
2/2✓ Branch 0 taken 13520036 times.
✓ Branch 1 taken 21446 times.
|
13541482 | for (uint32_t i = 0; i < block_count; i++) |
| 219 | { | ||
| 220 | 13520036 | const float *v = vecs + (size_t)i * dim; | |
| 221 |
2/2✓ Branch 0 taken 133054938 times.
✓ Branch 1 taken 13520036 times.
|
146574974 | for (uint32_t j = 0; j < nlist; j++) |
| 222 | { | ||
| 223 | 133054938 | const float *c = centroids + (size_t)j * dim; | |
| 224 | 133054938 | float dot = 0.0f; | |
| 225 |
2/2✓ Branch 0 taken 8952271717 times.
✓ Branch 1 taken 133054938 times.
|
9085326655 | for (uint32_t d = 0; d < dim; d++) |
| 226 | /* Unreachable | ||
| 227 | * unless the standalone arena is exhausted, which no | ||
| 228 | * caller handles. */ | ||
| 229 | /* NOLINTNEXTLINE(clang-analyzer-core.NullDereference) */ | ||
| 230 | 8952271717 | dot += v[d] * c[d]; | |
| 231 | 133054938 | dots[(size_t)i * nlist + j] = dot; | |
| 232 | } | ||
| 233 | } | ||
| 234 | 21446 | } | |
| 235 | |||
| 236 | /* | ||
| 237 | * Convert dot products to distances and find argmin per row. | ||
| 238 | */ | ||
| 239 | static void | ||
| 240 | 21446 | lloyd_dots_to_assignments( | |
| 241 | KMeansState *st, | ||
| 242 | float *dist, | ||
| 243 | uint32_t block_start, | ||
| 244 | uint32_t block_count) | ||
| 245 | { | ||
| 246 | 21446 | uint32_t nlist = st->nlist; | |
| 247 | |||
| 248 | /* Convert dot products to distances based on metric */ | ||
| 249 |
3/4✓ Branch 0 taken 8052 times.
✓ Branch 1 taken 374 times.
✓ Branch 2 taken 13020 times.
✗ Branch 3 not taken.
|
21446 | switch (st->metric) |
| 250 | { | ||
| 251 | 7357 | case DISTANCE_L2: | |
| 252 | /* dist = ||x||² + ||c||² - 2*dot(x,c) */ | ||
| 253 |
2/2✓ Branch 0 taken 13002830 times.
✓ Branch 1 taken 20141 times.
|
13022971 | for (uint32_t i = 0; i < block_count; i++) |
| 254 | { | ||
| 255 | 13002830 | float nx = st->norms_x[block_start + i]; | |
| 256 |
2/2✓ Branch 0 taken 129720771 times.
✓ Branch 1 taken 13002830 times.
|
142723601 | for (uint32_t j = 0; j < nlist; j++) |
| 257 | { | ||
| 258 | 129720771 | size_t idx = (size_t)i * nlist + j; | |
| 259 | 129720771 | dist[idx] = nx + st->norms_c[j] - 2.0f * dist[idx]; | |
| 260 | } | ||
| 261 | } | ||
| 262 | 7357 | break; | |
| 263 | 368 | case DISTANCE_INNER_PRODUCT: | |
| 264 | /* dist = -dot(x,c) */ | ||
| 265 |
2/2✓ Branch 0 taken 238650 times.
✓ Branch 1 taken 6368 times.
|
245018 | for (uint32_t i = 0; i < block_count; i++) |
| 266 |
2/2✓ Branch 0 taken 901932 times.
✓ Branch 1 taken 244644 times.
|
1146576 | for (uint32_t j = 0; j < nlist; j++) |
| 267 | { | ||
| 268 | 901932 | size_t idx = (size_t)i * nlist + j; | |
| 269 | 901932 | dist[idx] = -dist[idx]; | |
| 270 | } | ||
| 271 | 368 | break; | |
| 272 | 236 | case DISTANCE_COSINE: | |
| 273 | /* dist = 1 - dot(x,c) (vectors pre-normalized) */ | ||
| 274 |
2/2✓ Branch 0 taken 54407 times.
✓ Branch 1 taken 219086 times.
|
273493 | for (uint32_t i = 0; i < block_count; i++) |
| 275 |
2/2✓ Branch 0 taken 2432235 times.
✓ Branch 1 taken 272562 times.
|
2704797 | for (uint32_t j = 0; j < nlist; j++) |
| 276 | { | ||
| 277 | 2432235 | size_t idx = (size_t)i * nlist + j; | |
| 278 | 2432235 | dist[idx] = 1.0f - dist[idx]; | |
| 279 | } | ||
| 280 | 236 | break; | |
| 281 | } | ||
| 282 | |||
| 283 | /* Find argmin per row */ | ||
| 284 |
2/2✓ Branch 0 taken 13520036 times.
✓ Branch 1 taken 21446 times.
|
13541482 | for (uint32_t i = 0; i < block_count; i++) |
| 285 | { | ||
| 286 | 13520036 | uint32_t vec_idx = block_start + i; | |
| 287 | 13520036 | float min_dist = dist[(size_t)i * nlist]; | |
| 288 | 13520036 | uint32_t min_j = 0; | |
| 289 | |||
| 290 |
2/2✓ Branch 0 taken 119534902 times.
✓ Branch 1 taken 13520036 times.
|
133054938 | for (uint32_t j = 1; j < nlist; j++) |
| 291 | { | ||
| 292 | 119534902 | float d = dist[(size_t)i * nlist + j]; | |
| 293 |
2/2✓ Branch 0 taken 22098260 times.
✓ Branch 1 taken 97436642 times.
|
119534902 | if (d < min_dist) |
| 294 | { | ||
| 295 | 22098260 | min_dist = d; | |
| 296 | 22098260 | min_j = j; | |
| 297 | } | ||
| 298 | } | ||
| 299 | |||
| 300 | 13520036 | st->assignments[vec_idx] = min_j; | |
| 301 | 13520036 | st->total_cost += min_dist; | |
| 302 | } | ||
| 303 | 21446 | } | |
| 304 | |||
| 305 | /* | ||
| 306 | * Assignment step: builtin fallback for one block. | ||
| 307 | * | ||
| 308 | * For f32 input: uses batch dot-product + norm decomposition. | ||
| 309 | * For f16 input: preconverts block to f32, then uses the same f32 kernel. | ||
| 310 | */ | ||
| 311 | static void | ||
| 312 | 21434 | lloyd_assign_block_builtin_f32( | |
| 313 | KMeansState *st, uint32_t block_start, uint32_t block_count) | ||
| 314 | { | ||
| 315 | 21434 | uint32_t nlist = st->nlist; | |
| 316 | 21434 | uint32_t dim = st->dim; | |
| 317 | 21434 | float *dist = st->dist_block; | |
| 318 | |||
| 319 | 13485 | const float *block_vecs; | |
| 320 |
2/2✓ Branch 0 taken 17155 times.
✓ Branch 1 taken 4279 times.
|
21434 | if (st->indices != NULL) |
| 321 | { | ||
| 322 | /* Gather indexed vectors into contiguous block */ | ||
| 323 |
2/2✓ Branch 0 taken 7977793 times.
✓ Branch 1 taken 17155 times.
|
7994948 | for (uint32_t i = 0; i < block_count; i++) |
| 324 | { | ||
| 325 | 7977793 | uint32_t idx = st->indices[block_start + i]; | |
| 326 | /* Unreachable unless the standalone arena is exhausted, which | ||
| 327 | * no caller handles. */ | ||
| 328 | /* NOLINTNEXTLINE(clang-analyzer-core.NonNullParamChecker) */ | ||
| 329 | 7977793 | memcpy(st->vec_block + (size_t)i * dim, | |
| 330 | 7977793 | (const float *)st->vectors + (size_t)idx * dim, | |
| 331 | dim * sizeof(float)); | ||
| 332 | } | ||
| 333 | 17155 | block_vecs = st->vec_block; | |
| 334 | } | ||
| 335 | else | ||
| 336 | { | ||
| 337 | 4279 | block_vecs = (const float *)st->vectors + (size_t)block_start * dim; | |
| 338 | } | ||
| 339 | |||
| 340 | 21434 | lloyd_compute_dot_products( | |
| 341 | 21434 | block_vecs, st->centroids, dist, block_count, nlist, dim); | |
| 342 | 21434 | lloyd_dots_to_assignments(st, dist, block_start, block_count); | |
| 343 | 21434 | } | |
| 344 | |||
| 345 | static void | ||
| 346 | 12 | lloyd_assign_block_builtin_f16( | |
| 347 | KMeansState *st, uint32_t block_start, uint32_t block_count) | ||
| 348 | { | ||
| 349 | 12 | uint32_t nlist = st->nlist; | |
| 350 | 12 | uint32_t dim = st->dim; | |
| 351 | |||
| 352 | /* Preconvert f16 block to f32 — O(block*dim), saves O(block*K*dim) */ | ||
| 353 |
1/2✗ Branch 0 not taken.
✓ Branch 1 taken 12 times.
|
12 | if (st->indices != NULL) |
| 354 | { | ||
| 355 | ✗ | for (uint32_t i = 0; i < block_count; i++) | |
| 356 | { | ||
| 357 | ✗ | uint32_t idx = st->indices[block_start + i]; | |
| 358 | ✗ | const half *src = (const half *)st->vectors + (size_t)idx * dim; | |
| 359 | /* Unreachable | ||
| 360 | * unless the standalone arena is exhausted, which no | ||
| 361 | * caller handles. */ | ||
| 362 | /* NOLINTNEXTLINE(clang-analyzer-core.NullPointerArithm) */ | ||
| 363 | ✗ | vs_half_to_float_array(src, st->vec_block + (size_t)i * dim, dim); | |
| 364 | } | ||
| 365 | } | ||
| 366 | else | ||
| 367 | { | ||
| 368 | 12 | const half *src = (const half *)st->vectors + | |
| 369 | 12 | (size_t)block_start * dim; | |
| 370 | 12 | vs_half_to_float_array(src, st->vec_block, block_count * dim); | |
| 371 | } | ||
| 372 | |||
| 373 | 12 | lloyd_compute_dot_products( | |
| 374 | 12 | st->vec_block, | |
| 375 | 12 | st->centroids, | |
| 376 | st->dist_block, | ||
| 377 | block_count, | ||
| 378 | nlist, | ||
| 379 | dim); | ||
| 380 | 12 | lloyd_dots_to_assignments(st, st->dist_block, block_start, block_count); | |
| 381 | 12 | } | |
| 382 | |||
| 383 | /* Block assignment function pointer — selected once per lloyd_assign call */ | ||
| 384 | typedef void (*lloyd_block_fn)(KMeansState *, uint32_t, uint32_t); | ||
| 385 | |||
| 386 | static lloyd_block_fn | ||
| 387 | 2618 | lloyd_select_block_fn(KMeansState *st, bool use_cblas) | |
| 388 | { | ||
| 389 | #ifdef VS_HAVE_CBLAS | ||
| 390 | if (use_cblas) | ||
| 391 | { | ||
| 392 | switch (st->vec_type) | ||
| 393 | { | ||
| 394 | case VS_VEC_F32: | ||
| 395 | return lloyd_assign_block_cblas_f32; | ||
| 396 | #if defined(VS_F16C_SUPPORT) && !defined(VS_SIMD_NONE) | ||
| 397 | case VS_VEC_F16C: | ||
| 398 | return lloyd_assign_block_cblas_f16c; | ||
| 399 | #endif | ||
| 400 | default: | ||
| 401 | return lloyd_assign_block_cblas_f16; | ||
| 402 | } | ||
| 403 | } | ||
| 404 | #endif | ||
| 405 | 1512 | (void)use_cblas; | |
| 406 |
2/2✓ Branch 0 taken 1094 times.
✓ Branch 1 taken 12 times.
|
2618 | switch (st->vec_type) |
| 407 | { | ||
| 408 | 1094 | case VS_VEC_F32: | |
| 409 | 1094 | return lloyd_assign_block_builtin_f32; | |
| 410 | 12 | default: | |
| 411 | 12 | return lloyd_assign_block_builtin_f16; | |
| 412 | } | ||
| 413 | } | ||
| 414 | |||
| 415 | /* | ||
| 416 | * Adaptive block size for the parallel builtin path. | ||
| 417 | * | ||
| 418 | * The distance buffer is [block × nlist] floats per thread. With | ||
| 419 | * fixed block=4096, large nlist blows past L3 (e.g., 4096×8000×4 | ||
| 420 | * = 128MB per thread). Cap the buffer at ~2MB so the distance | ||
| 421 | * matrix, vector block, and centroid tile all fit in L3. | ||
| 422 | */ | ||
| 423 | static uint32_t | ||
| 424 | 2606 | lloyd_parallel_block_size(uint32_t nlist, uint32_t nvecs) | |
| 425 | { | ||
| 426 | 2606 | uint32_t block = KMEANS_BLOCK_SIZE; | |
| 427 | 2606 | uint32_t buf_sz = block * nlist * (uint32_t)sizeof(float); | |
| 428 |
1/2✗ Branch 0 not taken.
✓ Branch 1 taken 1094 times.
|
2606 | if (buf_sz > 8 * 1024 * 1024) |
| 429 | { | ||
| 430 | ✗ | block = 8 * 1024 * 1024 / (nlist * (uint32_t)sizeof(float)); | |
| 431 | ✗ | if (block < 64) | |
| 432 | ✗ | block = 64; | |
| 433 | } | ||
| 434 |
1/2✓ Branch 0 taken 1094 times.
✗ Branch 1 not taken.
|
2606 | if (block > nvecs) |
| 435 | 1094 | block = nvecs; | |
| 436 | 2606 | return block; | |
| 437 | } | ||
| 438 | |||
| 439 | static void | ||
| 440 | 20769 | lloyd_assign_range( | |
| 441 | KMeansState *st, | ||
| 442 | lloyd_block_fn block_fn, | ||
| 443 | float *dist_buf, | ||
| 444 | uint32_t block_size, | ||
| 445 | uint32_t range_start, | ||
| 446 | uint32_t range_end, | ||
| 447 | float *cost_out) | ||
| 448 | { | ||
| 449 | 20769 | float *saved_dist = st->dist_block; | |
| 450 | 20769 | st->dist_block = dist_buf; | |
| 451 | |||
| 452 | 20769 | float cost = 0.0f; | |
| 453 |
2/2✓ Branch 0 taken 21446 times.
✓ Branch 1 taken 20769 times.
|
42215 | for (uint32_t start = range_start; start < range_end; start += block_size) |
| 454 | { | ||
| 455 | 21446 | uint32_t count = range_end - start; | |
| 456 |
1/2✗ Branch 0 not taken.
✓ Branch 1 taken 7961 times.
|
21446 | if (count > block_size) |
| 457 | ✗ | count = block_size; | |
| 458 | 21446 | st->total_cost = 0.0f; | |
| 459 | 21446 | block_fn(st, start, count); | |
| 460 | 21446 | cost += st->total_cost; | |
| 461 | } | ||
| 462 | 20769 | *cost_out = cost; | |
| 463 | 20769 | st->dist_block = saved_dist; | |
| 464 | 20769 | } | |
| 465 | |||
| 466 | typedef struct | ||
| 467 | { | ||
| 468 | KMeansState *st; | ||
| 469 | lloyd_block_fn block_fn; | ||
| 470 | float *dist_bufs; /* [nthreads * block * nlist] */ | ||
| 471 | float *costs; /* [nthreads] */ | ||
| 472 | float *vec_bufs; /* [nthreads * block * dim] or NULL */ | ||
| 473 | uint32_t block; | ||
| 474 | } LloydParCtx; | ||
| 475 | |||
| 476 | static void | ||
| 477 | 2606 | lloyd_par_worker(uint32_t thread_id, uint32_t start, uint32_t end, void *arg) | |
| 478 | { | ||
| 479 | 2606 | LloydParCtx *ctx = (LloydParCtx *)arg; | |
| 480 | 2606 | uint32_t nlist = ctx->st->nlist; | |
| 481 | 2606 | uint32_t dim = ctx->st->dim; | |
| 482 | 2606 | uint32_t block = ctx->block; | |
| 483 | |||
| 484 | 2606 | KMeansState local = *ctx->st; | |
| 485 | 2606 | local.dist_block = ctx->dist_bufs + (size_t)thread_id * block * nlist; | |
| 486 |
2/2✓ Branch 0 taken 2175 times.
✓ Branch 1 taken 431 times.
|
2606 | if (ctx->vec_bufs) |
| 487 | 2175 | local.vec_block = ctx->vec_bufs + (size_t)thread_id * block * dim; | |
| 488 | |||
| 489 | 2606 | lloyd_assign_range( | |
| 490 | &local, | ||
| 491 | ctx->block_fn, | ||
| 492 | local.dist_block, | ||
| 493 | ctx->block, | ||
| 494 | start, | ||
| 495 | end, | ||
| 496 | 2606 | &ctx->costs[thread_id]); | |
| 497 | 2606 | } | |
| 498 | |||
| 499 | void | ||
| 500 | 12 | lloyd_assign(KMeansState *st, bool use_cblas) | |
| 501 | { | ||
| 502 | 12 | st->total_cost = 0.0f; | |
| 503 | |||
| 504 |
1/2✓ Branch 0 taken 12 times.
✗ Branch 1 not taken.
|
12 | if (st->metric == DISTANCE_L2) |
| 505 | 12 | precompute_norms_c(st); | |
| 506 | |||
| 507 |
0/2✗ Branch 0 not taken.
✗ Branch 1 not taken.
|
12 | lloyd_block_fn block_fn = lloyd_select_block_fn(st, use_cblas); |
| 508 | |||
| 509 | 12 | lloyd_assign_range( | |
| 510 | st, | ||
| 511 | block_fn, | ||
| 512 | st->dist_block, | ||
| 513 | KMEANS_BLOCK_SIZE, | ||
| 514 | 0, | ||
| 515 | st->nvecs, | ||
| 516 | &st->total_cost); | ||
| 517 | 12 | } | |
| 518 | |||
| 519 | /* ---------------------------------------------------------------- | ||
| 520 | * Fused iterative assign + update via iterate callback | ||
| 521 | * ---------------------------------------------------------------- */ | ||
| 522 | |||
| 523 | typedef struct | ||
| 524 | { | ||
| 525 | KMeansState *st; | ||
| 526 | lloyd_block_fn block_fn; | ||
| 527 | uint32_t nthreads; | ||
| 528 | uint32_t block; | ||
| 529 | |||
| 530 | /* Per-thread buffers */ | ||
| 531 | float *dist_bufs; /* [nthreads * block * nlist] */ | ||
| 532 | float *vec_bufs; /* [nthreads * block * dim] or NULL */ | ||
| 533 | float *centroid_sums; /* [nthreads * nlist * dim] */ | ||
| 534 | uint32_t *centroid_cnts; /* [nthreads * nlist] */ | ||
| 535 | float *costs; /* [nthreads] */ | ||
| 536 | |||
| 537 | /* Convergence */ | ||
| 538 | float *old_centroids; /* [nlist * dim] saved before iteration */ | ||
| 539 | float tolerance; | ||
| 540 | bool verbose; | ||
| 541 | uint32_t completed_iters; | ||
| 542 | } LloydIterCtx; | ||
| 543 | |||
| 544 | /* | ||
| 545 | * Work function: assign vectors [start, end) to nearest centroids, | ||
| 546 | * then accumulate per-thread centroid sums for the update step. | ||
| 547 | * | ||
| 548 | * Each thread gets its own dist/vec buffers and centroid accumulators | ||
| 549 | * so there is no shared mutable state during the parallel phase. | ||
| 550 | */ | ||
| 551 | static void | ||
| 552 | 18151 | lloyd_iter_work(uint32_t thread_id, uint32_t start, uint32_t end, void *arg) | |
| 553 | { | ||
| 554 | 18151 | LloydIterCtx *ctx = (LloydIterCtx *)arg; | |
| 555 | 18151 | KMeansState *st = ctx->st; | |
| 556 | 18151 | uint32_t nlist = st->nlist; | |
| 557 | 18151 | uint32_t dim = st->dim; | |
| 558 | 18151 | uint32_t block = ctx->block; | |
| 559 | |||
| 560 | /* 1. Set up thread-local KMeansState with private buffers */ | ||
| 561 | 18151 | KMeansState local = *st; | |
| 562 | 18151 | local.dist_block = ctx->dist_bufs + (size_t)thread_id * block * nlist; | |
| 563 |
2/2✓ Branch 0 taken 14980 times.
✓ Branch 1 taken 3171 times.
|
18151 | if (ctx->vec_bufs) |
| 564 | 14980 | local.vec_block = ctx->vec_bufs + (size_t)thread_id * block * dim; | |
| 565 | |||
| 566 | /* 2. Assign: compute distances and find nearest centroid */ | ||
| 567 | 18151 | lloyd_assign_range( | |
| 568 | &local, | ||
| 569 | ctx->block_fn, | ||
| 570 | local.dist_block, | ||
| 571 | block, | ||
| 572 | start, | ||
| 573 | end, | ||
| 574 | 18151 | &ctx->costs[thread_id]); | |
| 575 | |||
| 576 | /* 3. Accumulate: add each vector to its assigned centroid's sum */ | ||
| 577 | 18151 | float *my_sums = ctx->centroid_sums + (size_t)thread_id * nlist * dim; | |
| 578 | 18151 | uint32_t *my_cnts = ctx->centroid_cnts + (size_t)thread_id * nlist; | |
| 579 | |||
| 580 |
2/2✓ Branch 0 taken 12713258 times.
✓ Branch 1 taken 18151 times.
|
12731409 | for (uint32_t i = start; i < end; i++) |
| 581 | { | ||
| 582 | 12713258 | ClusterId c = st->assignments[i]; | |
| 583 |
2/2✓ Branch 0 taken 7515454 times.
✓ Branch 1 taken 5197804 times.
|
12713258 | uint32_t idx = st->indices ? st->indices[i] : i; |
| 584 | 12713258 | const float *vec = (const float *)st->vectors + (size_t)idx * dim; | |
| 585 | 12713258 | float *sum = my_sums + (size_t)c * dim; | |
| 586 |
2/2✓ Branch 0 taken 776043970 times.
✓ Branch 1 taken 12713258 times.
|
788757228 | for (uint32_t d = 0; d < dim; d++) |
| 587 | 776043970 | sum[d] += vec[d]; | |
| 588 | 12713258 | my_cnts[c]++; | |
| 589 | } | ||
| 590 | 18151 | } | |
| 591 | |||
| 592 | /* | ||
| 593 | * Reduce function: merge per-thread results, update centroids, | ||
| 594 | * and check convergence. Runs on the leader thread between | ||
| 595 | * barrier-synchronized iterations. | ||
| 596 | * | ||
| 597 | * Returns true to continue iterating, false to stop. | ||
| 598 | */ | ||
| 599 | static bool | ||
| 600 | 18151 | lloyd_iter_reduce(void *arg, uint32_t iteration) | |
| 601 | { | ||
| 602 | 18151 | LloydIterCtx *ctx = (LloydIterCtx *)arg; | |
| 603 | 18151 | KMeansState *st = ctx->st; | |
| 604 | 18151 | uint32_t nlist = st->nlist; | |
| 605 | 18151 | uint32_t dim = st->dim; | |
| 606 | 18151 | uint32_t nt = ctx->nthreads; | |
| 607 | |||
| 608 | /* 1. Sum per-thread costs into total */ | ||
| 609 | 18151 | st->total_cost = 0.0f; | |
| 610 |
2/2✓ Branch 0 taken 18151 times.
✓ Branch 1 taken 18151 times.
|
36302 | for (uint32_t t = 0; t < nt; t++) |
| 611 | 18151 | st->total_cost += ctx->costs[t]; | |
| 612 | |||
| 613 | /* 2. Merge per-thread centroid accumulators into new_centroids */ | ||
| 614 | 18151 | memset(st->new_centroids, 0, (size_t)nlist * dim * sizeof(float)); | |
| 615 | 18151 | memset(st->cluster_sizes, 0, nlist * sizeof(uint32_t)); | |
| 616 | |||
| 617 |
2/2✓ Branch 0 taken 18151 times.
✓ Branch 1 taken 18151 times.
|
36302 | for (uint32_t t = 0; t < nt; t++) |
| 618 | { | ||
| 619 | 18151 | float *sums = ctx->centroid_sums + (size_t)t * nlist * dim; | |
| 620 | 18151 | uint32_t *cnts = ctx->centroid_cnts + (size_t)t * nlist; | |
| 621 | |||
| 622 |
2/2✓ Branch 0 taken 129756 times.
✓ Branch 1 taken 18151 times.
|
147907 | for (uint32_t c = 0; c < nlist; c++) |
| 623 | { | ||
| 624 | 129756 | st->cluster_sizes[c] += cnts[c]; | |
| 625 | 129756 | float *dst = st->new_centroids + (size_t)c * dim; | |
| 626 | 129756 | float *src = sums + (size_t)c * dim; | |
| 627 |
2/2✓ Branch 0 taken 10073205 times.
✓ Branch 1 taken 129756 times.
|
10202961 | for (uint32_t d = 0; d < dim; d++) |
| 628 | 10073205 | dst[d] += src[d]; | |
| 629 | } | ||
| 630 | } | ||
| 631 | |||
| 632 | /* 3. Compute new centroids as mean of assigned vectors */ | ||
| 633 |
2/2✓ Branch 0 taken 129756 times.
✓ Branch 1 taken 18151 times.
|
147907 | for (uint32_t c = 0; c < nlist; c++) |
| 634 | { | ||
| 635 |
2/2✓ Branch 0 taken 566 times.
✓ Branch 1 taken 129190 times.
|
129756 | if (st->cluster_sizes[c] == 0) |
| 636 | 566 | continue; | |
| 637 | 129190 | float inv = 1.0f / (float)st->cluster_sizes[c]; | |
| 638 | 129190 | float *cent = st->new_centroids + (size_t)c * dim; | |
| 639 |
2/2✓ Branch 0 taken 10060441 times.
✓ Branch 1 taken 129190 times.
|
10189631 | for (uint32_t d = 0; d < dim; d++) |
| 640 | 10060441 | cent[d] *= inv; | |
| 641 | } | ||
| 642 | |||
| 643 | /* 4. Re-normalize centroids for cosine metric */ | ||
| 644 |
2/2✓ Branch 0 taken 861 times.
✓ Branch 1 taken 17290 times.
|
18151 | if (st->metric == DISTANCE_COSINE) |
| 645 | { | ||
| 646 |
2/2✓ Branch 0 taken 4744 times.
✓ Branch 1 taken 861 times.
|
5605 | for (uint32_t c = 0; c < nlist; c++) |
| 647 | { | ||
| 648 |
1/2✗ Branch 0 not taken.
✓ Branch 1 taken 4744 times.
|
4744 | if (st->cluster_sizes[c] == 0) |
| 649 | ✗ | continue; | |
| 650 | 4744 | float *cent = st->new_centroids + (size_t)c * dim; | |
| 651 | 4744 | float norm = vs_l2_norm(cent, dim); | |
| 652 |
1/2✓ Branch 0 taken 4744 times.
✗ Branch 1 not taken.
|
4744 | if (norm > 1e-10f) |
| 653 | 4744 | vec32_scale(cent, 1.0f / norm, cent, dim); | |
| 654 | } | ||
| 655 | } | ||
| 656 | |||
| 657 | /* 5. Swap old and new centroids */ | ||
| 658 | 18151 | float *tmp = st->centroids; | |
| 659 | 18151 | st->centroids = st->new_centroids; | |
| 660 | 18151 | st->new_centroids = tmp; | |
| 661 | |||
| 662 | /* 6. Check convergence: max centroid movement (squared) */ | ||
| 663 | 29447 | float shift_sq = kmeans_max_centroid_shift_between( | |
| 664 | 18151 | st->centroids, ctx->old_centroids, nlist, dim); | |
| 665 | |||
| 666 |
2/2✓ Branch 0 taken 8 times.
✓ Branch 1 taken 18143 times.
|
18151 | if (ctx->verbose) |
| 667 |
0/2✗ Branch 1 not taken.
✗ Branch 2 not taken.
|
8 | vs_log(" iter %u: cost=%.4f, max_shift=%.6f\n", |
| 668 | iteration, | ||
| 669 | st->total_cost, | ||
| 670 | sqrtf(shift_sq)); | ||
| 671 | |||
| 672 | 18151 | ctx->completed_iters = iteration + 1; | |
| 673 | |||
| 674 | 18151 | float tol_sq = ctx->tolerance * ctx->tolerance; | |
| 675 |
2/2✓ Branch 0 taken 11012 times.
✓ Branch 1 taken 7139 times.
|
18151 | if (shift_sq < tol_sq) |
| 676 | 1010 | return false; | |
| 677 | |||
| 678 | /* 7. Prepare for next iteration: save centroids, reset accumulators */ | ||
| 679 | 15847 | memcpy(ctx->old_centroids, | |
| 680 |
2/2✓ Branch 0 taken 9398 times.
✓ Branch 1 taken 604 times.
|
15847 | st->centroids, |
| 681 | 5845 | (size_t)nlist * dim * sizeof(float)); | |
| 682 |
2/2✓ Branch 0 taken 9398 times.
✓ Branch 1 taken 604 times.
|
15847 | memset(ctx->centroid_sums, 0, (size_t)nt * nlist * dim * sizeof(float)); |
| 683 | 15847 | memset(ctx->centroid_cnts, 0, (size_t)nt * nlist * sizeof(uint32_t)); | |
| 684 | 15847 | memset(ctx->costs, 0, nt * sizeof(float)); | |
| 685 | |||
| 686 | /* 8. Precompute centroid norms for next iteration's distance calc */ | ||
| 687 |
2/2✓ Branch 0 taken 14723 times.
✓ Branch 1 taken 1124 times.
|
15847 | if (st->metric == DISTANCE_L2) |
| 688 | 14723 | precompute_norms_c(st); | |
| 689 | |||
| 690 | 5845 | return true; | |
| 691 | } | ||
| 692 | |||
| 693 | uint32_t | ||
| 694 | 2606 | lloyd_iterate(KMeansState *st, bool use_cblas, const KMeansOptions *opts) | |
| 695 | { | ||
| 696 | 1512 | (void)use_cblas; | |
| 697 | |||
| 698 | 2606 | uint32_t nt = 1; /* serial: the driver owns parallelism, not k-means */ | |
| 699 | 2606 | uint32_t nlist = st->nlist; | |
| 700 | 2606 | uint32_t dim = st->dim; | |
| 701 |
1/2✗ Branch 0 not taken.
✓ Branch 1 taken 1512 times.
|
2606 | uint32_t block = lloyd_parallel_block_size(nlist, st->nvecs); |
| 702 | |||
| 703 |
1/2✗ Branch 0 not taken.
✓ Branch 1 taken 1512 times.
|
2606 | lloyd_block_fn block_fn = lloyd_select_block_fn(st, use_cblas); |
| 704 | |||
| 705 | /* Precompute norms before first iteration */ | ||
| 706 |
2/2✓ Branch 0 taken 2508 times.
✓ Branch 1 taken 98 times.
|
2606 | if (st->metric == DISTANCE_L2) |
| 707 | 2508 | precompute_norms_c(st); | |
| 708 | |||
| 709 | 2606 | VS_MEMCTX_SCOPE(iter_ctx); | |
| 710 | 2606 | VsMemCtx old_ctx = vs_memctx_switch(iter_ctx); | |
| 711 | |||
| 712 | 2606 | float *dist_bufs = vs_alloc((size_t)nt * block * nlist * sizeof(float)); | |
| 713 | 2606 | float *vec_bufs = NULL; | |
| 714 |
2/2✓ Branch 0 taken 2175 times.
✓ Branch 1 taken 431 times.
|
2606 | if (st->vec_block != NULL) |
| 715 | 2175 | vec_bufs = vs_alloc((size_t)nt * block * dim * sizeof(float)); | |
| 716 | |||
| 717 | 2606 | float *centroid_sums = vs_alloc0((size_t)nt * nlist * dim * sizeof(float)); | |
| 718 | 2606 | uint32_t *centroid_cnts = vs_alloc0((size_t)nt * nlist * sizeof(uint32_t)); | |
| 719 | 2606 | float *costs = vs_alloc0(nt * sizeof(float)); | |
| 720 | 2606 | float *old_centroids = vs_alloc((size_t)nlist * dim * sizeof(float)); | |
| 721 | |||
| 722 | 2606 | vs_memctx_switch(old_ctx); | |
| 723 | |||
| 724 | 2606 | memcpy(old_centroids, st->centroids, (size_t)nlist * dim * sizeof(float)); | |
| 725 | |||
| 726 | 2606 | LloydIterCtx ctx = { | |
| 727 | .st = st, | ||
| 728 | .block_fn = block_fn, | ||
| 729 | .nthreads = nt, | ||
| 730 | .block = block, | ||
| 731 | .dist_bufs = dist_bufs, | ||
| 732 | .vec_bufs = vec_bufs, | ||
| 733 | .centroid_sums = centroid_sums, | ||
| 734 | .centroid_cnts = centroid_cnts, | ||
| 735 | .costs = costs, | ||
| 736 | .old_centroids = old_centroids, | ||
| 737 | 2606 | .tolerance = opts->tolerance, | |
| 738 | 2606 | .verbose = opts->verbose, | |
| 739 | .completed_iters = 0, | ||
| 740 | }; | ||
| 741 | |||
| 742 |
2/2✓ Branch 0 taken 18151 times.
✓ Branch 1 taken 302 times.
|
18453 | for (uint32_t iter = 0; iter < opts->max_iterations; iter++) |
| 743 | { | ||
| 744 | 18151 | lloyd_iter_work(0, 0, st->nvecs, &ctx); | |
| 745 |
2/2✓ Branch 1 taken 11012 times.
✓ Branch 2 taken 7139 times.
|
18151 | if (!lloyd_iter_reduce(&ctx, iter)) |
| 746 | 1010 | break; | |
| 747 | } | ||
| 748 | |||
| 749 | /* Final assignment (reduce already did the last centroid update) */ | ||
| 750 | 2606 | memset(costs, 0, nt * sizeof(float)); | |
| 751 | |||
| 752 | 2606 | LloydParCtx par_ctx = { | |
| 753 | .st = st, | ||
| 754 | .block_fn = block_fn, | ||
| 755 | .dist_bufs = dist_bufs, | ||
| 756 | .costs = costs, | ||
| 757 | .vec_bufs = vec_bufs, | ||
| 758 | .block = block, | ||
| 759 | }; | ||
| 760 | |||
| 761 | 2606 | lloyd_par_worker(0, 0, st->nvecs, &par_ctx); | |
| 762 | 2606 | st->total_cost = 0.0f; | |
| 763 |
2/2✓ Branch 0 taken 2606 times.
✓ Branch 1 taken 2606 times.
|
5212 | for (uint32_t t = 0; t < nt; t++) |
| 764 | 2606 | st->total_cost += costs[t]; | |
| 765 | |||
| 766 |
1/2✓ Branch 0 taken 1512 times.
✗ Branch 1 not taken.
|
2606 | return ctx.completed_iters; |
| 767 | } | ||
| 768 |