101 static_assert(std::is_same_v<T, float> || std::is_same_v<T, double>,
102 "LloydFusedGemm<T> requires T to be float or double");
105 : m_centroidsOld({0, 0}), m_cSqNorms({0}), m_sums({0, 0}), m_counts({0}), m_minDistSq({0}),
106 m_shiftSq({0}), m_partialSums({0}), m_partialComps({0}), m_partialCounts({0}),
107 m_foldComp({0}), m_packedB({0}), m_packedCSqNorms({0}), m_distsChunk({0, 0}),
108 m_gemmApArena({0}), m_packedXAp({0}), m_xNormsSq({0}), m_varSum({0}), m_varSumSq({0}),
109 m_varPartialSum({0}), m_varPartialSumSq({0}), m_u({0}), m_l({0}), m_shiftEuclidean({0}),
110 m_halfDistToNearestOther({0}), m_elkanBounds({0, 0}), m_centerDist({0, 0}) {}
112#ifdef CLUSTERING_KMEANS_KAHAN_N_THRESHOLD
121 static constexpr std::size_t
kahanNThreshold = CLUSTERING_KMEANS_KAHAN_N_THRESHOLD;
148 std::size_t k, std::size_t maxIter, T tol,
math::Pool pool,
150 bool &outConverged) {
151 const std::size_t n = X.
dim(0);
152 const std::size_t d = X.
dim(1);
160 if (n == 0 || d == 0) {
167 const std::size_t workerCount = pool.
workerCount();
168 ensureShape(n, d, k, workerCount);
178 if (m_xstatsCachedXData == X.
data() && m_xstatsCachedN == n && m_xstatsCachedD == d) {
179 meanVar = m_xstatsCachedMeanVar;
181 meanVar = computeXStatistics(X, pool);
182 m_xstatsCachedXData = X.
data();
185 m_xstatsCachedMeanVar = meanVar;
187 const T shiftSqThreshold = tol * meanVar;
190 refreshCentroidSqNorms(centroids);
192 std::size_t iter = 0;
193 bool converged =
false;
203 const bool directHamerly =
204 (d * k >= kHamerlyMinDirectScanDims) && (workerCount <= kHamerlyDirectWorkerCap);
205 const bool hamerlyEligible =
212 (k <= kElkanMaxK) && (n * k <= kElkanNKLimit) && (k >= 2);
214 bool ranPlex =
false;
215#ifdef CLUSTERING_USE_AVX2
216 if constexpr (std::is_same_v<T, float>) {
224 constexpr std::size_t kMinPlexElems = std::size_t{1} << 16;
226 if (pool.
pool !=
nullptr && workerCount > 1 && maxIter > 0 && (n * d >= kMinPlexElems) &&
227 (assignmentProducesDirectMinDistSq(X, centroids) ||
228 assignmentUsesFusedArgmin(X, centroids) || plexChunkedHamerly)) {
229 runPlexLoop(X, centroids, outLabels, k, maxIter, shiftSqThreshold, useKahan,
230 hamerlyEligible, pool, iter, converged);
236 while (!ranPlex && iter < maxIter) {
237 std::memcpy(m_centroidsOld.data(), centroids.
data(),
238 centroids.
dim(0) * centroids.
dim(1) *
sizeof(T));
240 if (hamerlyEligible && iter > 0) {
241 runHamerlyAssignmentAndScatter(X, centroids, outLabels, k, useKahan, pool);
242 }
else if (elkanEligible && iter > 0) {
243 runElkanAssignmentAndScatter(X, centroids, outLabels, k, useKahan, pool);
245 runAssignmentAndScatter(X, centroids, outLabels, k, useKahan, pool);
246 if (hamerlyEligible && iter == 0 && assignmentUsesFusedArgmin(X, centroids)) {
247 seedHamerlyBoundsFromAssignedMinDist(outLabels, k, d, pool);
251 (void)::clustering::kmeans::detail::reseedEmptyClusters<T>(X, centroids, m_sums, m_counts,
253 finalizeMeans(centroids);
254 refreshCentroidSqNorms(centroids);
257 const T totalShift = ::clustering::kmeans::detail::totalShiftSqKahan<T>(m_shiftSq);
260 if (totalShift <= shiftSqThreshold) {
268 bool finalMinDistExact = assignmentProducesDirectMinDistSq(X, centroids);
269 if (hamerlyEligible && iter > 0) {
270 runHamerlyAssignment(X, centroids, outLabels, pool,
true);
271 finalMinDistExact =
true;
272 }
else if (elkanEligible && iter > 0) {
273 runElkanAssignment(X, centroids, outLabels,
math::Pool{});
275 runAssignment(X, centroids, outLabels, pool);
277 if (!finalMinDistExact) {
278 recomputeMinDistSqDirect(X, centroids, outLabels, pool);
281 outInertia = inertiaKahan(n, pool);
283 outConverged = converged;
289 struct HamerlyShiftTop2 {
292 std::size_t argMax = 0;
297 [[gnu::always_inline]]
void scatterRowToSlab(
const T *xBase,
const std::int32_t *labelsBase,
298 std::size_t i, std::size_t slot, std::size_t k,
299 std::size_t d,
bool useKahan)
noexcept {
300 const std::int32_t lbl = labelsBase[i];
301 if (lbl < 0 || std::cmp_greater_equal(lbl, k)) {
304 const auto row =
static_cast<std::size_t
>(lbl);
305 const T *xRow = xBase + (i * d);
306 T *sumRow = m_partialSums.data() + (((slot * k) + row) * d);
307 std::int32_t *cslab = m_partialCounts.data() + (slot * k);
309 T *compRow = m_partialComps.data() + (((slot * k) + row) * d);
310 math::detail::kahanAddRow<T>(xRow, d, sumRow, compRow);
312 for (std::size_t t = 0; t < d; ++t) {
313 sumRow[t] += xRow[t];
327 [[nodiscard]] T computeXStatistics(
const NDArray<T, 2, Layout::Contig> &X, math::Pool pool) {
328 const std::size_t n = X.dim(0);
329 const std::size_t d = X.dim(1);
330 if (n == 0 || d == 0) {
333 const T *xData = X.data();
335 if (m_varSum.dim(0) != d) {
336 m_varSum = NDArray<T, 1>({d});
337 m_varSumSq = NDArray<T, 1>({d});
339 T *colSum = m_varSum.data();
340 T *colSumSq = m_varSumSq.data();
342 const std::size_t workers = pool.workerCount();
343 const bool willParallelize = workers > 1;
344 if (willParallelize) {
345 const std::size_t partialSize = workers * d;
346 if (m_varPartialSum.dim(0) != partialSize) {
347 m_varPartialSum = NDArray<T, 1>({partialSize});
348 m_varPartialSumSq = NDArray<T, 1>({partialSize});
350 T *partialSum = m_varPartialSum.data();
351 T *partialSumSq = m_varPartialSumSq.data();
352 for (std::size_t e = 0; e < partialSize; ++e) {
353 partialSum[e] = T{0};
354 partialSumSq[e] = T{0};
356 pool.parallelForExactBlocksWithSlot<citor::HintsDefaults>(
357 std::size_t{0}, n, workers,
358 [&](std::size_t lo, std::size_t hi, std::size_t slot)
noexcept {
359 T *localSum = partialSum + (slot * d);
360 T *localSumSq = partialSumSq + (slot * d);
361 for (std::size_t i = lo; i < hi; ++i) {
362 const T *row = xData + (i * d);
363 math::detail::columnwiseAccumSumSq<T>(row, d, localSum, localSumSq);
367 for (std::size_t t = 0; t < d; ++t) {
370 for (std::size_t w = 0; w < workers; ++w) {
371 s += partialSum[(w * d) + t];
372 ss += partialSumSq[(w * d) + t];
378 for (std::size_t t = 0; t < d; ++t) {
382 for (std::size_t i = 0; i < n; ++i) {
383 const T *row = xData + (i * d);
384 math::detail::columnwiseAccumSumSq<T>(row, d, colSum, colSumSq);
389 const auto nInv =
static_cast<T
>(1) /
static_cast<T
>(n);
391 for (std::size_t t = 0; t < d; ++t) {
392 const T mean = colSum[t] * nInv;
393 acc += (colSumSq[t] * nInv) - (mean * mean);
395 return acc /
static_cast<T
>(d);
398 [[nodiscard]]
double inertiaKahan(std::size_t n, math::Pool pool) {
399 const T *minDist = m_minDistSq.data();
400 auto sumRange = [minDist](std::size_t lo, std::size_t hi)
noexcept {
403 for (std::size_t i = lo; i < hi; ++i) {
404 const auto addend =
static_cast<double>(minDist[i]);
405 const double y = addend - comp;
406 const double t =
sum + y;
407 comp = (t -
sum) - y;
412 return pool.parallelReduce<citor::HintsDefaults>(
413 std::size_t{0}, n, 0.0, sumRange,
414 [](
double lhs,
double rhs)
noexcept {
return lhs + rhs; });
417 [[nodiscard]]
static constexpr bool
418 packedXApCacheEnabledForShape(std::size_t n, std::size_t d, std::size_t workerCount)
noexcept {
419 if constexpr (std::is_same_v<T, float>) {
420 constexpr std::size_t kPackedXMaxElements = std::size_t{4} << 20;
421 constexpr std::size_t kPackedXMaxWorkers = 4;
423 d != 0 && n <= (kPackedXMaxElements / d);
432 void resetPackedXApCacheMetadata() noexcept {
433 m_packedXApCachedXData =
nullptr;
434 m_packedXApCachedN = 0;
435 m_packedXApCachedD = 0;
438 void ensureShape(std::size_t n, std::size_t d, std::size_t k, std::size_t workerCount) {
439 const bool shapeChanged = (n != m_n) || (d != m_d) || (k != m_k);
440 const bool workerChanged = (workerCount != m_workerCount);
441 if (!shapeChanged && !workerChanged) {
447 const std::size_t blocks = workerCount == 0 ? std::size_t{1} : workerCount;
448 const bool directHamerlyScratch =
449 (d * k >= kHamerlyMinDirectScanDims) && (blocks <= kHamerlyDirectWorkerCap);
451 (k <= kHamerlyMaxK) && (k >= 2);
452 const std::size_t partialBlocks = hamerlyScratch ? hamerlyScatterBlocks(n, blocks) : blocks;
455 m_centroidsOld = NDArray<T, 2, Layout::Contig>({k, d});
456 m_cSqNorms = NDArray<T, 1>({k});
457 m_sums = NDArray<T, 2, Layout::Contig>({k, d});
458 m_counts = NDArray<std::int32_t, 1>({k});
459 m_minDistSq = NDArray<T, 1>({n});
460 m_xNormsSq = NDArray<T, 1>({n});
461 m_shiftSq = NDArray<T, 1>({k});
462 m_foldComp = NDArray<T, 1>({k * d});
466 m_u = NDArray<T, 1>({n});
467 m_l = NDArray<T, 1>({n});
468 m_shiftEuclidean = NDArray<T, 1>({k});
469 m_halfDistToNearestOther = NDArray<T, 1>({k});
474 const bool elkanCanFire = (k > kHamerlyMaxK) && (k <= kElkanMaxK) && (n * k <= kElkanNKLimit);
476 m_elkanBounds = NDArray<T, 2, Layout::Contig>({n, k});
477 m_centerDist = NDArray<T, 2, Layout::Contig>({k, k});
479 m_elkanBounds = NDArray<T, 2, Layout::Contig>({0, 0});
480 m_centerDist = NDArray<T, 2, Layout::Contig>({0, 0});
486 const std::size_t packedBSize = needsChunk
487 ? math::detail::packedBScratchSizeFloatsTiled<T>(k, d)
488 : math::detail::packedBScratchSizeFloats(k, d);
489 const std::size_t packedNormsSize = math::detail::packedCSqNormsScratchSizeFloats(k);
490 m_packedB = NDArray<T, 1>({packedBSize == 0 ? std::size_t{1} : packedBSize});
491 m_packedCSqNorms = NDArray<T, 1>({packedNormsSize == 0 ? std::size_t{1} : packedNormsSize});
494 const std::size_t distRows = needsChunk ? (blocks * chunkCap) : std::size_t{1};
495 const std::size_t safeK = (k == 0) ? std::size_t{1} : k;
496 const std::size_t distCols = needsChunk ? safeK : std::size_t{1};
497 m_distsChunk = NDArray<T, 2, Layout::Contig>({distRows, distCols});
498 }
else if (workerChanged) {
501 const std::size_t distRows = blocks * chunkCap;
502 const std::size_t distCols = (k == 0 ? std::size_t{1} : k);
503 m_distsChunk = NDArray<T, 2, Layout::Contig>({distRows, distCols});
509 m_partialSums = NDArray<T, 1>({partialBlocks * k * d});
510 m_partialComps = NDArray<T, 1>({partialBlocks * k * d});
511 m_partialCounts = NDArray<std::int32_t, 1>({partialBlocks * k});
515 const std::size_t apSize = blocks * math::detail::kMc<T> * math::detail::kKc<T>;
516 m_gemmApArena = NDArray<T, 1>({needsChunk ? apSize : std::size_t{1}});
518 if (shapeChanged || workerChanged) {
519 resetPackedXApCacheMetadata();
520 if (packedXApCacheEnabledForShape(n, d, workerCount)) {
521 const std::size_t numChunks = (n + chunkCap - 1) / chunkCap;
522 m_packedXApChunkRows = chunkCap;
523 m_packedXApPerChunk = math::detail::packedAScratchSizeForRows<T>(chunkCap, d);
524 m_packedXAp = NDArray<T, 1>({numChunks * m_packedXApPerChunk});
526 m_packedXApChunkRows = 0;
527 m_packedXApPerChunk = 0;
528 m_packedXAp = NDArray<T, 1>({std::size_t{1}});
535 m_workerCount = workerCount;
538 void refreshCentroidSqNorms(
const NDArray<T, 2, Layout::Contig> ¢roids)
noexcept {
539 const std::size_t k = centroids.dim(0);
540 const std::size_t d = centroids.dim(1);
541 for (std::size_t c = 0; c < k; ++c) {
542 const T *row = centroids.data() + (c * d);
544 for (std::size_t t = 0; t < d; ++t) {
545 s += row[t] * row[t];
551 void finalizeMeans(NDArray<T, 2, Layout::Contig> ¢roids)
noexcept {
552 const std::size_t k = centroids.dim(0);
553 const std::size_t d = centroids.dim(1);
554 for (std::size_t c = 0; c < k; ++c) {
555 const std::int32_t cnt = m_counts(c);
559 const T inv = T{1} /
static_cast<T
>(cnt);
560 const T *src = m_sums.data() + (c * d);
561 T *dst = centroids.data() + (c * d);
562 for (std::size_t t = 0; t < d; ++t) {
563 dst[t] = src[t] * inv;
575 void runAssignment(
const NDArray<T, 2, Layout::Contig> &X,
576 const NDArray<T, 2, Layout::Contig> ¢roids,
577 NDArray<std::int32_t, 1> &labels, math::Pool pool) {
578#ifdef CLUSTERING_USE_AVX2
579 if constexpr (std::is_same_v<T, float>) {
580 const std::size_t d = X.dim(1);
581 if (X.template isAligned<32>() && centroids.template isAligned<32>() && d != 0) {
583 math::detail::pairwiseArgminDirectSmallDF32(X, centroids, labels, m_minDistSq, pool);
587 math::detail::pairwiseArgminOuterAvx2F32WithScratch(X, centroids, m_cSqNorms, labels,
588 m_minDistSq, m_packedB.data(),
589 m_packedCSqNorms.data(), pool);
595 runChunkedMaterializedAssignment(X, centroids, labels, pool);
604 assignmentProducesDirectMinDistSq(
const NDArray<T, 2, Layout::Contig> &X,
605 const NDArray<T, 2, Layout::Contig> &C)
noexcept {
606#ifdef CLUSTERING_USE_AVX2
607 if constexpr (std::is_same_v<T, float>) {
608 const std::size_t d = X.dim(1);
609 return X.template isAligned<32>() && C.template isAligned<32>() && d != 0 &&
623 [[nodiscard]]
bool assignmentUsesFusedArgmin(
const NDArray<T, 2, Layout::Contig> &X,
624 const NDArray<T, 2, Layout::Contig> &C)
noexcept {
625#ifdef CLUSTERING_USE_AVX2
626 if constexpr (std::is_same_v<T, float>) {
627 const std::size_t d = X.dim(1);
628 return X.template isAligned<32>() && C.template isAligned<32>() &&
650 void packCentroidsTiled(
const NDArray<T, 2, Layout::Contig> ¢roids)
noexcept {
651 constexpr std::size_t kNr = math::detail::kKernelNr<T>;
652 constexpr std::size_t kKcVal = math::detail::kKc<T>;
653 constexpr std::size_t kNcVal = math::detail::kNc<T>;
654 const std::size_t k = centroids.dim(0);
655 const std::size_t d = centroids.dim(1);
656 const auto cTransposed = centroids.t();
657 const auto cDesc = ::clustering::detail::describeMatrix(cTransposed);
658 T *bp = m_packedB.data();
659 std::size_t jcBase = 0;
660 for (std::size_t jc = 0; jc < k; jc += kNcVal) {
661 const std::size_t nc = (jc + kNcVal <= k) ? kNcVal : (k - jc);
662 const std::size_t roundedNc = ((nc + kNr - 1) / kNr) * kNr;
663 std::size_t pcOffInJc = 0;
664 for (std::size_t pc = 0; pc < d; pc += kKcVal) {
665 const std::size_t kc = (pc + kKcVal <= d) ? kKcVal : (d - pc);
666 math::detail::packB<T>(cDesc, pc, kc, jc, nc, bp + jcBase + pcOffInJc);
667 pcOffInJc += kc * roundedNc;
669 jcBase += d * roundedNc;
673 [[nodiscard]]
bool ensurePackedXAp(
const NDArray<T, 2, Layout::Contig> &X)
noexcept {
674 if constexpr (std::is_same_v<T, float>) {
675 const std::size_t n = X.dim(0);
676 const std::size_t d = X.dim(1);
678 if (!packedXApCacheEnabledForShape(n, d, m_workerCount) || m_packedXApPerChunk == 0 ||
679 m_packedXApChunkRows != chunkCap) {
682 if (m_packedXApCachedXData == X.data() && m_packedXApCachedN == n &&
683 m_packedXApCachedD == d) {
687 const T *xBase = X.data();
688 const std::size_t numChunks = (n + chunkCap - 1) / chunkCap;
689 for (std::size_t c = 0; c < numChunks; ++c) {
690 const std::size_t iBase = c * chunkCap;
691 const std::size_t chunkRows = (iBase + chunkCap <= n) ? chunkCap : (n - iBase);
694 const auto xDesc = ::clustering::detail::describeMatrix(xChunk);
695 math::detail::packAChunk<T>(xDesc, chunkRows, d, chunkCap,
696 m_packedXAp.data() + (c * m_packedXApPerChunk));
699 m_packedXApCachedXData = X.data();
700 m_packedXApCachedN = n;
701 m_packedXApCachedD = d;
709 void runChunkedMaterializedAssignment(
const NDArray<T, 2, Layout::Contig> &X,
710 const NDArray<T, 2, Layout::Contig> ¢roids,
711 NDArray<std::int32_t, 1> &labels,
712 math::Pool pool)
noexcept {
713 const std::size_t n = X.dim(0);
714 const std::size_t k = centroids.dim(0);
715 const std::size_t d = X.dim(1);
716 if (n == 0 || k == 0) {
720 packCentroidsTiled(centroids);
722 constexpr std::size_t kMcVal = math::detail::kMc<T>;
723 constexpr std::size_t kKcVal = math::detail::kKc<T>;
725 const std::size_t numChunks = (n + chunkCap - 1) / chunkCap;
726 const T *bp = m_packedB.data();
727 T *apArena = m_gemmApArena.data();
728 T *distsBase = m_distsChunk.data();
729 const T *cNormsBase = m_cSqNorms.data();
730 T *minDistBase = m_minDistSq.data();
731 std::int32_t *labelsBase = labels.data();
732 const T *xBase = X.data();
735 T *uBase = m_u.data();
736 T *lBase = m_l.data();
740 T *elkanBoundsBase = m_elkanBounds.dim(0) == n ? m_elkanBounds.data() :
nullptr;
742 auto runOneChunk = [&](std::size_t chunkIdx)
noexcept {
743 const std::size_t iBase = chunkIdx * chunkCap;
744 const std::size_t chunkRows = (iBase + chunkCap <= n) ? chunkCap : (n - iBase);
746 T *distsChunk = distsBase + (w * chunkCap * k);
747 T *apSlice = apArena + (w * kMcVal * kKcVal);
752 const auto xDesc = ::clustering::detail::describeMatrix(xChunk);
753 auto distsDesc = ::clustering::detail::describeMatrixMut(distsView);
755 math::detail::gemmRunPrepacked<T>(xDesc, bp, d, k, distsDesc, T{-2}, T{0}, apSlice,
758 const T *xNormsChunk = m_xNormsSq.data() + iBase;
759 for (std::size_t i = 0; i < chunkRows; ++i) {
760 const T xn = xNormsChunk[i];
761 const T *row = distsChunk + (i * k);
762 T *elkanRow = elkanBoundsBase !=
nullptr ? elkanBoundsBase + ((iBase + i) * k) : nullptr;
763 T bestVal = std::numeric_limits<T>::infinity();
764 T secondVal = std::numeric_limits<T>::infinity();
765 std::int32_t bestIdx = 0;
766 for (std::size_t j = 0; j < k; ++j) {
767 T v = row[j] + xn + cNormsBase[j];
771 if (elkanRow !=
nullptr) {
772 elkanRow[j] = std::sqrt(v);
777 bestIdx =
static_cast<std::int32_t
>(j);
778 }
else if (v < secondVal) {
782 minDistBase[iBase + i] = bestVal;
783 labelsBase[iBase + i] = bestIdx;
784 uBase[iBase + i] = std::sqrt(bestVal);
785 lBase[iBase + i] = std::sqrt(secondVal);
789 pool.parallelForBlocks(std::size_t{0}, numChunks, std::size_t{0},
790 [&](std::size_t lo, std::size_t hi) {
791 for (std::size_t c = lo; c < hi; ++c) {
808 void runAssignmentAndScatter(
const NDArray<T, 2, Layout::Contig> &X,
809 const NDArray<T, 2, Layout::Contig> ¢roids,
810 NDArray<std::int32_t, 1> &labels, std::size_t k,
bool useKahan,
812 const std::size_t n = X.dim(0);
813 const std::size_t d = X.dim(1);
814 const std::size_t workers = pool.workerCount();
815 if (n == 0 || k == 0 || d == 0) {
819 preZeroPartialSlabs(useKahan, workers, k, d);
821#ifdef CLUSTERING_USE_AVX2
822 const bool aligned32 = X.template isAligned<32>() && centroids.template isAligned<32>();
824 const bool aligned32 =
false;
828 const bool useChunked = !useDirect && !useFused;
830#ifdef CLUSTERING_USE_AVX2
831 if constexpr (std::is_same_v<T, float>) {
833 math::detail::packCentroidsForFusedArgminF32(centroids, k, d, m_packedB.data());
834 math::detail::packCSqNorms<float>(m_cSqNorms.data(), k, m_packedCSqNorms.data());
838 bool usePackedXAp =
false;
840 packCentroidsTiled(centroids);
841 usePackedXAp = ensurePackedXAp(X);
844#ifdef CLUSTERING_USE_AVX2
845 if constexpr (std::is_same_v<T, float>) {
847 constexpr std::size_t kMr8 = 8;
848 constexpr std::size_t kMr16 = 16;
849 const bool useWideTile = workers > 1;
850 const std::size_t mTiles =
851 useWideTile ? ((n + kMr16 - 1) / kMr16) : ((n + kMr8 - 1) / kMr8);
852 pool.parallelForExactBlocksWithSlot<citor::HintsDefaults>(
853 std::size_t{0}, mTiles, workers,
854 [&](std::size_t lo, std::size_t hi, std::size_t slot)
noexcept {
856 assignScatterDirect16Tiles(X, centroids, labels, k, useKahan, lo, hi, slot);
858 assignScatterDirectTiles(X, centroids, labels, k, useKahan, lo, hi, slot);
861 foldPartialSlabs(useKahan, workers, k, d);
865 constexpr std::size_t kMr8 = math::detail::kKernelMr<float>;
866 constexpr std::size_t kMr16 = 16;
867 const bool useWideTile = workers > 1;
868 const std::size_t mTiles =
869 useWideTile ? ((n + kMr16 - 1) / kMr16) : ((n + kMr8 - 1) / kMr8);
870 pool.parallelForExactBlocksWithSlot<citor::HintsDefaults>(
871 std::size_t{0}, mTiles, workers,
872 [&](std::size_t lo, std::size_t hi, std::size_t slot)
noexcept {
874 assignScatterFused16Tiles(X, labels, k, useKahan, lo, hi, slot);
876 assignScatterFusedTiles(X, labels, k, useKahan, lo, hi, slot);
879 foldPartialSlabs(useKahan, workers, k, d);
887 const std::size_t numChunks = (n + chunkCap - 1) / chunkCap;
888 pool.parallelForExactBlocksWithSlot<citor::HintsDefaults>(
889 std::size_t{0}, numChunks, workers,
890 [&](std::size_t chunkLo, std::size_t chunkHi, std::size_t slot)
noexcept {
891 assignScatterChunkRange(X, labels, k, useKahan, chunkLo, chunkHi, slot, usePackedXAp);
894 foldPartialSlabs(useKahan, workers, k, d);
901 void assignScatterChunkRange(
const NDArray<T, 2, Layout::Contig> &X,
902 NDArray<std::int32_t, 1> &labels, std::size_t k,
bool useKahan,
903 std::size_t chunkLo, std::size_t chunkHi, std::size_t slot,
904 bool usePackedXAp)
noexcept {
905 constexpr std::size_t kMcVal = math::detail::kMc<T>;
906 constexpr std::size_t kKcVal = math::detail::kKc<T>;
907 const std::size_t n = X.dim(0);
908 const std::size_t d = X.dim(1);
910 const T *xBase = X.data();
911 std::int32_t *labelsBase = labels.data();
912 const T *bp = m_packedB.data();
913 const T *cNormsBase = m_cSqNorms.data();
914 T *minDistBase = m_minDistSq.data();
915 T *uBase = m_u.data();
916 T *lBase = m_l.data();
917 T *elkanBoundsBase = m_elkanBounds.dim(0) == n ? m_elkanBounds.data() :
nullptr;
918 T *distsChunk = m_distsChunk.data() + (slot * chunkCap * k);
919 T *apSlice = m_gemmApArena.data() + (slot * kMcVal * kKcVal);
920 const T *packedXBase = usePackedXAp ? m_packedXAp.data() :
nullptr;
921 const std::size_t packedXStride = m_packedXApPerChunk;
923 for (std::size_t c = chunkLo; c < chunkHi; ++c) {
924 const std::size_t iBase = c * chunkCap;
925 const std::size_t chunkRows = (iBase + chunkCap <= n) ? chunkCap : (n - iBase);
928 auto distsDesc = ::clustering::detail::describeMatrixMut(distsView);
930 if (packedXBase !=
nullptr) {
931 math::detail::gemmRunPrepackedAB<T>(packedXBase + (c * packedXStride), chunkRows, chunkCap,
932 bp, d, k, distsDesc, T{-2}, T{0});
936 const auto xDesc = ::clustering::detail::describeMatrix(xChunk);
937 math::detail::gemmRunPrepacked<T>(xDesc, bp, d, k, distsDesc, T{-2}, T{0}, apSlice,
941 const T *xNormsChunk = m_xNormsSq.data() + iBase;
942 for (std::size_t i = 0; i < chunkRows; ++i) {
943 const T xn = xNormsChunk[i];
944 const T *row = distsChunk + (i * k);
945 T *elkanRow = elkanBoundsBase !=
nullptr ? elkanBoundsBase + ((iBase + i) * k) : nullptr;
946 T bestVal = std::numeric_limits<T>::infinity();
947 T secondVal = std::numeric_limits<T>::infinity();
948 std::int32_t bestIdx = 0;
949 for (std::size_t j = 0; j < k; ++j) {
950 T v = row[j] + xn + cNormsBase[j];
954 if (elkanRow !=
nullptr) {
955 elkanRow[j] = std::sqrt(v);
960 bestIdx =
static_cast<std::int32_t
>(j);
961 }
else if (v < secondVal) {
965 minDistBase[iBase + i] = bestVal;
966 labelsBase[iBase + i] = bestIdx;
967 uBase[iBase + i] = std::sqrt(bestVal);
968 lBase[iBase + i] = std::sqrt(secondVal);
969 scatterRowToSlab(xBase, labelsBase, iBase + i, slot, k, d, useKahan);
974#ifdef CLUSTERING_USE_AVX2
976 void assignScatterDirectTiles(
const NDArray<T, 2, Layout::Contig> &X,
977 const NDArray<T, 2, Layout::Contig> ¢roids,
978 NDArray<std::int32_t, 1> &labels, std::size_t k,
bool useKahan,
979 std::size_t tileLo, std::size_t tileHi, std::size_t slot)
noexcept {
980 constexpr std::size_t kMr = 8;
981 const std::size_t n = X.dim(0);
982 const std::size_t d = X.dim(1);
983 const T *xBase = X.data();
984 const std::int32_t *labelsBase = labels.data();
985 for (std::size_t t = tileLo; t < tileHi; ++t) {
986 math::detail::argminDirectMTileF32(X, centroids, labels, m_minDistSq, t, n, k, d);
987 const std::size_t iBase = t * kMr;
988 const std::size_t mc = (iBase + kMr <= n) ? kMr : (n - iBase);
989 for (std::size_t r = 0; r < mc; ++r) {
990 scatterRowToSlab(xBase, labelsBase, iBase + r, slot, k, d, useKahan);
996 void assignScatterDirect16Tiles(
const NDArray<T, 2, Layout::Contig> &X,
997 const NDArray<T, 2, Layout::Contig> ¢roids,
998 NDArray<std::int32_t, 1> &labels, std::size_t k,
bool useKahan,
999 std::size_t tileLo, std::size_t tileHi,
1000 std::size_t slot)
noexcept {
1001 constexpr std::size_t kMr = 16;
1002 const std::size_t n = X.dim(0);
1003 const std::size_t d = X.dim(1);
1004 const T *xBase = X.data();
1005 const std::int32_t *labelsBase = labels.data();
1006 for (std::size_t t = tileLo; t < tileHi; ++t) {
1007 math::detail::argminDirectM16TileF32(X, centroids, labels, m_minDistSq, t, n, k, d);
1008 const std::size_t iBase = t * kMr;
1009 const std::size_t mc = (iBase + kMr <= n) ? kMr : (n - iBase);
1010 for (std::size_t r = 0; r < mc; ++r) {
1011 scatterRowToSlab(xBase, labelsBase, iBase + r, slot, k, d, useKahan);
1018 void assignScatterFusedTiles(
const NDArray<T, 2, Layout::Contig> &X,
1019 NDArray<std::int32_t, 1> &labels, std::size_t k,
bool useKahan,
1020 std::size_t tileLo, std::size_t tileHi, std::size_t slot)
noexcept {
1021 constexpr std::size_t kMr = math::detail::kKernelMr<float>;
1022 const std::size_t n = X.dim(0);
1023 const std::size_t d = X.dim(1);
1024 const T *xBase = X.data();
1025 const std::int32_t *labelsBase = labels.data();
1026 const float *bpacked = m_packedB.data();
1027 const float *normsPacked = m_packedCSqNorms.data();
1028 for (std::size_t t = tileLo; t < tileHi; ++t) {
1029 math::detail::argminFusedMTileF32(X, bpacked, normsPacked, labels, m_minDistSq, t, n, k, d);
1030 const std::size_t iBase = t * kMr;
1031 const std::size_t mc = (iBase + kMr <= n) ? kMr : (n - iBase);
1032 for (std::size_t r = 0; r < mc; ++r) {
1033 scatterRowToSlab(xBase, labelsBase, iBase + r, slot, k, d, useKahan);
1039 void assignScatterFused16Tiles(
const NDArray<T, 2, Layout::Contig> &X,
1040 NDArray<std::int32_t, 1> &labels, std::size_t k,
bool useKahan,
1041 std::size_t tileLo, std::size_t tileHi,
1042 std::size_t slot)
noexcept {
1043 constexpr std::size_t kMr = 16;
1044 const std::size_t n = X.dim(0);
1045 const std::size_t d = X.dim(1);
1046 const T *xBase = X.data();
1047 const std::int32_t *labelsBase = labels.data();
1048 const float *bpacked = m_packedB.data();
1049 const float *normsPacked = m_packedCSqNorms.data();
1050 for (std::size_t t = tileLo; t < tileHi; ++t) {
1051 math::detail::argminFusedM16TileF32(X, bpacked, normsPacked, labels, m_minDistSq, t, n, k, d);
1052 const std::size_t iBase = t * kMr;
1053 const std::size_t mc = (iBase + kMr <= n) ? kMr : (n - iBase);
1054 for (std::size_t r = 0; r < mc; ++r) {
1055 scatterRowToSlab(xBase, labelsBase, iBase + r, slot, k, d, useKahan);
1076 void runPlexLoop(
const NDArray<T, 2, Layout::Contig> &X, NDArray<T, 2, Layout::Contig> ¢roids,
1077 NDArray<std::int32_t, 1> &labels, std::size_t k, std::size_t maxIter,
1078 T shiftSqThreshold,
bool useKahan,
bool hamerlyEligible, math::Pool pool,
1079 std::size_t &iter,
bool &converged) {
1080 const std::size_t n = X.dim(0);
1081 const std::size_t d = X.dim(1);
1082 const std::size_t workers = pool.workerCount();
1088 const std::size_t units = (n + unit - 1) / unit;
1090 HamerlyShiftTop2 top2{};
1091 bool usePackedXAp =
false;
1092 bool stopPhases =
false;
1093 auto plexTok = citor::CancellationToken::makeOwned();
1095 auto packFusedCentroids = [&]()
noexcept {
1096 math::detail::packCentroidsForFusedArgminF32(centroids, k, d, m_packedB.data());
1097 math::detail::packCSqNorms<float>(m_cSqNorms.data(), k, m_packedCSqNorms.data());
1103 auto iterationGlue = [&](math::Pool gluePool) {
1104 foldPartialSlabs(useKahan, workers, k, d);
1105 (void)::clustering::kmeans::detail::reseedEmptyClusters<T>(X, centroids, m_sums, m_counts,
1107 finalizeMeans(centroids);
1108 refreshCentroidSqNorms(centroids);
1110 const T totalShift = ::clustering::kmeans::detail::totalShiftSqKahan<T>(m_shiftSq);
1112 if (totalShift <= shiftSqThreshold) {
1117 auto prePhase = [&](std::size_t phaseIdx) {
1118 if (phaseIdx == 0) {
1119 std::memcpy(m_centroidsOld.data(), centroids.data(), k * d *
sizeof(T));
1120 preZeroPartialSlabs(useKahan, workers, k, d);
1122 packCentroidsTiled(centroids);
1123 usePackedXAp = ensurePackedXAp(X);
1124 }
else if (!useDirect) {
1125 packFusedCentroids();
1129 iterationGlue(math::Pool{});
1132 plexTok.request_stop();
1135 std::memcpy(m_centroidsOld.data(), centroids.data(), k * d *
sizeof(T));
1136 preZeroPartialSlabs(useKahan, workers, k, d);
1137 if (hamerlyEligible) {
1138 top2 = prepareHamerlyGeometry(centroids, k, d);
1139 }
else if (useChunked) {
1140 packCentroidsTiled(centroids);
1141 usePackedXAp = ensurePackedXAp(X);
1142 }
else if (!useDirect) {
1143 packFusedCentroids();
1147 auto phase = [&](std::size_t phaseIdx, std::uint32_t slot, std::size_t lo, std::size_t hi,
1148 void * =
nullptr)
noexcept {
1152 const auto s =
static_cast<std::size_t
>(slot);
1153 if (hamerlyEligible && phaseIdx > 0) {
1154 hamerlyAssignScatterRange(X, centroids, labels, k, useKahan, top2, std::min(lo * unit, n),
1155 std::min(hi * unit, n), s);
1160 assignScatterChunkRange(X, labels, k, useKahan, lo, hi, s, usePackedXAp);
1164 assignScatterDirect16Tiles(X, centroids, labels, k, useKahan, lo, hi, s);
1166 assignScatterFused16Tiles(X, labels, k, useKahan, lo, hi, s);
1168 if (hamerlyEligible && phaseIdx == 0) {
1169 seedHamerlyBoundsFromMinDistRange(labels, k, d, std::min(lo * unit, n),
1170 std::min(hi * unit, n));
1174 pool.parallelRunPlex<citor::HintsDefaults>(maxIter, units, std::move(phase),
1175 std::move(prePhase), plexTok);
1180 iterationGlue(pool);
1185 void preZeroPartialSlabs(
bool useKahan, std::size_t numBlocks, std::size_t k,
1186 std::size_t d)
noexcept {
1187 T *partialSums = m_partialSums.data();
1188 std::int32_t *partialCounts = m_partialCounts.data();
1190 for (std::size_t c = 0; c < k; ++c) {
1192 for (std::size_t t = 0; t < d; ++t) {
1193 m_sums(c, t) = T{0};
1197 T *foldComp = m_foldComp.data();
1198 for (std::size_t e = 0; e < k * d; ++e) {
1201 T *partialComps = m_partialComps.data();
1202 for (std::size_t b = 0; b < numBlocks; ++b) {
1203 T *slab = partialSums + (b * k * d);
1204 T *cslab = partialComps + (b * k * d);
1205 std::int32_t *nslab = partialCounts + (b * k);
1206 for (std::size_t e = 0; e < k * d; ++e) {
1210 for (std::size_t c = 0; c < k; ++c) {
1215 for (std::size_t b = 0; b < numBlocks; ++b) {
1216 T *slab = partialSums + (b * k * d);
1217 std::int32_t *cslab = partialCounts + (b * k);
1218 for (std::size_t e = 0; e < k * d; ++e) {
1221 for (std::size_t c = 0; c < k; ++c) {
1228 void foldPartialSlabs(
bool useKahan, std::size_t numBlocks, std::size_t k,
1229 std::size_t d)
noexcept {
1230 const T *partialSums = m_partialSums.data();
1231 const std::int32_t *partialCounts = m_partialCounts.data();
1233 const T *partialComps = m_partialComps.data();
1234 T *foldComp = m_foldComp.data();
1235 for (std::size_t b = 0; b < numBlocks; ++b) {
1236 const T *slab = partialSums + (b * k * d);
1237 const T *cslab = partialComps + (b * k * d);
1238 const std::int32_t *nslab = partialCounts + (b * k);
1239 for (std::size_t c = 0; c < k; ++c) {
1240 m_counts(c) += nslab[c];
1241 const T *src = slab + (c * d);
1242 const T *comp = cslab + (c * d);
1243 T *dstRow = &m_sums(c, 0);
1244 T *foldRow = foldComp + (c * d);
1245 for (std::size_t t = 0; t < d; ++t) {
1246 const T addend = src[t] - comp[t];
1247 const T y = addend - foldRow[t];
1248 const T tVal = dstRow[t] + y;
1249 foldRow[t] = (tVal - dstRow[t]) - y;
1255 for (std::size_t b = 0; b < numBlocks; ++b) {
1256 const T *slab = partialSums + (b * k * d);
1257 const std::int32_t *cslab = partialCounts + (b * k);
1258 for (std::size_t c = 0; c < k; ++c) {
1259 m_counts(c) += cslab[c];
1260 const T *src = slab + (c * d);
1261 T *dstRow = &m_sums(c, 0);
1262 for (std::size_t t = 0; t < d; ++t) {
1263 dstRow[t] += src[t];
1270 [[nodiscard]]
static std::size_t hamerlyScatterBlocks(std::size_t n,
1271 std::size_t workers)
noexcept {
1272 if (workers <= 1 || n == 0) {
1273 return std::max<std::size_t>(workers, std::size_t{1});
1275 constexpr std::size_t kMinRowsPerBlock = 256;
1276 const std::size_t byRows = std::max<std::size_t>(1, n / kMinRowsPerBlock);
1277 const std::size_t blocks = std::min(workers * 8, byRows);
1278 return std::max(blocks, workers);
1281 void recomputeMinDistSqDirect(
const NDArray<T, 2, Layout::Contig> &X,
1282 const NDArray<T, 2, Layout::Contig> ¢roids,
1283 const NDArray<std::int32_t, 1> &labels, math::Pool pool)
noexcept {
1284 const std::size_t n = X.dim(0);
1285 const std::size_t d = X.dim(1);
1286 const std::size_t k = centroids.dim(0);
1287 if (n == 0 || d == 0 || k == 0) {
1291 auto runRowRange = [&](std::size_t lo, std::size_t hi)
noexcept {
1292 for (std::size_t i = lo; i < hi; ++i) {
1293 const std::int32_t lbl = labels(i);
1294 if (lbl < 0 || std::cmp_greater_equal(lbl, k)) {
1295 m_minDistSq(i) = T{0};
1298 const T *xRow = X.data() + (i * d);
1299 const T *cRow = centroids.data() + (
static_cast<std::size_t
>(lbl) * d);
1300 m_minDistSq(i) = math::detail::sqEuclideanRowPtr<T>(xRow, cRow, d);
1304 pool.parallelForBlocks(std::size_t{0}, n, std::size_t{0},
1305 [&](std::size_t lo, std::size_t hi) { runRowRange(lo, hi); });
1309 void seedHamerlyBoundsFromMinDistRange(
const NDArray<std::int32_t, 1> &labels, std::size_t k,
1310 std::size_t d, std::size_t lo, std::size_t hi)
noexcept {
1313 const T slackScale =
static_cast<T
>(8) * std::numeric_limits<T>::epsilon() *
static_cast<T
>(d);
1314 for (std::size_t i = lo; i < hi; ++i) {
1315 const std::int32_t lbl = labels(i);
1316 if (lbl < 0 || std::cmp_greater_equal(lbl, k)) {
1317 m_minDistSq(i) = T{0};
1318 m_u(i) = std::numeric_limits<T>::infinity();
1322 T tightSq = m_minDistSq(i);
1323 if (tightSq < T{0}) {
1325 m_minDistSq(i) = T{0};
1327 m_minDistSq(i) = tightSq;
1328 const T u = std::sqrt(tightSq);
1329 m_u(i) = u + ((u + T{1}) * slackScale);
1334 void seedHamerlyBoundsFromAssignedMinDist(
const NDArray<std::int32_t, 1> &labels, std::size_t k,
1335 std::size_t d, math::Pool pool)
noexcept {
1336 const std::size_t n = labels.dim(0);
1337 if (n == 0 || d == 0 || k == 0) {
1341 pool.parallelForBlocks(std::size_t{0}, n, std::size_t{0}, [&](std::size_t lo, std::size_t hi) {
1342 seedHamerlyBoundsFromMinDistRange(labels, k, d, lo, hi);
1352 static constexpr std::size_t kHamerlyMaxK = 64;
1362 static constexpr std::size_t kHamerlyMinDirectScanDims = 128;
1372 static constexpr std::size_t kHamerlyDirectWorkerCap = 4;
1380 static constexpr std::size_t kElkanMaxK = 4096;
1389 static constexpr std::size_t kElkanNKLimit = std::size_t{32} << 20;
1401 [[nodiscard]] HamerlyShiftTop2
1402 prepareHamerlyGeometry(
const NDArray<T, 2, Layout::Contig> ¢roids, std::size_t k,
1403 std::size_t d)
noexcept {
1404 HamerlyShiftTop2 top2{};
1405 const T *cData = centroids.data();
1406 T *shiftData = m_shiftEuclidean.data();
1407 for (std::size_t c = 0; c < k; ++c) {
1408 const T s = std::sqrt(m_shiftSq(c));
1410 if (s > top2.sMax) {
1411 top2.s2Max = top2.sMax;
1414 }
else if (s > top2.s2Max) {
1419 T *halfDistData = m_halfDistToNearestOther.data();
1420 for (std::size_t c = 0; c < k; ++c) {
1421 T nearestSq = std::numeric_limits<T>::infinity();
1422 const T *caRow = cData + (c * d);
1423 for (std::size_t cp = 0; cp < k; ++cp) {
1427 const T dsq = math::detail::sqEuclideanRowPtr<T>(caRow, cData + (cp * d), d);
1428 if (dsq < nearestSq) {
1432 halfDistData[c] = T{0.5} * std::sqrt(nearestSq);
1440 void hamerlyAssignScatterRange(
const NDArray<T, 2, Layout::Contig> &X,
1441 const NDArray<T, 2, Layout::Contig> ¢roids,
1442 NDArray<std::int32_t, 1> &labels, std::size_t k,
bool useKahan,
1443 const HamerlyShiftTop2 &top2, std::size_t lo, std::size_t hi,
1444 std::size_t slot)
noexcept {
1445 const std::size_t d = X.dim(1);
1446 const T *xData = X.data();
1447 const T *cData = centroids.data();
1448 T *uData = m_u.data();
1449 T *lData = m_l.data();
1450 T *minDistData = m_minDistSq.data();
1451 std::int32_t *labelsData = labels.data();
1452 const T *shiftData = m_shiftEuclidean.data();
1453 const T *halfDistData = m_halfDistToNearestOther.data();
1454 T *slabSum = m_partialSums.data() + (slot * k * d);
1455 T *slabComp = m_partialComps.data() + (slot * k * d);
1456 std::int32_t *slabCnt = m_partialCounts.data() + (slot * k);
1458 std::array<T, kHamerlyMaxK> distBuf{};
1459 for (std::size_t i = lo; i < hi; ++i) {
1460 const std::int32_t a = labelsData[i];
1461 if (a < 0 || std::cmp_greater_equal(a, k)) {
1464 const auto au =
static_cast<std::size_t
>(a);
1465 T ui = uData[i] + shiftData[au];
1466 T li = lData[i] - ((au == top2.argMax) ? top2.s2Max : top2.sMax);
1467 std::int32_t bestLabel = a;
1468 bool labelDecided =
false;
1470 if (ui <= li || ui <= halfDistData[au]) {
1473 labelDecided =
true;
1476 if (!labelDecided) {
1477 const T *xi = xData + (i * d);
1478 const T *caRow = cData + (au * d);
1479 const T tightSq = math::detail::sqEuclideanRowPtr<T>(xi, caRow, d);
1480 ui = std::sqrt(tightSq);
1485 minDistData[i] = tightSq;
1486 labelDecided =
true;
1489 T best = std::numeric_limits<T>::infinity();
1490 T second = std::numeric_limits<T>::infinity();
1491 std::int32_t bestIdx = 0;
1492 for (std::size_t j = 0; j < k; ++j) {
1493 const T v = distBuf[j];
1497 bestIdx =
static_cast<std::int32_t
>(j);
1498 }
else if (v < second) {
1502 bestLabel = bestIdx;
1503 labelsData[i] = bestIdx;
1504 minDistData[i] = best;
1505 uData[i] = std::sqrt(best);
1506 lData[i] = std::sqrt(second);
1510 if (bestLabel < 0 || std::cmp_greater_equal(bestLabel, k)) {
1513 const auto row =
static_cast<std::size_t
>(bestLabel);
1514 const T *xRow = xData + (i * d);
1515 T *sumRow = slabSum + (row * d);
1517 T *compRow = slabComp + (row * d);
1518 math::detail::kahanAddRow<T>(xRow, d, sumRow, compRow);
1520 for (std::size_t t = 0; t < d; ++t) {
1521 sumRow[t] += xRow[t];
1537 void runHamerlyAssignmentAndScatter(
const NDArray<T, 2, Layout::Contig> &X,
1538 const NDArray<T, 2, Layout::Contig> ¢roids,
1539 NDArray<std::int32_t, 1> &labels, std::size_t k,
1540 bool useKahan, math::Pool pool)
noexcept {
1541 const std::size_t n = X.dim(0);
1542 const std::size_t d = X.dim(1);
1543 if (n == 0 || d == 0 || k == 0 || k > kHamerlyMaxK) {
1546 const std::size_t workers = pool.workerCount();
1547 const std::size_t blocks = hamerlyScatterBlocks(n, workers);
1548 preZeroPartialSlabs(useKahan, blocks, k, d);
1550 const HamerlyShiftTop2 top2 = prepareHamerlyGeometry(centroids, k, d);
1552 pool.parallelForExactBlocksWithSlot<citor::HintsDefaults>(
1553 std::size_t{0}, n, blocks, [&](std::size_t lo, std::size_t hi, std::size_t slot)
noexcept {
1554 hamerlyAssignScatterRange(X, centroids, labels, k, useKahan, top2, lo, hi, slot);
1557 foldPartialSlabs(useKahan, blocks, k, d);
1566 void runElkanAssignmentAndScatter(
const NDArray<T, 2, Layout::Contig> &X,
1567 const NDArray<T, 2, Layout::Contig> ¢roids,
1568 NDArray<std::int32_t, 1> &labels, std::size_t k,
bool useKahan,
1569 math::Pool pool)
noexcept {
1570 const std::size_t n = X.dim(0);
1571 const std::size_t d = X.dim(1);
1572 if (n == 0 || d == 0 || k == 0 || m_elkanBounds.dim(0) != n || m_elkanBounds.dim(1) != k) {
1575 const std::size_t workers = pool.workerCount();
1576 preZeroPartialSlabs(useKahan, workers, k, d);
1578 const T *xData = X.data();
1579 const T *cData = centroids.data();
1580 T *uData = m_u.data();
1581 T *boundsData = m_elkanBounds.data();
1582 T *minDistData = m_minDistSq.data();
1583 std::int32_t *labelsData = labels.data();
1585 T *shiftData = m_shiftEuclidean.data();
1586 for (std::size_t c = 0; c < k; ++c) {
1587 shiftData[c] = std::sqrt(m_shiftSq(c));
1590 T *centerDistData = m_centerDist.data();
1591 T *halfDistData = m_halfDistToNearestOther.data();
1592 for (std::size_t c = 0; c < k; ++c) {
1593 centerDistData[(c * k) + c] = T{0};
1594 T nearest = std::numeric_limits<T>::infinity();
1595 for (std::size_t cp = 0; cp < k; ++cp) {
1601 const T dsq = math::detail::sqEuclideanRowPtr<T>(cData + (c * d), cData + (cp * d), d);
1602 dist = std::sqrt(dsq);
1603 centerDistData[(c * k) + cp] = dist;
1604 centerDistData[(cp * k) + c] = dist;
1606 dist = centerDistData[(c * k) + cp];
1608 if (dist < nearest) {
1612 halfDistData[c] = T{0.5} * nearest;
1615 T *partialSums = m_partialSums.data();
1616 T *partialComps = m_partialComps.data();
1617 std::int32_t *partialCounts = m_partialCounts.data();
1619 pool.parallelForBlocks<citor::HintsDefaults>(
1620 std::size_t{0}, n, std::size_t{0}, [&](std::size_t lo, std::size_t hi)
noexcept {
1622 T *slabSum = partialSums + (slot * k * d);
1623 T *slabComp = partialComps + (slot * k * d);
1624 std::int32_t *slabCnt = partialCounts + (slot * k);
1625 for (std::size_t i = lo; i < hi; ++i) {
1626 std::int32_t a = labelsData[i];
1627 if (a < 0 || std::cmp_greater_equal(a, k)) {
1630 auto au =
static_cast<std::size_t
>(a);
1631 T u = uData[i] + shiftData[au];
1632 T *lRow = boundsData + (i * k);
1633 for (std::size_t c = 0; c < k; ++c) {
1634 T lnew = lRow[c] - shiftData[c];
1641 if (u <= halfDistData[au]) {
1644 bool uTight =
false;
1645 const T *xi = xData + (i * d);
1646 for (std::size_t c = 0; c < k; ++c) {
1650 const T lc = lRow[c];
1651 const T half = T{0.5} * centerDistData[(au * k) + c];
1652 if (u <= lc || u <= half) {
1656 const T tightSq = math::detail::sqEuclideanRowPtr<T>(xi, cData + (au * d), d);
1657 u = std::sqrt(tightSq);
1658 minDistData[i] = tightSq;
1660 if (u <= lc || u <= half) {
1664 const T dSq = math::detail::sqEuclideanRowPtr<T>(xi, cData + (c * d), d);
1665 const T dEuc = std::sqrt(dSq);
1669 a =
static_cast<std::int32_t
>(c);
1671 minDistData[i] = dSq;
1678 const auto row =
static_cast<std::size_t
>(a);
1679 const T *xRow = xData + (i * d);
1680 T *sumRow = slabSum + (row * d);
1682 T *compRow = slabComp + (row * d);
1683 math::detail::kahanAddRow<T>(xRow, d, sumRow, compRow);
1685 for (std::size_t t = 0; t < d; ++t) {
1686 sumRow[t] += xRow[t];
1693 foldPartialSlabs(useKahan, workers, k, d);
1706 void runHamerlyAssignment(
const NDArray<T, 2, Layout::Contig> &X,
1707 const NDArray<T, 2, Layout::Contig> ¢roids,
1708 NDArray<std::int32_t, 1> &labels, math::Pool pool,
1709 bool refreshAssignedMinDist =
false) noexcept {
1710 const std::size_t n = X.dim(0);
1711 const std::size_t d = X.dim(1);
1712 const std::size_t k = centroids.dim(0);
1713 if (n == 0 || d == 0 || k == 0 || k > kHamerlyMaxK) {
1716 const T *xData = X.data();
1717 const T *cData = centroids.data();
1718 T *uData = m_u.data();
1719 T *lData = m_l.data();
1720 T *minDistData = m_minDistSq.data();
1721 std::int32_t *labelsData = labels.data();
1726 const HamerlyShiftTop2 top2 = prepareHamerlyGeometry(centroids, k, d);
1727 const T *shiftData = m_shiftEuclidean.data();
1728 const T *halfDistData = m_halfDistToNearestOther.data();
1730 auto processRange = [&](std::size_t lo, std::size_t hi)
noexcept {
1731 std::array<T, kHamerlyMaxK> distBuf{};
1732 for (std::size_t i = lo; i < hi; ++i) {
1733 const std::int32_t a = labelsData[i];
1734 if (a < 0 || std::cmp_greater_equal(a, k)) {
1737 const auto au =
static_cast<std::size_t
>(a);
1738 T ui = uData[i] + shiftData[au];
1739 T li = lData[i] - ((au == top2.argMax) ? top2.s2Max : top2.sMax);
1741 const bool assignedByLower = ui <= li;
1742 const bool assignedByCenter = !assignedByLower && ui <= halfDistData[au];
1743 if (assignedByLower || assignedByCenter) {
1744 if (refreshAssignedMinDist) {
1745 const T *xi = xData + (i * d);
1746 const T *caRow = cData + (au * d);
1747 const T tightSq = math::detail::sqEuclideanRowPtr<T>(xi, caRow, d);
1748 minDistData[i] = tightSq;
1749 ui = std::sqrt(tightSq);
1756 const T *xi = xData + (i * d);
1757 const T *caRow = cData + (au * d);
1758 const T tightSq = math::detail::sqEuclideanRowPtr<T>(xi, caRow, d);
1759 ui = std::sqrt(tightSq);
1764 minDistData[i] = tightSq;
1769 T best = std::numeric_limits<T>::infinity();
1770 T second = std::numeric_limits<T>::infinity();
1771 std::int32_t bestIdx = 0;
1772 for (std::size_t j = 0; j < k; ++j) {
1773 const T v = distBuf[j];
1777 bestIdx =
static_cast<std::int32_t
>(j);
1778 }
else if (v < second) {
1782 labelsData[i] = bestIdx;
1783 minDistData[i] = best;
1784 uData[i] = std::sqrt(best);
1785 lData[i] = std::sqrt(second);
1789 pool.parallelForBlocks(std::size_t{0}, n, std::size_t{0},
1790 [&](std::size_t lo, std::size_t hi) { processRange(lo, hi); });
1803 void runElkanAssignment(
const NDArray<T, 2, Layout::Contig> &X,
1804 const NDArray<T, 2, Layout::Contig> ¢roids,
1805 NDArray<std::int32_t, 1> &labels, math::Pool pool)
noexcept {
1806 const std::size_t n = X.dim(0);
1807 const std::size_t d = X.dim(1);
1808 const std::size_t k = centroids.dim(0);
1809 if (n == 0 || d == 0 || k == 0 || m_elkanBounds.dim(0) != n || m_elkanBounds.dim(1) != k) {
1812 const T *xData = X.data();
1813 const T *cData = centroids.data();
1814 T *uData = m_u.data();
1815 T *boundsData = m_elkanBounds.data();
1816 T *minDistData = m_minDistSq.data();
1817 std::int32_t *labelsData = labels.data();
1820 T *shiftData = m_shiftEuclidean.data();
1821 for (std::size_t c = 0; c < k; ++c) {
1822 shiftData[c] = std::sqrt(m_shiftSq(c));
1827 T *centerDistData = m_centerDist.data();
1828 T *halfDistData = m_halfDistToNearestOther.data();
1829 for (std::size_t c = 0; c < k; ++c) {
1830 centerDistData[(c * k) + c] = T{0};
1831 T nearest = std::numeric_limits<T>::infinity();
1832 for (std::size_t cp = 0; cp < k; ++cp) {
1838 const T dsq = math::detail::sqEuclideanRowPtr<T>(cData + (c * d), cData + (cp * d), d);
1839 dist = std::sqrt(dsq);
1840 centerDistData[(c * k) + cp] = dist;
1841 centerDistData[(cp * k) + c] = dist;
1843 dist = centerDistData[(c * k) + cp];
1845 if (dist < nearest) {
1849 halfDistData[c] = T{0.5} * nearest;
1852 auto processRange = [&](std::size_t lo, std::size_t hi)
noexcept {
1853 for (std::size_t i = lo; i < hi; ++i) {
1854 std::int32_t a = labelsData[i];
1855 if (a < 0 || std::cmp_greater_equal(a, k)) {
1858 auto au =
static_cast<std::size_t
>(a);
1859 T u = uData[i] + shiftData[au];
1860 T *lRow = boundsData + (i * k);
1863 for (std::size_t c = 0; c < k; ++c) {
1864 T lnew = lRow[c] - shiftData[c];
1871 if (u <= halfDistData[au]) {
1876 bool uTight =
false;
1877 const T *xi = xData + (i * d);
1878 for (std::size_t c = 0; c < k; ++c) {
1882 const T lc = lRow[c];
1883 const T half = T{0.5} * centerDistData[(au * k) + c];
1884 if (u <= lc || u <= half) {
1888 const T tightSq = math::detail::sqEuclideanRowPtr<T>(xi, cData + (au * d), d);
1889 u = std::sqrt(tightSq);
1890 minDistData[i] = tightSq;
1892 if (u <= lc || u <= half) {
1896 const T dSq = math::detail::sqEuclideanRowPtr<T>(xi, cData + (c * d), d);
1897 const T dEuc = std::sqrt(dSq);
1901 a =
static_cast<std::int32_t
>(c);
1903 minDistData[i] = dSq;
1911 if (pool.shouldParallelize(n, 64, 2)) {
1912 pool.parallelForBlocks(std::size_t{0}, n, std::size_t{0},
1913 [&](std::size_t lo, std::size_t hi) { processRange(lo, hi); });
1919 NDArray<T, 2, Layout::Contig> m_centroidsOld;
1920 NDArray<T, 1> m_cSqNorms;
1921 NDArray<T, 2, Layout::Contig> m_sums;
1922 NDArray<std::int32_t, 1> m_counts;
1923 NDArray<T, 1> m_minDistSq;
1924 NDArray<T, 1> m_shiftSq;
1925 NDArray<T, 1> m_partialSums;
1926 NDArray<T, 1> m_partialComps;
1927 NDArray<std::int32_t, 1> m_partialCounts;
1928 NDArray<T, 1> m_foldComp;
1929 NDArray<T, 1> m_packedB;
1930 NDArray<T, 1> m_packedCSqNorms;
1931 NDArray<T, 2, Layout::Contig> m_distsChunk;
1932 NDArray<T, 1> m_gemmApArena;
1933 NDArray<T, 1> m_packedXAp;
1934 NDArray<T, 1> m_xNormsSq;
1935 NDArray<T, 1> m_varSum;
1936 NDArray<T, 1> m_varSumSq;
1937 NDArray<T, 1> m_varPartialSum;
1938 NDArray<T, 1> m_varPartialSumSq;
1945 NDArray<T, 1> m_shiftEuclidean;
1949 NDArray<T, 1> m_halfDistToNearestOther;
1953 NDArray<T, 2, Layout::Contig> m_elkanBounds;
1956 NDArray<T, 2, Layout::Contig> m_centerDist;
1958 std::size_t m_n = 0;
1959 std::size_t m_d = 0;
1960 std::size_t m_k = 0;
1961 std::size_t m_workerCount = 0;
1963 const T *m_packedXApCachedXData =
nullptr;
1964 std::size_t m_packedXApCachedN = 0;
1965 std::size_t m_packedXApCachedD = 0;
1966 std::size_t m_packedXApChunkRows = 0;
1967 std::size_t m_packedXApPerChunk = 0;
1972 const T *m_xstatsCachedXData =
nullptr;
1973 std::size_t m_xstatsCachedN = 0;
1974 std::size_t m_xstatsCachedD = 0;
1975 T m_xstatsCachedMeanVar{0};