Clustering
C++20 header-only: DBSCAN, HDBSCAN, k-means.
Loading...
Searching...
No Matches
greedy_kmpp_seeder.h
Go to the documentation of this file.
1#pragma once
2
3#include <algorithm>
4#include <array>
5#include <cmath>
6#include <cstddef>
7#include <cstdint>
8#include <cstring>
9#include <limits>
10#include <type_traits>
11#include <vector>
12
14#include "clustering/math/detail/avx2_helpers.h"
15#include "clustering/math/detail/gemm_outer.h"
16#include "clustering/math/detail/inverse_cdf_blocks.h"
17#include "clustering/math/detail/matrix_desc.h"
18#include "clustering/math/detail/sq_distances_block.h"
20#include "clustering/math/rng.h"
22#include "clustering/ndarray.h"
23
24#ifdef CLUSTERING_USE_AVX2
25#include <immintrin.h>
26
27#include "clustering/math/detail/kmpp_score_avx2.h"
28#endif
29
30namespace clustering::kmeans {
31
32namespace detail {
33
34using math::detail::sqEuclideanRowPtr;
35
45[[nodiscard]] inline std::size_t greedyKmppLocalTrials(std::size_t k) noexcept {
46 if (k <= 1) {
47 return 1;
48 }
49 const auto lnK = std::log(static_cast<double>(k));
50 return 2 + static_cast<std::size_t>(lnK);
51}
52
60[[nodiscard]] constexpr std::size_t greedyKmppTransposedWidth(std::size_t L) noexcept {
61 constexpr std::size_t kChunk = 8;
62 return ((L + kChunk - 1) / kChunk) * kChunk;
63}
64
72[[nodiscard]] inline std::size_t greedyKmppSweepBlocks(math::Pool pool, std::size_t rows,
73 std::size_t opsPerRow) noexcept {
74 constexpr std::size_t kMinOpsPerBlock = std::size_t{1} << 15;
75 if (pool.pool == nullptr || rows == 0) {
76 return 1;
77 }
78 if (pool.shouldParallelize(rows, 1024, 2)) {
79 return pool.stealBlocks(rows);
80 }
81 const std::size_t cap = std::max<std::size_t>(1, (rows * opsPerRow) / kMinOpsPerBlock);
82 return std::min(pool.stealBlocks(rows), cap);
83}
84
85#ifdef CLUSTERING_USE_AVX2
86
96template <std::size_t B>
97[[gnu::always_inline]] inline void
98sqEuclideanRowToBatchAvx2Fixed(const float *x, const float *candData, std::size_t d,
99 float *out) noexcept {
100 static_assert(B >= 1 && B <= 8, "B must lie in [1, 8] -- 8 ymm regs hold the batch");
101 // Double accumulator set (2 * B YMMs) over a 2x-unrolled K loop. Halves the per-iter fmadd
102 // dependency chain so Zen5's 4-FMA-per-cycle throughput isn't latency-bound on the 4-cycle
103 // fmadd round-trip; also gives the register allocator enough explicit live ranges to keep
104 // accumulators in YMM registers rather than spilling to the stack (measured: 8 GFLOPS with
105 // the original 1x loop, ~2x post-unroll on the seeder's B=4 hot path).
106 std::array<__m256, B> acc0{};
107 std::array<__m256, B> acc1{};
108 for (std::size_t t = 0; t < B; ++t) {
109 acc0[t] = _mm256_setzero_ps();
110 acc1[t] = _mm256_setzero_ps();
111 }
112 std::size_t k = 0;
113 for (; k + 16 <= d; k += 16) {
114 const __m256 vx0 = _mm256_loadu_ps(x + k);
115 const __m256 vx1 = _mm256_loadu_ps(x + k + 8);
116 for (std::size_t t = 0; t < B; ++t) {
117 const __m256 vc0 = _mm256_loadu_ps(candData + (t * d) + k);
118 const __m256 vc1 = _mm256_loadu_ps(candData + (t * d) + k + 8);
119 const __m256 diff0 = _mm256_sub_ps(vx0, vc0);
120 const __m256 diff1 = _mm256_sub_ps(vx1, vc1);
121 acc0[t] = _mm256_fmadd_ps(diff0, diff0, acc0[t]);
122 acc1[t] = _mm256_fmadd_ps(diff1, diff1, acc1[t]);
123 }
124 }
125 // 8-lane tail.
126 for (; k + 8 <= d; k += 8) {
127 const __m256 vx = _mm256_loadu_ps(x + k);
128 for (std::size_t t = 0; t < B; ++t) {
129 const __m256 vc = _mm256_loadu_ps(candData + (t * d) + k);
130 const __m256 diff = _mm256_sub_ps(vx, vc);
131 acc0[t] = _mm256_fmadd_ps(diff, diff, acc0[t]);
132 }
133 }
134 std::array<float, B> tail{};
135 for (std::size_t t = 0; t < B; ++t) {
136 tail[t] = 0.0F;
137 }
138 for (std::size_t kt = k; kt < d; ++kt) {
139 const float xk = x[kt];
140 for (std::size_t t = 0; t < B; ++t) {
141 const float diff = xk - candData[(t * d) + kt];
142 tail[t] += diff * diff;
143 }
144 }
145 for (std::size_t t = 0; t < B; ++t) {
146 const __m256 sum = _mm256_add_ps(acc0[t], acc1[t]);
147 out[t] = math::detail::horizontalSumAvx2(sum) + tail[t];
148 }
149}
150
151template <std::size_t B>
152inline void sqEuclideanRowToBatchAvx2Fixed(const double *x, const double *candData, std::size_t d,
153 double *out) noexcept {
154 static_assert(B >= 1 && B <= 8, "B must lie in [1, 8] -- 8 ymm regs hold the batch");
155 std::array<__m256d, B> acc{};
156 for (std::size_t t = 0; t < B; ++t) {
157 acc[t] = _mm256_setzero_pd();
158 }
159 std::size_t k = 0;
160 for (; k + 4 <= d; k += 4) {
161 const __m256d vx = _mm256_loadu_pd(x + k);
162 for (std::size_t t = 0; t < B; ++t) {
163 const __m256d vc = _mm256_loadu_pd(candData + (t * d) + k);
164 const __m256d diff = _mm256_sub_pd(vx, vc);
165 acc[t] = _mm256_fmadd_pd(diff, diff, acc[t]);
166 }
167 }
168 std::array<double, B> tail{};
169 for (std::size_t t = 0; t < B; ++t) {
170 tail[t] = 0.0;
171 }
172 for (std::size_t kt = k; kt < d; ++kt) {
173 const double xk = x[kt];
174 for (std::size_t t = 0; t < B; ++t) {
175 const double diff = xk - candData[(t * d) + kt];
176 tail[t] += diff * diff;
177 }
178 }
179 for (std::size_t t = 0; t < B; ++t) {
180 out[t] = math::detail::horizontalSumAvx2(acc[t]) + tail[t];
181 }
182}
183
192template <class T>
193inline void sqEuclideanRowToBatchAvx2(const T *x, const T *candData, std::size_t L, std::size_t d,
194 T *out) noexcept {
195 std::size_t base = 0;
196 while (base + 8 <= L) {
197 sqEuclideanRowToBatchAvx2Fixed<8>(x, candData + (base * d), d, out + base);
198 base += 8;
199 }
200 switch (L - base) {
201 case 0:
202 break;
203 case 1:
204 sqEuclideanRowToBatchAvx2Fixed<1>(x, candData + (base * d), d, out + base);
205 break;
206 case 2:
207 sqEuclideanRowToBatchAvx2Fixed<2>(x, candData + (base * d), d, out + base);
208 break;
209 case 3:
210 sqEuclideanRowToBatchAvx2Fixed<3>(x, candData + (base * d), d, out + base);
211 break;
212 case 4:
213 sqEuclideanRowToBatchAvx2Fixed<4>(x, candData + (base * d), d, out + base);
214 break;
215 case 5:
216 sqEuclideanRowToBatchAvx2Fixed<5>(x, candData + (base * d), d, out + base);
217 break;
218 case 6:
219 sqEuclideanRowToBatchAvx2Fixed<6>(x, candData + (base * d), d, out + base);
220 break;
221 case 7:
222 sqEuclideanRowToBatchAvx2Fixed<7>(x, candData + (base * d), d, out + base);
223 break;
224 default:
225 break;
226 }
227}
228
238inline void sqEuclideanRowAgainst8Transposed(const float *x, const float *candData, std::size_t d,
239 float *out) noexcept {
240 __m256 acc = _mm256_setzero_ps();
241 for (std::size_t k = 0; k < d; ++k) {
242 const __m256 cv = _mm256_load_ps(candData + (k * 8));
243 const __m256 xv = _mm256_set1_ps(x[k]);
244 const __m256 diff = _mm256_sub_ps(xv, cv);
245 acc = _mm256_fmadd_ps(diff, diff, acc);
246 }
247 _mm256_storeu_ps(out, acc);
248}
249
257inline __m256 sqEuclideanRowAgainst8TransposedReg(const float *x, const float *candData,
258 std::size_t d) noexcept {
259 __m256 acc = _mm256_setzero_ps();
260 for (std::size_t k = 0; k < d; ++k) {
261 const __m256 cv = _mm256_load_ps(candData + (k * 8));
262 const __m256 xv = _mm256_set1_ps(x[k]);
263 const __m256 diff = _mm256_sub_ps(xv, cv);
264 acc = _mm256_fmadd_ps(diff, diff, acc);
265 }
266 return acc;
267}
268
274inline std::pair<__m256, __m256> sqEuclideanRowAgainst16TransposedReg(const float *x,
275 const float *candData,
276 std::size_t d) noexcept {
277 __m256 accLo = _mm256_setzero_ps();
278 __m256 accHi = _mm256_setzero_ps();
279 for (std::size_t k = 0; k < d; ++k) {
280 const __m256 cLo = _mm256_load_ps(candData + (k * 16));
281 const __m256 cHi = _mm256_load_ps(candData + (k * 16) + 8);
282 const __m256 xv = _mm256_set1_ps(x[k]);
283 const __m256 diffLo = _mm256_sub_ps(xv, cLo);
284 const __m256 diffHi = _mm256_sub_ps(xv, cHi);
285 accLo = _mm256_fmadd_ps(diffLo, diffLo, accLo);
286 accHi = _mm256_fmadd_ps(diffHi, diffHi, accHi);
287 }
288 return {accLo, accHi};
289}
290
299inline void sqEuclideanRowAgainst16Transposed(const float *x, const float *candData, std::size_t d,
300 float *out) noexcept {
301 __m256 accLo = _mm256_setzero_ps();
302 __m256 accHi = _mm256_setzero_ps();
303 for (std::size_t k = 0; k < d; ++k) {
304 const __m256 cLo = _mm256_load_ps(candData + (k * 16));
305 const __m256 cHi = _mm256_load_ps(candData + (k * 16) + 8);
306 const __m256 xv = _mm256_set1_ps(x[k]);
307 const __m256 diffLo = _mm256_sub_ps(xv, cLo);
308 const __m256 diffHi = _mm256_sub_ps(xv, cHi);
309 accLo = _mm256_fmadd_ps(diffLo, diffLo, accLo);
310 accHi = _mm256_fmadd_ps(diffHi, diffHi, accHi);
311 }
312 _mm256_storeu_ps(out, accLo);
313 _mm256_storeu_ps(out + 8, accHi);
314}
315
324inline void sqEuclideanRowAgainst8TransposedStrided(const float *x, const float *candData,
325 std::size_t d, std::size_t rowStride,
326 float *out) noexcept {
327 __m256 acc = _mm256_setzero_ps();
328 for (std::size_t k = 0; k < d; ++k) {
329 const __m256 cv = _mm256_loadu_ps(candData + (k * rowStride));
330 const __m256 xv = _mm256_set1_ps(x[k]);
331 const __m256 diff = _mm256_sub_ps(xv, cv);
332 acc = _mm256_fmadd_ps(diff, diff, acc);
333 }
334 _mm256_storeu_ps(out, acc);
335}
336
337#endif // CLUSTERING_USE_AVX2
338
347template <class T>
348inline void sqEuclideanRowToBatch(const T *x, const T *candData, std::size_t L, std::size_t d,
349 T *out) noexcept {
350#ifdef CLUSTERING_USE_AVX2
351 if constexpr (std::is_same_v<T, float> || std::is_same_v<T, double>) {
353 sqEuclideanRowToBatchAvx2(x, candData, L, d, out);
354 return;
355 }
356 }
357#endif
358 for (std::size_t t = 0; t < L; ++t) {
359 out[t] = sqEuclideanRowPtr(x, candData + (t * d), d);
360 }
361}
362
363} // namespace detail
364
381template <class T> class GreedyKmppSeeder {
382public:
383 static_assert(std::is_same_v<T, float> || std::is_same_v<T, double>,
384 "GreedyKmppSeeder<T> requires T to be float or double");
385
387 : m_candRows({0, 0}), m_candRowsT({0, 0}), m_candDistSq({0, 0}), m_sweepSums({0}),
388 m_minSq({0}), m_distsFlat({0, 0}), m_xNormsSq({0}), m_candNormsSq({0}), m_gemmApArena({0}),
389 m_gemmBpArena({0}), m_localScores({0}) {}
390
401 void run(const NDArray<T, 2, Layout::Contig> &X, std::size_t k, std::uint64_t seed,
402 math::Pool pool, NDArray<T, 2, Layout::Contig> &outCentroids) {
403 const std::size_t n = X.dim(0);
404 const std::size_t d = X.dim(1);
405
406 CLUSTERING_ALWAYS_ASSERT(outCentroids.isMutable());
407 CLUSTERING_ALWAYS_ASSERT(outCentroids.dim(0) == k);
408 CLUSTERING_ALWAYS_ASSERT(outCentroids.dim(1) == d);
411
412 (void)pool;
413
414 const std::size_t nLocalTrials = detail::greedyKmppLocalTrials(k);
415 // Per-worker score slab padded to a cache line so adjacent workers do not false-share it
416 // while accumulating. An unpadded slab packs several workers onto one line and serialises
417 // the sweep on cross-die coherence traffic.
418 constexpr std::size_t kScoreSlabFloats = 16; // 64 bytes / sizeof(float)
419 const std::size_t scoreSlab =
420 ((nLocalTrials + kScoreSlabFloats - 1) / kScoreSlabFloats) * kScoreSlabFloats;
421 ensureShape(n, d, nLocalTrials, pool.workerCount());
422
423 math::pcg64 rng;
424 rng.seed(seed);
425
426 const T *xData = X.data();
427 T *centroidsData = outCentroids.data();
428 T *minSq = m_minSq.data();
429 T *candRowsData = m_candRows.data();
430 T *sweepSums = m_sweepSums.data();
431
432 // GEMM scoring wins only when the candidate width L is >= one kNr panel (6). Below that
433 // the 8x6 kernel's fixed 48-FMA body over-computes the 8xL useful tile; the per-row
434 // streaming kernel with L parallel accumulators is tighter. Gate on L >= kNr<float>.
435 constexpr std::size_t kNrF = math::detail::kKernelNr<float>;
436 const bool useGemmScoring = (d >= 32) && (nLocalTrials >= kNrF);
437 if (useGemmScoring) {
438 T *xNormsData = m_xNormsSq.data();
439 for (std::size_t i = 0; i < n; ++i) {
440 xNormsData[i] = math::detail::sqNormRow<T, Layout::Contig>(X, i);
441 }
442 }
443
444 // Step 1: first centroid uniformly. randUniformU64 is the deterministic primitive; the
445 // modulo map carries a tiny bias for very large n but is the standard sklearn convention.
446 const auto first = static_cast<std::size_t>(math::randUniformU64(rng) % n);
447 std::memcpy(centroidsData, xData + (first * d), d * sizeof(T));
448
449 // Every sweep that mutates minSq banks its block's weight into sweepSums, so candidate
450 // sampling never needs a separate pass over the array. The block partition is the
451 // deterministic sliceLo split of parallelForExactBlocksWithSlot.
452 const std::size_t sweepBlocks = detail::greedyKmppSweepBlocks(pool, n, d + 8);
453
454#ifdef CLUSTERING_USE_AVX2
455 if constexpr (std::is_same_v<T, float>) {
456 // The init sweep and every pick round run inside one persistent-worker plex: workers
457 // stay spin-resident across the per-round scoring and refresh passes while the serial
458 // glue rides the pre-phase hook, dropping the two fork-joins each round pays in the
459 // dispatch loop below. GEMM-scoring shapes keep the dispatch loop -- their scoring is
460 // one whole-matrix GEMM, not a range-invocable body -- as do runs too small to
461 // amortize the per-phase epoch cost.
462 constexpr std::size_t kMinPlexElems = std::size_t{1} << 16;
463 if (pool.pool != nullptr && pool.workerCount() > 1 && k >= 2 && !useGemmScoring &&
464 (n * d >= kMinPlexElems)) {
465 runPlexRounds(X, k, nLocalTrials, scoreSlab, sweepBlocks, rng, pool, outCentroids);
466 return;
467 }
468 }
469#endif
470
471 {
472 const T *firstRow = centroidsData;
474 std::size_t{0}, n, sweepBlocks,
475 [&](std::size_t lo, std::size_t hi, std::size_t s) noexcept {
476 initSweepBlock(firstRow, xData, d, lo, hi, s);
477 });
478 }
479
480 if (k == 1) {
481 return;
482 }
483
484 std::vector<std::size_t> candidates(nLocalTrials, 0);
485 std::vector<T> scores(nLocalTrials, T{0});
486
487 for (std::size_t c = 1; c < k; ++c) {
488 // The fold visits blocks in slot order, so the total is deterministic no matter which
489 // worker refreshed which block.
490 T total = T{0};
491 for (std::size_t s = 0; s < sweepBlocks; ++s) {
492 total += sweepSums[s];
493 }
494
495 // Degenerate guard: when every chosen centroid coincides with every remaining point the
496 // total collapses to ~0; pick the next centroid uniformly so the routine cannot stall.
497 if (!(total > T{0})) {
498 const auto pick = static_cast<std::size_t>(math::randUniformU64(rng) % n);
499 std::memcpy(centroidsData + (c * d), xData + (pick * d), d * sizeof(T));
500 const T *cRow = centroidsData + (c * d);
502 std::size_t{0}, n, sweepBlocks,
503 [&](std::size_t lo, std::size_t hi, std::size_t s) noexcept {
504 refreshSweepBlock(cRow, xData, d, lo, hi, s);
505 });
506 continue;
507 }
508
509 // Draw nLocalTrials candidates by inverse-CDF sampling. An empty or zero-weight block
510 // can never straddle a draw in `[0, total)` because the block fold above accumulates in
511 // the same order. Identical seed + identical n produces identical candidate sets.
512 drawCandidates(rng, total, n, sweepBlocks, nLocalTrials, candidates.data());
513
514 // Pack the L candidate rows into a contiguous (L, d) buffer so the batched scoring kernel
515 // can stream x once across L accumulators. The L*d pack is negligible against the n-pass
516 // scoring it amortizes.
517 packCandidates(xData, d, nLocalTrials, candidates.data());
518
519 for (std::size_t t = 0; t < nLocalTrials; ++t) {
520 scores[t] = T{0};
521 }
522 constexpr std::size_t kMaxLocalTrials = 32;
523 CLUSTERING_ALWAYS_ASSERT(nLocalTrials <= kMaxLocalTrials);
524
525 const std::size_t transposedWidth = detail::greedyKmppTransposedWidth(nLocalTrials);
526 bool scoredViaTransposed = false;
527#ifdef CLUSTERING_USE_AVX2
528 // Low-d hot path: at d <= kAvx2Lanes the (L, d) row-batched kernel either falls into the
529 // scalar K-tail (d < 8) or pays @c L horizontal-sum reductions for one K-iter of work
530 // (d == 8). The transposed `(d, W)` layout puts the same-feature components of every
531 // candidate in consecutive 8-lane YMM registers, so each broadcast-of-x[k] + FMA pair
532 // folds 8 (or 16, for the 16-lane unroll) distances at once.
533 if constexpr (std::is_same_v<T, float>) {
534 if (d > 0 && d <= math::detail::kAvx2Lanes<float>) {
535 packCandidatesTransposed(d, nLocalTrials, transposedWidth);
536 // Per-worker score accumulators are reused from @ref m_localScores; the candDistSq
537 // writes are row-local so partitioning by `i` is aliasing-free. The fan-out width
538 // is capped by the sweep's kernel work; a width of one keeps the sweep on the
539 // calling thread.
540 const std::size_t blocksT = detail::greedyKmppSweepBlocks(pool, n, d * transposedWidth);
541 const bool willParallelizeT = blocksT > 1;
542 const std::size_t workersT = willParallelizeT ? pool.workerCount() : std::size_t{1};
543 T *localScoresT = m_localScores.data();
544 zeroScoreSlabs(workersT, scoreSlab);
545
546 if (transposedWidth == 16) {
547 if (willParallelizeT) {
548 pool.parallelForBlocks(std::size_t{0}, n, blocksT,
549 [&](std::size_t lo, std::size_t hi) {
550 const std::size_t w = math::Pool::workerIndex();
551 scoreTransposed16Range(xData, d, nLocalTrials, lo, hi,
552 localScoresT + (w * scoreSlab));
553 });
554 } else {
555 scoreTransposed16Range(xData, d, nLocalTrials, std::size_t{0}, n, localScoresT);
556 }
557 } else if (transposedWidth == 8) {
558 if (willParallelizeT) {
559 pool.parallelForBlocks(std::size_t{0}, n, blocksT,
560 [&](std::size_t lo, std::size_t hi) {
561 const std::size_t w = math::Pool::workerIndex();
562 scoreTransposed8Range(xData, d, nLocalTrials, lo, hi,
563 localScoresT + (w * scoreSlab));
564 });
565 } else {
566 scoreTransposed8Range(xData, d, nLocalTrials, std::size_t{0}, n, localScoresT);
567 }
568 } else {
569 // Generic chunked path for L > 16 (very high k). Walk the transposed layout 8 lanes
570 // at a time so each chunk stays on the fully unrolled 8-wide kernel.
571 (void)candDistRows(n, transposedWidth);
572 if (willParallelizeT) {
574 std::size_t{0}, n, blocksT, [&](std::size_t lo, std::size_t hi) {
575 const std::size_t w = math::Pool::workerIndex();
576 scoreTransposedChunkedRange(xData, d, nLocalTrials, transposedWidth, lo, hi,
577 localScoresT + (w * scoreSlab));
578 });
579 } else {
580 scoreTransposedChunkedRange(xData, d, nLocalTrials, transposedWidth, std::size_t{0},
581 n, localScoresT);
582 }
583 }
584 foldScoreSlabs(workersT, scoreSlab, nLocalTrials, scores.data());
585 scoredViaTransposed = true;
586 }
587 }
588#endif
589
590 if (!scoredViaTransposed) {
591 // GEMM-based batch distance for moderate-to-high d: compute X * cand^T via the core
592 // GEMM (alpha=-2, beta=0), then add pre-computed per-row ||x||^2 and per-candidate
593 // ||c||^2 in one min+sum fold. BLAS-style GEMM is the decisive win at d >= ~16 where
594 // the per-row streaming kernel bottlenecks on L1/L2 bandwidth.
595 if (useGemmScoring) {
596 auto candView = NDArray<T, 2, Layout::Contig>::borrow(candRowsData, {nLocalTrials, d});
597 auto xView = NDArray<T, 2, Layout::Contig>::borrow(const_cast<T *>(xData), {n, d});
598 auto distsView = NDArray<T, 2>::borrow(m_distsFlat.data(), {n, nLocalTrials});
599 auto candT = candView.t();
600 // Direct gemmRunReference with caller-owned scratch so the seeder's per-pick GEMM
601 // leaves the shape-stable allocation footprint in place (no per-call arena alloc).
602 const auto xDesc = ::clustering::detail::describeMatrix(xView);
603 const auto candDesc = ::clustering::detail::describeMatrix(candT);
604 auto distsDesc = ::clustering::detail::describeMatrixMut(distsView);
605 math::detail::gemmRunReference<T>(xDesc, candDesc, distsDesc, T{-2}, T{0},
606 m_gemmApArena.data(), m_gemmBpArena.data(), pool);
607 // Candidate norms once per pick.
608 T *candNorms = m_candNormsSq.data();
609 for (std::size_t t = 0; t < nLocalTrials; ++t) {
610 candNorms[t] = math::detail::sqNormRow<T, Layout::Contig>(candView, t);
611 }
612 const T *xNorms = m_xNormsSq.data();
613 const T *distsFlat = m_distsFlat.data();
614 T *candDistSqData = candDistRows(n, transposedWidth);
615 for (std::size_t i = 0; i < n; ++i) {
616 const T mi = minSq[i];
617 const T xn = xNorms[i];
618 const T *distRowI = distsFlat + (i * nLocalTrials);
619 T *dstRow = candDistSqData + (i * transposedWidth);
620 for (std::size_t t = 0; t < nLocalTrials; ++t) {
621 T v = distRowI[t] + xn + candNorms[t];
622 if (v < T{0}) {
623 v = T{0};
624 }
625 dstRow[t] = v;
626 scores[t] += (v < mi) ? v : mi;
627 }
628 }
629 } else {
630 // Fused scoring: for each x row, compute L distances against the candidate pack and
631 // update L parallel running sums in one pass. The single-x-stream path is the load-
632 // bearing win at envelope shapes where n*d far exceeds L2 -- one stream is the
633 // difference between bandwidth-bound and bandwidth-bound times L. Parallelized over
634 // X rows via per-worker score slabs reduced at the end; candDistSqData writes are
635 // row-local so no aliasing across workers.
636 const bool willParallelize = pool.shouldParallelize(n, 1024, 2);
637 bool scoredViaSoa = false;
638#ifdef CLUSTERING_USE_AVX2
639 if constexpr (std::is_same_v<T, float>) {
640 // SoA 8-row M-tile kernel: streams X AoS through an in-register 8x8 transpose so 8
641 // rows' features land in feature-major YMM accumulators, folds L distances per row
642 // without per-row horizontal reductions, writes the per-(row, cand) distances to
643 // @c outDist, and accumulates min-capped scores. The kernel handles arbitrary row
644 // counts, so per-worker row ranges slot in under the same parallel fan-out that
645 // feeds the fallback path.
646 // Score-only path: skip the (n, L) cand-dist materialization. The winner row
647 // is recomputed by the commit-step refresh below, trading L stores per row for
648 // one sqEuclideanRowPtr per row at pick time.
649 const bool soaEligible = (d >= 8) && (nLocalTrials >= 1) && (nLocalTrials <= 6);
650 if (soaEligible) {
651 if (willParallelize) {
652 const std::size_t workers = pool.workerCount();
653 T *localScores = m_localScores.data();
654 zeroScoreSlabs(workers, scoreSlab);
655 pool.parallelForBlocks(std::size_t{0}, n, pool.stealBlocks(n),
656 [&](std::size_t lo, std::size_t hi) {
657 const std::size_t w = math::Pool::workerIndex();
658 scoreSoaRange(xData, d, nLocalTrials, transposedWidth, lo,
659 hi, localScores + (w * scoreSlab));
660 });
661 foldScoreSlabs(workers, scoreSlab, nLocalTrials, scores.data());
662 } else {
663 scoreSoaRange(xData, d, nLocalTrials, transposedWidth, std::size_t{0}, n,
664 scores.data());
665 }
666 scoredViaSoa = true;
667 }
668 }
669#endif
670 if (!scoredViaSoa) {
671 (void)candDistRows(n, transposedWidth);
672 if (willParallelize) {
673 const std::size_t workers = pool.workerCount();
674 T *localScores = m_localScores.data();
675 zeroScoreSlabs(workers, scoreSlab);
676 pool.parallelForBlocks(std::size_t{0}, n, pool.stealBlocks(n),
677 [&](std::size_t lo, std::size_t hi) {
678 const std::size_t w = math::Pool::workerIndex();
679 scoreScalarRange(xData, d, nLocalTrials, transposedWidth, lo,
680 hi, localScores + (w * scoreSlab));
681 });
682 foldScoreSlabs(workers, scoreSlab, nLocalTrials, scores.data());
683 } else {
684 scoreScalarRange(xData, d, nLocalTrials, transposedWidth, std::size_t{0}, n,
685 scores.data());
686 }
687 }
688 }
689 }
690
691 std::size_t bestT = 0;
692 T bestScore = scores[0];
693 for (std::size_t t = 1; t < nLocalTrials; ++t) {
694 if (scores[t] < bestScore) {
695 bestScore = scores[t];
696 bestT = t;
697 }
698 }
699 const std::size_t bestCandidate = candidates[bestT];
700
701 // Commit best candidate: copy its row into outCentroids, then refresh @c minSq with a
702 // fresh O(n*d) scan against the winner row. We deliberately DO NOT materialize the
703 // full (n, L) candidate-distance plane during the score sweep -- at the d=2 envelope
704 // its per-row plane write traffic dominated the seeder's runtime. Recomputing one
705 // column for the winner trades 1 row-distance call per row against L row-distance
706 // writes per row in the score sweep.
707 const T *winnerRow = xData + (bestCandidate * d);
708 std::memcpy(centroidsData + (c * d), winnerRow, d * sizeof(T));
710 std::size_t{0}, n, sweepBlocks,
711 [&](std::size_t lo, std::size_t hi, std::size_t s) noexcept {
712 refreshSweepBlock(winnerRow, xData, d, lo, hi, s);
713 });
714 }
715 }
716
717private:
725 T *candDistRows(std::size_t n, std::size_t w) {
726 if (m_candDistSq.dim(0) != n || m_candDistSq.dim(1) != w) {
727 m_candDistSq = NDArray<T, 2, Layout::Contig>({n == 0 ? std::size_t{1} : n, w});
728 }
729 return m_candDistSq.data();
730 }
731
734 void initSweepBlock(const T *firstRow, const T *xData, std::size_t d, std::size_t lo,
735 std::size_t hi, std::size_t s) noexcept {
736 T *minSq = m_minSq.data();
737 for (std::size_t i = lo; i < hi; ++i) {
738 minSq[i] = std::numeric_limits<T>::infinity();
739 }
740 m_sweepSums.data()[s] =
741 math::detail::refreshMinSqAgainstRow(firstRow, xData + (lo * d), hi - lo, d, minSq + lo);
742 }
743
746 void refreshSweepBlock(const T *row, const T *xData, std::size_t d, std::size_t lo,
747 std::size_t hi, std::size_t s) noexcept {
748 m_sweepSums.data()[s] = math::detail::refreshMinSqAgainstRow(row, xData + (lo * d), hi - lo, d,
749 m_minSq.data() + lo);
750 }
751
754 void drawCandidates(math::pcg64 &rng, T total, std::size_t n, std::size_t sweepBlocks,
755 std::size_t nLocalTrials, std::size_t *candidates) noexcept {
756 const T *sweepSums = m_sweepSums.data();
757 const T *minSq = m_minSq.data();
758 for (std::size_t t = 0; t < nLocalTrials; ++t) {
759 const T u = math::randUnit<T>(rng) * total;
760 std::size_t s = 0;
761 T acc = T{0};
762 while (s + 1 < sweepBlocks && acc + sweepSums[s] <= u) {
763 acc += sweepSums[s];
764 ++s;
765 }
766 candidates[t] = math::detail::inverseCdfPickInRange(minSq, (n * s) / sweepBlocks,
767 (n * (s + 1)) / sweepBlocks, u - acc);
768 }
769 }
770
772 void packCandidates(const T *xData, std::size_t d, std::size_t nLocalTrials,
773 const std::size_t *candidates) noexcept {
774 T *candRowsData = m_candRows.data();
775 for (std::size_t t = 0; t < nLocalTrials; ++t) {
776 std::memcpy(candRowsData + (t * d), xData + (candidates[t] * d), d * sizeof(T));
777 }
778 }
779
782 void packCandidatesTransposed(std::size_t d, std::size_t nLocalTrials,
783 std::size_t transposedWidth) noexcept {
784 const T *candRowsData = m_candRows.data();
785 T *candRowsTData = m_candRowsT.data();
786 for (std::size_t kk = 0; kk < d; ++kk) {
787 T *dstK = candRowsTData + (kk * transposedWidth);
788 for (std::size_t t = 0; t < nLocalTrials; ++t) {
789 dstK[t] = candRowsData[(t * d) + kk];
790 }
791 for (std::size_t t = nLocalTrials; t < transposedWidth; ++t) {
792 dstK[t] = T{0};
793 }
794 }
795 }
796
797 void zeroScoreSlabs(std::size_t slabs, std::size_t scoreSlab) noexcept {
798 T *localScores = m_localScores.data();
799 for (std::size_t e = 0; e < slabs * scoreSlab; ++e) {
800 localScores[e] = T{0};
801 }
802 }
803
806 void foldScoreSlabs(std::size_t slabs, std::size_t scoreSlab, std::size_t nLocalTrials,
807 T *scores) const noexcept {
808 const T *localScores = m_localScores.data();
809 for (std::size_t w = 0; w < slabs; ++w) {
810 const T *row = localScores + (w * scoreSlab);
811 for (std::size_t t = 0; t < nLocalTrials; ++t) {
812 scores[t] += row[t];
813 }
814 }
815 }
816
820 void scoreScalarRange(const T *xData, std::size_t d, std::size_t nLocalTrials,
821 std::size_t transposedWidth, std::size_t lo, std::size_t hi,
822 T *dst) noexcept {
823 const T *candRowsData = m_candRows.data();
824 const T *minSq = m_minSq.data();
825 T *candDistSqData = m_candDistSq.data();
826 std::array<T, 32> distRowLocal{};
827 for (std::size_t i = lo; i < hi; ++i) {
828 const T *xi = xData + (i * d);
829 const T mi = minSq[i];
830 detail::sqEuclideanRowToBatch<T>(xi, candRowsData, nLocalTrials, d, distRowLocal.data());
831 T *dstRow = candDistSqData + (i * transposedWidth);
832 for (std::size_t t = 0; t < nLocalTrials; ++t) {
833 dstRow[t] = distRowLocal[t];
834 dst[t] += (distRowLocal[t] < mi) ? distRowLocal[t] : mi;
835 }
836 }
837 }
838
839#ifdef CLUSTERING_USE_AVX2
842 void scoreTransposed16Range(const T *xData, std::size_t d, std::size_t nLocalTrials,
843 std::size_t lo, std::size_t hi, T *dst) noexcept {
844 const T *minSq = m_minSq.data();
845 const T *candRowsTData = m_candRowsT.data();
846 __m256 scoresLoAcc = _mm256_setzero_ps();
847 __m256 scoresHiAcc = _mm256_setzero_ps();
848 for (std::size_t i = lo; i < hi; ++i) {
849 const float *xi = xData + (i * d);
850 const __m256 miVec = _mm256_set1_ps(minSq[i]);
851 const auto [dLo, dHi] = detail::sqEuclideanRowAgainst16TransposedReg(xi, candRowsTData, d);
852 scoresLoAcc = _mm256_add_ps(scoresLoAcc, _mm256_min_ps(dLo, miVec));
853 scoresHiAcc = _mm256_add_ps(scoresHiAcc, _mm256_min_ps(dHi, miVec));
854 }
855 std::array<float, 16> tmp{};
856 _mm256_storeu_ps(tmp.data(), scoresLoAcc);
857 _mm256_storeu_ps(tmp.data() + 8, scoresHiAcc);
858 for (std::size_t t = 0; t < nLocalTrials; ++t) {
859 dst[t] += tmp[t];
860 }
861 }
862
864 void scoreTransposed8Range(const T *xData, std::size_t d, std::size_t nLocalTrials,
865 std::size_t lo, std::size_t hi, T *dst) noexcept {
866 const T *minSq = m_minSq.data();
867 const T *candRowsTData = m_candRowsT.data();
868 __m256 scoresAcc = _mm256_setzero_ps();
869 for (std::size_t i = lo; i < hi; ++i) {
870 const float *xi = xData + (i * d);
871 const __m256 miVec = _mm256_set1_ps(minSq[i]);
872 const __m256 dv = detail::sqEuclideanRowAgainst8TransposedReg(xi, candRowsTData, d);
873 scoresAcc = _mm256_add_ps(scoresAcc, _mm256_min_ps(dv, miVec));
874 }
875 std::array<float, 8> tmp{};
876 _mm256_storeu_ps(tmp.data(), scoresAcc);
877 for (std::size_t t = 0; t < nLocalTrials; ++t) {
878 dst[t] += tmp[t];
879 }
880 }
881
884 void scoreTransposedChunkedRange(const T *xData, std::size_t d, std::size_t nLocalTrials,
885 std::size_t transposedWidth, std::size_t lo, std::size_t hi,
886 T *dst) noexcept {
887 const T *minSq = m_minSq.data();
888 const T *candRowsTData = m_candRowsT.data();
889 T *candDistSqData = m_candDistSq.data();
890 for (std::size_t i = lo; i < hi; ++i) {
891 const float *xi = xData + (i * d);
892 const float mi = minSq[i];
893 float *dstRow = candDistSqData + (i * transposedWidth);
894 for (std::size_t base = 0; base < transposedWidth; base += 8) {
895 detail::sqEuclideanRowAgainst8TransposedStrided(xi, candRowsTData + base, d,
896 transposedWidth, dstRow + base);
897 }
898 for (std::size_t t = 0; t < nLocalTrials; ++t) {
899 dst[t] += (dstRow[t] < mi) ? dstRow[t] : mi;
900 }
901 }
902 }
903
906 void scoreSoaRange(const T *xData, std::size_t d, std::size_t nLocalTrials,
907 std::size_t transposedWidth, std::size_t lo, std::size_t hi, T *dst) noexcept {
908 const std::size_t rangeN = hi - lo;
909 const float *xSlice = xData + (lo * d);
910 const float *minSlice = m_minSq.data() + lo;
911 const float *candRowsData = m_candRows.data();
912 switch (nLocalTrials) {
913 case 1:
914 math::detail::kmppScoreSoaRowsAvx2F32<1, /*WriteOutDist=*/false>(
915 xSlice, rangeN, d, candRowsData, minSlice, nullptr, transposedWidth, dst);
916 break;
917 case 2:
918 math::detail::kmppScoreSoaRowsAvx2F32<2, /*WriteOutDist=*/false>(
919 xSlice, rangeN, d, candRowsData, minSlice, nullptr, transposedWidth, dst);
920 break;
921 case 3:
922 math::detail::kmppScoreSoaRowsAvx2F32<3, /*WriteOutDist=*/false>(
923 xSlice, rangeN, d, candRowsData, minSlice, nullptr, transposedWidth, dst);
924 break;
925 case 4:
926 math::detail::kmppScoreSoaRowsAvx2F32<4, /*WriteOutDist=*/false>(
927 xSlice, rangeN, d, candRowsData, minSlice, nullptr, transposedWidth, dst);
928 break;
929 case 5:
930 math::detail::kmppScoreSoaRowsAvx2F32<5, /*WriteOutDist=*/false>(
931 xSlice, rangeN, d, candRowsData, minSlice, nullptr, transposedWidth, dst);
932 break;
933 case 6:
934 math::detail::kmppScoreSoaRowsAvx2F32<6, /*WriteOutDist=*/false>(
935 xSlice, rangeN, d, candRowsData, minSlice, nullptr, transposedWidth, dst);
936 break;
937 default:
938 break;
939 }
940 }
941
944 enum class ScoreKernel : std::uint8_t {
945 kTransposed16,
946 kTransposed8,
947 kTransposedChunked,
948 kSoa,
949 kScalar
950 };
951
954 [[nodiscard]] static ScoreKernel pickScoreKernel(std::size_t d,
955 std::size_t nLocalTrials) noexcept {
956 if (d > 0 && d <= math::detail::kAvx2Lanes<float>) {
957 const std::size_t w = detail::greedyKmppTransposedWidth(nLocalTrials);
958 if (w == 8) {
959 return ScoreKernel::kTransposed8;
960 }
961 if (w == 16) {
962 return ScoreKernel::kTransposed16;
963 }
964 return ScoreKernel::kTransposedChunked;
965 }
966 if (d >= 8 && nLocalTrials >= 1 && nLocalTrials <= 6) {
967 return ScoreKernel::kSoa;
968 }
969 return ScoreKernel::kScalar;
970 }
971
973 void scoreRange(ScoreKernel kernel, const T *xData, std::size_t d, std::size_t nLocalTrials,
974 std::size_t transposedWidth, std::size_t lo, std::size_t hi, T *dst) noexcept {
975 switch (kernel) {
976 case ScoreKernel::kTransposed16:
977 scoreTransposed16Range(xData, d, nLocalTrials, lo, hi, dst);
978 break;
979 case ScoreKernel::kTransposed8:
980 scoreTransposed8Range(xData, d, nLocalTrials, lo, hi, dst);
981 break;
982 case ScoreKernel::kTransposedChunked:
983 scoreTransposedChunkedRange(xData, d, nLocalTrials, transposedWidth, lo, hi, dst);
984 break;
985 case ScoreKernel::kSoa:
986 scoreSoaRange(xData, d, nLocalTrials, transposedWidth, lo, hi, dst);
987 break;
988 case ScoreKernel::kScalar:
989 scoreScalarRange(xData, d, nLocalTrials, transposedWidth, lo, hi, dst);
990 break;
991 }
992 }
993
1007 void runPlexRounds(const NDArray<T, 2, Layout::Contig> &X, std::size_t k,
1008 std::size_t nLocalTrials, std::size_t scoreSlab, std::size_t sweepBlocks,
1009 math::pcg64 &rng, math::Pool pool,
1010 NDArray<T, 2, Layout::Contig> &outCentroids) {
1011 const std::size_t n = X.dim(0);
1012 const std::size_t d = X.dim(1);
1013 const T *xData = X.data();
1014 T *centroidsData = outCentroids.data();
1015 const T *sweepSums = m_sweepSums.data();
1016 const std::size_t workers = pool.workerCount();
1017 const std::size_t transposedWidth = detail::greedyKmppTransposedWidth(nLocalTrials);
1018 const ScoreKernel kernel = pickScoreKernel(d, nLocalTrials);
1019 auto sweepLo = [n, sweepBlocks](std::size_t s) noexcept { return (n * s) / sweepBlocks; };
1020
1021 constexpr std::size_t kMaxLocalTrials = 32;
1022 CLUSTERING_ALWAYS_ASSERT(nLocalTrials <= kMaxLocalTrials);
1023
1024 // Plane-writing kernels grow the (n, W) plane before workers go plex-resident.
1025 if (kernel == ScoreKernel::kTransposedChunked || kernel == ScoreKernel::kScalar) {
1026 (void)candDistRows(n, transposedWidth);
1027 }
1028
1029 std::vector<std::size_t> candidates(nLocalTrials, 0);
1030 std::vector<T> scores(nLocalTrials, T{0});
1031 const T *refreshRow = centroidsData; // Phase 0 sweeps against the first centroid.
1032 bool roundDegenerate = false;
1033
1034 auto prePhase = [&](std::size_t phaseIdx) noexcept {
1035 if (phaseIdx == 0) {
1036 return;
1037 }
1038 const std::size_t c = (phaseIdx + 1) / 2;
1039 if ((phaseIdx & 1U) != 0) {
1040 // Scoring glue: fold the banked block sums, then either record the degenerate
1041 // uniform pick or draw and pack this round's candidates.
1042 T total = T{0};
1043 for (std::size_t s = 0; s < sweepBlocks; ++s) {
1044 total += sweepSums[s];
1045 }
1046 roundDegenerate = !(total > T{0});
1047 if (roundDegenerate) {
1048 const auto pick = static_cast<std::size_t>(math::randUniformU64(rng) % n);
1049 std::memcpy(centroidsData + (c * d), xData + (pick * d), d * sizeof(T));
1050 refreshRow = centroidsData + (c * d);
1051 return;
1052 }
1053 drawCandidates(rng, total, n, sweepBlocks, nLocalTrials, candidates.data());
1054 packCandidates(xData, d, nLocalTrials, candidates.data());
1055 if (kernel != ScoreKernel::kSoa && kernel != ScoreKernel::kScalar) {
1056 packCandidatesTransposed(d, nLocalTrials, transposedWidth);
1057 }
1058 zeroScoreSlabs(workers, scoreSlab);
1059 return;
1060 }
1061 // Refresh glue: fold the slot slabs, pick the winner, and stage its row for the sweep.
1062 if (roundDegenerate) {
1063 return;
1064 }
1065 for (std::size_t t = 0; t < nLocalTrials; ++t) {
1066 scores[t] = T{0};
1067 }
1068 foldScoreSlabs(workers, scoreSlab, nLocalTrials, scores.data());
1069 std::size_t bestT = 0;
1070 T bestScore = scores[0];
1071 for (std::size_t t = 1; t < nLocalTrials; ++t) {
1072 if (scores[t] < bestScore) {
1073 bestScore = scores[t];
1074 bestT = t;
1075 }
1076 }
1077 const T *winnerRow = xData + (candidates[bestT] * d);
1078 std::memcpy(centroidsData + (c * d), winnerRow, d * sizeof(T));
1079 refreshRow = winnerRow;
1080 };
1081
1082 auto phase = [&](std::size_t phaseIdx, std::uint32_t slot, std::size_t lo, std::size_t hi,
1083 void * /*tlsArena*/ = nullptr) noexcept {
1084 if (lo >= hi) {
1085 return;
1086 }
1087 if (phaseIdx == 0) {
1088 for (std::size_t s = lo; s < hi; ++s) {
1089 initSweepBlock(refreshRow, xData, d, sweepLo(s), sweepLo(s + 1), s);
1090 }
1091 return;
1092 }
1093 if ((phaseIdx & 1U) != 0) {
1094 if (roundDegenerate) {
1095 return;
1096 }
1097 scoreRange(kernel, xData, d, nLocalTrials, transposedWidth, sweepLo(lo), sweepLo(hi),
1098 m_localScores.data() + (static_cast<std::size_t>(slot) * scoreSlab));
1099 return;
1100 }
1101 for (std::size_t s = lo; s < hi; ++s) {
1102 refreshSweepBlock(refreshRow, xData, d, sweepLo(s), sweepLo(s + 1), s);
1103 }
1104 };
1105
1106 pool.parallelRunPlex<citor::HintsDefaults>(1 + (2 * (k - 1)), sweepBlocks, std::move(phase),
1107 std::move(prePhase));
1108 }
1109#endif // CLUSTERING_USE_AVX2
1110
1111 void ensureShape(std::size_t n, std::size_t d, std::size_t L, std::size_t workers) {
1112 const std::size_t w = detail::greedyKmppTransposedWidth(L == 0 ? std::size_t{1} : L);
1113 if (m_candRows.dim(0) != L || m_candRows.dim(1) != d) {
1114 m_candRows = NDArray<T, 2, Layout::Contig>({L, d});
1115 }
1116 if (m_candRowsT.dim(0) != d || m_candRowsT.dim(1) != w) {
1117 m_candRowsT = NDArray<T, 2, Layout::Contig>({d == 0 ? std::size_t{1} : d, w});
1118 }
1119 const std::size_t sweepSlots = std::max<std::size_t>(workers, std::size_t{1}) * 8;
1120 if (m_sweepSums.dim(0) != sweepSlots) {
1121 m_sweepSums = NDArray<T, 1>({sweepSlots});
1122 }
1123 if (m_minSq.dim(0) != n) {
1124 m_minSq = NDArray<T, 1>({n});
1125 }
1126 // GEMM-scoring-only scratch (distsFlat, xNormsSq, candNormsSq, gemmApArena, gemmBpArena).
1127 // The GEMM path fires at `d >= 32` && L >= kKernelNr<float>; outside that envelope we keep
1128 // unit-sized placeholders so @c .data() stays dereferenceable without paying the @c kKc*kNc
1129 // envelope tax (@c Bp alone is several MB).
1130 constexpr std::size_t kNrForGemm = math::detail::kKernelNr<float>;
1131 const bool gemmScoringUsed = std::is_same_v<T, float> && (d >= 32) && (L >= kNrForGemm);
1132 const std::size_t nSafe = (n == 0) ? std::size_t{1} : n;
1133 const std::size_t lSafe = (L == 0) ? std::size_t{1} : L;
1134 const std::size_t distsFlatRows = gemmScoringUsed ? nSafe : std::size_t{1};
1135 const std::size_t distsFlatCols = gemmScoringUsed ? lSafe : std::size_t{1};
1136 if (m_distsFlat.dim(0) != distsFlatRows || m_distsFlat.dim(1) != distsFlatCols) {
1137 m_distsFlat = NDArray<T, 2, Layout::Contig>({distsFlatRows, distsFlatCols});
1138 }
1139 const std::size_t xNormsLen = gemmScoringUsed ? nSafe : std::size_t{1};
1140 if (m_xNormsSq.dim(0) != xNormsLen) {
1141 m_xNormsSq = NDArray<T, 1>({xNormsLen});
1142 }
1143 const std::size_t candNormsLen = gemmScoringUsed ? lSafe : std::size_t{1};
1144 if (m_candNormsSq.dim(0) != candNormsLen) {
1145 m_candNormsSq = NDArray<T, 1>({candNormsLen});
1146 }
1147 const std::size_t workersClamped = workers == 0 ? std::size_t{1} : workers;
1148 // @c gemmRunReference parallelizes the Mc-tile loop, with each worker owning a per-worker
1149 // slice of the A-pack arena at offset `(worker * kMc * kKc)`. Sizing the arena for just
1150 // one worker was fine while the seeder's envelope kept the GEMM path off (k=16, L=4 fell
1151 // into the SoA kernel), but the Elkan-eligible shapes push L >= kNrF where the GEMM scoring
1152 // activates and multiple workers collide into the same slice.
1153 const std::size_t apSize = gemmScoringUsed
1154 ? (workersClamped * math::detail::kMc<T> * math::detail::kKc<T>)
1155 : std::size_t{1};
1156 const std::size_t bpSize =
1157 gemmScoringUsed ? (math::detail::kKc<T> * math::detail::kNc<T>) : std::size_t{1};
1158 if (m_gemmApArena.dim(0) != apSize) {
1159 m_gemmApArena = NDArray<T, 1>({apSize});
1160 }
1161 if (m_gemmBpArena.dim(0) != bpSize) {
1162 m_gemmBpArena = NDArray<T, 1>({bpSize});
1163 }
1164 const std::size_t scoreSlab = ((lSafe + 15U) / 16U) * 16U; // cache-line-padded per-worker slab
1165 const std::size_t lsLen = workersClamped * scoreSlab;
1166 if (m_localScores.dim(0) != lsLen) {
1167 m_localScores = NDArray<T, 1>({lsLen});
1168 }
1169 }
1170
1172 NDArray<T, 2, Layout::Contig> m_candRows;
1177 NDArray<T, 2, Layout::Contig> m_candRowsT;
1180 NDArray<T, 2, Layout::Contig> m_candDistSq;
1183 NDArray<T, 1> m_sweepSums;
1186 NDArray<T, 1> m_minSq;
1188 NDArray<T, 2, Layout::Contig> m_distsFlat;
1190 NDArray<T, 1> m_xNormsSq;
1192 NDArray<T, 1> m_candNormsSq;
1194 NDArray<T, 1> m_gemmApArena;
1196 NDArray<T, 1> m_gemmBpArena;
1199 NDArray<T, 1> m_localScores;
1200};
1201
1202} // namespace clustering::kmeans
#define CLUSTERING_ALWAYS_ASSERT(cond)
Release-active assertion: evaluates cond in every build configuration.
Represents a multidimensional array (NDArray) of a fixed number of dimensions N and element type T.
Definition ndarray.h:136
size_t dim(std::size_t index) const noexcept
Returns the size of a specific dimension of the NDArray.
Definition ndarray.h:462
static NDArray borrow(T *ptr, std::array< std::size_t, N > shape) noexcept
Borrows a contiguous buffer as an NDArray without taking ownership.
Definition ndarray.h:571
const T * data() const noexcept
Provides read-only access to the internal data array.
Definition ndarray.h:504
bool isMutable() const noexcept
Reports whether writes through operator(), Accessor, or flatIndex are allowed.
Definition ndarray.h:489
void run(const NDArray< T, 2, Layout::Contig > &X, std::size_t k, std::uint64_t seed, math::Pool pool, NDArray< T, 2, Layout::Contig > &outCentroids)
Seed k centroids from X into outCentroids.
void sqEuclideanRowAgainst8TransposedStrided(const float *x, const float *candData, std::size_t d, std::size_t rowStride, float *out) noexcept
Compute one 8-way squared distance slab against an (d, W) transposed candidate layout with an explici...
void sqEuclideanRowToBatchAvx2Fixed(const float *x, const float *candData, std::size_t d, float *out) noexcept
Compile-time batched scoring kernel: stream x once across B parallel AVX2 accumulators to compute B s...
void sqEuclideanRowToBatchAvx2(const T *x, const T *candData, std::size_t L, std::size_t d, T *out) noexcept
Compute L squared Euclidean distances against an (L, d) row-batched candidate layout in a single stre...
void sqEuclideanRowAgainst8Transposed(const float *x, const float *candData, std::size_t d, float *out) noexcept
Compute L squared distances against an (d, 8) transposed candidate layout with one streaming pass ove...
constexpr std::size_t greedyKmppTransposedWidth(std::size_t L) noexcept
Round L up to the nearest multiple of 8 used by the transposed scoring layout.
__m256 sqEuclideanRowAgainst8TransposedReg(const float *x, const float *candData, std::size_t d) noexcept
Register-only variant of sqEuclideanRowAgainst8Transposed.
void sqEuclideanRowAgainst16Transposed(const float *x, const float *candData, std::size_t d, float *out) noexcept
Compute two 8-way squared distance slabs against an (d, 16) transposed candidate layout in one stream...
std::size_t greedyKmppSweepBlocks(math::Pool pool, std::size_t rows, std::size_t opsPerRow) noexcept
Fan-out width for one of the seeder's per-round O(n*d) sweeps.
std::pair< __m256, __m256 > sqEuclideanRowAgainst16TransposedReg(const float *x, const float *candData, std::size_t d) noexcept
Register-only 16-wide variant of sqEuclideanRowAgainst16Transposed.
std::size_t greedyKmppLocalTrials(std::size_t k) noexcept
Compute the local-trials count used by greedy k-means++.
void sqEuclideanRowToBatch(const T *x, const T *candData, std::size_t L, std::size_t d, T *out) noexcept
Squared Euclidean distance from one x row to a batch of L candidate rows.
float horizontalSumAvx2(__m256 v) noexcept
Definition pairwise.h:42
constexpr std::size_t kAvx2Lanes
Definition pairwise.h:131
T sqNormRow(const NDArray< T, 2, LX > &X, std::size_t i) noexcept
Definition pairwise.h:205
T randUnit(Rng &rng) noexcept
Draw a uniform variate in the half-open unit interval [0, 1).
Definition rng.h:152
std::uint64_t randUniformU64(Rng &rng) noexcept
Draw a 64-bit unsigned integer uniformly at random from the full u64 range.
Definition rng.h:139
Thin compile-time-templated wrapper around the underlying OwnedPool.
Definition thread.h:109
static std::size_t workerIndex() noexcept
Stable index of the calling worker thread within the owning pool.
Definition thread.h:131
OwnedPool * pool
Underlying pool, or nullptr to force serial execution.
Definition thread.h:111
std::size_t workerCount() const noexcept
Number of worker threads available, or 1 in serial mode.
Definition thread.h:118
void parallelForBlocks(std::size_t first, std::size_t last, std::size_t numBlocks, Body body)
Run body in parallel over [first, last) partitioned into numBlocks blocks.
Definition thread.h:239
std::size_t stealBlocks(std::size_t n, std::size_t minRowsPerBlock=256) const noexcept
Block count for a fan-out over n rows that lets work-stealing balance heterogeneous cores.
Definition thread.h:191
bool shouldParallelize(std::size_t totalWork, std::size_t minChunk, std::size_t minTasksPerWorker=2) const noexcept
Decide whether totalWork warrants parallel dispatch.
Definition thread.h:147
void parallelForExactBlocksWithSlot(std::size_t first, std::size_t last, std::size_t numBlocks, Body body)
Slot-aware variant of parallelForExactBlocks.
Definition thread.h:307
128-bit state for the PCG-XSL-RR 64-bit output generator (Melissa O'Neill).
Definition rng.h:30
void seed(std::uint64_t seedValue, std::uint64_t stream=0) noexcept
Initialize the generator per PCG's canonical seeding procedure.
Definition rng.h:46