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");
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}) {}
403 const std::size_t n = X.
dim(0);
404 const std::size_t d = X.
dim(1);
418 constexpr std::size_t kScoreSlabFloats = 16;
419 const std::size_t scoreSlab =
420 ((nLocalTrials + kScoreSlabFloats - 1) / kScoreSlabFloats) * kScoreSlabFloats;
421 ensureShape(n, d, nLocalTrials, pool.
workerCount());
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();
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) {
447 std::memcpy(centroidsData, xData + (first * d), d *
sizeof(T));
454#ifdef CLUSTERING_USE_AVX2
455 if constexpr (std::is_same_v<T, float>) {
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);
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);
484 std::vector<std::size_t> candidates(nLocalTrials, 0);
485 std::vector<T> scores(nLocalTrials, T{0});
487 for (std::size_t c = 1; c < k; ++c) {
491 for (std::size_t s = 0; s < sweepBlocks; ++s) {
492 total += sweepSums[s];
497 if (!(total > T{0})) {
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);
512 drawCandidates(rng, total, n, sweepBlocks, nLocalTrials, candidates.data());
517 packCandidates(xData, d, nLocalTrials, candidates.data());
519 for (std::size_t t = 0; t < nLocalTrials; ++t) {
522 constexpr std::size_t kMaxLocalTrials = 32;
526 bool scoredViaTransposed =
false;
527#ifdef CLUSTERING_USE_AVX2
533 if constexpr (std::is_same_v<T, float>) {
535 packCandidatesTransposed(d, nLocalTrials, 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);
546 if (transposedWidth == 16) {
547 if (willParallelizeT) {
549 [&](std::size_t lo, std::size_t hi) {
551 scoreTransposed16Range(xData, d, nLocalTrials, lo, hi,
552 localScoresT + (w * scoreSlab));
555 scoreTransposed16Range(xData, d, nLocalTrials, std::size_t{0}, n, localScoresT);
557 }
else if (transposedWidth == 8) {
558 if (willParallelizeT) {
560 [&](std::size_t lo, std::size_t hi) {
562 scoreTransposed8Range(xData, d, nLocalTrials, lo, hi,
563 localScoresT + (w * scoreSlab));
566 scoreTransposed8Range(xData, d, nLocalTrials, std::size_t{0}, n, localScoresT);
571 (void)candDistRows(n, transposedWidth);
572 if (willParallelizeT) {
574 std::size_t{0}, n, blocksT, [&](std::size_t lo, std::size_t hi) {
576 scoreTransposedChunkedRange(xData, d, nLocalTrials, transposedWidth, lo, hi,
577 localScoresT + (w * scoreSlab));
580 scoreTransposedChunkedRange(xData, d, nLocalTrials, transposedWidth, std::size_t{0},
584 foldScoreSlabs(workersT, scoreSlab, nLocalTrials, scores.data());
585 scoredViaTransposed =
true;
590 if (!scoredViaTransposed) {
595 if (useGemmScoring) {
599 auto candT = candView.t();
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);
608 T *candNorms = m_candNormsSq.data();
609 for (std::size_t t = 0; t < nLocalTrials; ++t) {
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];
626 scores[t] += (v < mi) ? v : mi;
637 bool scoredViaSoa =
false;
638#ifdef CLUSTERING_USE_AVX2
639 if constexpr (std::is_same_v<T, float>) {
649 const bool soaEligible = (d >= 8) && (nLocalTrials >= 1) && (nLocalTrials <= 6);
651 if (willParallelize) {
653 T *localScores = m_localScores.data();
654 zeroScoreSlabs(workers, scoreSlab);
656 [&](std::size_t lo, std::size_t hi) {
658 scoreSoaRange(xData, d, nLocalTrials, transposedWidth, lo,
659 hi, localScores + (w * scoreSlab));
661 foldScoreSlabs(workers, scoreSlab, nLocalTrials, scores.data());
663 scoreSoaRange(xData, d, nLocalTrials, transposedWidth, std::size_t{0}, n,
671 (void)candDistRows(n, transposedWidth);
672 if (willParallelize) {
674 T *localScores = m_localScores.data();
675 zeroScoreSlabs(workers, scoreSlab);
677 [&](std::size_t lo, std::size_t hi) {
679 scoreScalarRange(xData, d, nLocalTrials, transposedWidth, lo,
680 hi, localScores + (w * scoreSlab));
682 foldScoreSlabs(workers, scoreSlab, nLocalTrials, scores.data());
684 scoreScalarRange(xData, d, nLocalTrials, transposedWidth, std::size_t{0}, n,
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];
699 const std::size_t bestCandidate = candidates[bestT];
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);
725 T *candDistRows(std::size_t n, std::size_t w) {
726 if (m_candDistSq.
dim(0) != n || m_candDistSq.
dim(1) != w) {
729 return m_candDistSq.data();
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();
740 m_sweepSums.data()[s] =
741 math::detail::refreshMinSqAgainstRow(firstRow, xData + (lo * d), hi - lo, d, minSq + lo);
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);
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) {
762 while (s + 1 < sweepBlocks && acc + sweepSums[s] <= u) {
766 candidates[t] = math::detail::inverseCdfPickInRange(minSq, (n * s) / sweepBlocks,
767 (n * (s + 1)) / sweepBlocks, u - acc);
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));
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];
791 for (std::size_t t = nLocalTrials; t < transposedWidth; ++t) {
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};
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) {
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,
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];
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;
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]);
852 scoresLoAcc = _mm256_add_ps(scoresLoAcc, _mm256_min_ps(dLo, miVec));
853 scoresHiAcc = _mm256_add_ps(scoresHiAcc, _mm256_min_ps(dHi, miVec));
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) {
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]);
873 scoresAcc = _mm256_add_ps(scoresAcc, _mm256_min_ps(dv, miVec));
875 std::array<float, 8> tmp{};
876 _mm256_storeu_ps(tmp.data(), scoresAcc);
877 for (std::size_t t = 0; t < nLocalTrials; ++t) {
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,
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) {
896 transposedWidth, dstRow + base);
898 for (std::size_t t = 0; t < nLocalTrials; ++t) {
899 dst[t] += (dstRow[t] < mi) ? dstRow[t] : mi;
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) {
914 math::detail::kmppScoreSoaRowsAvx2F32<1,
false>(
915 xSlice, rangeN, d, candRowsData, minSlice,
nullptr, transposedWidth, dst);
918 math::detail::kmppScoreSoaRowsAvx2F32<2,
false>(
919 xSlice, rangeN, d, candRowsData, minSlice,
nullptr, transposedWidth, dst);
922 math::detail::kmppScoreSoaRowsAvx2F32<3,
false>(
923 xSlice, rangeN, d, candRowsData, minSlice,
nullptr, transposedWidth, dst);
926 math::detail::kmppScoreSoaRowsAvx2F32<4,
false>(
927 xSlice, rangeN, d, candRowsData, minSlice,
nullptr, transposedWidth, dst);
930 math::detail::kmppScoreSoaRowsAvx2F32<5,
false>(
931 xSlice, rangeN, d, candRowsData, minSlice,
nullptr, transposedWidth, dst);
934 math::detail::kmppScoreSoaRowsAvx2F32<6,
false>(
935 xSlice, rangeN, d, candRowsData, minSlice,
nullptr, transposedWidth, dst);
944 enum class ScoreKernel : std::uint8_t {
954 [[nodiscard]]
static ScoreKernel pickScoreKernel(std::size_t d,
955 std::size_t nLocalTrials)
noexcept {
959 return ScoreKernel::kTransposed8;
962 return ScoreKernel::kTransposed16;
964 return ScoreKernel::kTransposedChunked;
966 if (d >= 8 && nLocalTrials >= 1 && nLocalTrials <= 6) {
967 return ScoreKernel::kSoa;
969 return ScoreKernel::kScalar;
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 {
976 case ScoreKernel::kTransposed16:
977 scoreTransposed16Range(xData, d, nLocalTrials, lo, hi, dst);
979 case ScoreKernel::kTransposed8:
980 scoreTransposed8Range(xData, d, nLocalTrials, lo, hi, dst);
982 case ScoreKernel::kTransposedChunked:
983 scoreTransposedChunkedRange(xData, d, nLocalTrials, transposedWidth, lo, hi, dst);
985 case ScoreKernel::kSoa:
986 scoreSoaRange(xData, d, nLocalTrials, transposedWidth, lo, hi, dst);
988 case ScoreKernel::kScalar:
989 scoreScalarRange(xData, d, nLocalTrials, transposedWidth, lo, hi, dst);
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();
1018 const ScoreKernel kernel = pickScoreKernel(d, nLocalTrials);
1019 auto sweepLo = [n, sweepBlocks](std::size_t s)
noexcept {
return (n * s) / sweepBlocks; };
1021 constexpr std::size_t kMaxLocalTrials = 32;
1025 if (kernel == ScoreKernel::kTransposedChunked || kernel == ScoreKernel::kScalar) {
1026 (void)candDistRows(n, transposedWidth);
1029 std::vector<std::size_t> candidates(nLocalTrials, 0);
1030 std::vector<T> scores(nLocalTrials, T{0});
1031 const T *refreshRow = centroidsData;
1032 bool roundDegenerate =
false;
1034 auto prePhase = [&](std::size_t phaseIdx)
noexcept {
1035 if (phaseIdx == 0) {
1038 const std::size_t c = (phaseIdx + 1) / 2;
1039 if ((phaseIdx & 1U) != 0) {
1043 for (std::size_t s = 0; s < sweepBlocks; ++s) {
1044 total += sweepSums[s];
1046 roundDegenerate = !(total > T{0});
1047 if (roundDegenerate) {
1049 std::memcpy(centroidsData + (c * d), xData + (pick * d), d *
sizeof(T));
1050 refreshRow = centroidsData + (c * d);
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);
1058 zeroScoreSlabs(workers, scoreSlab);
1062 if (roundDegenerate) {
1065 for (std::size_t t = 0; t < nLocalTrials; ++t) {
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];
1077 const T *winnerRow = xData + (candidates[bestT] * d);
1078 std::memcpy(centroidsData + (c * d), winnerRow, d *
sizeof(T));
1079 refreshRow = winnerRow;
1082 auto phase = [&](std::size_t phaseIdx, std::uint32_t slot, std::size_t lo, std::size_t hi,
1083 void * =
nullptr)
noexcept {
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);
1093 if ((phaseIdx & 1U) != 0) {
1094 if (roundDegenerate) {
1097 scoreRange(kernel, xData, d, nLocalTrials, transposedWidth, sweepLo(lo), sweepLo(hi),
1098 m_localScores.data() + (
static_cast<std::size_t
>(slot) * scoreSlab));
1101 for (std::size_t s = lo; s < hi; ++s) {
1102 refreshSweepBlock(refreshRow, xData, d, sweepLo(s), sweepLo(s + 1), s);
1106 pool.parallelRunPlex<citor::HintsDefaults>(1 + (2 * (k - 1)), sweepBlocks, std::move(phase),
1107 std::move(prePhase));
1111 void ensureShape(std::size_t n, std::size_t d, std::size_t L, std::size_t workers) {
1113 if (m_candRows.dim(0) != L || m_candRows.dim(1) != d) {
1114 m_candRows = NDArray<T, 2, Layout::Contig>({L, d});
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});
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});
1123 if (m_minSq.dim(0) != n) {
1124 m_minSq = NDArray<T, 1>({n});
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});
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});
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});
1147 const std::size_t workersClamped = workers == 0 ? std::size_t{1} : workers;
1153 const std::size_t apSize = gemmScoringUsed
1154 ? (workersClamped * math::detail::kMc<T> * math::detail::kKc<T>)
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});
1161 if (m_gemmBpArena.dim(0) != bpSize) {
1162 m_gemmBpArena = NDArray<T, 1>({bpSize});
1164 const std::size_t scoreSlab = ((lSafe + 15U) / 16U) * 16U;
1165 const std::size_t lsLen = workersClamped * scoreSlab;
1166 if (m_localScores.dim(0) != lsLen) {
1167 m_localScores = NDArray<T, 1>({lsLen});
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;