102 std::uint64_t seedFirst = 0, std::size_t nInit = 1) {
103 const std::size_t n = X.
dim(0);
104 const std::size_t d = X.
dim(1);
110 ensureOutputShape(n, d);
112 if (n == 0 || d == 0) {
119 const math::Pool pool{m_pool};
122 m_seeder.run(X, m_k, seedFirst, pool, m_centroids);
123 m_lloyd.run(X, m_centroids, m_k, maxIter, tol, pool, m_labels, m_inertia, m_nIter,
128 ensureBestOfScratch(n, d);
129 bool anyImprovement =
false;
131 for (std::size_t i = 0; i < nInit; ++i) {
132 const std::uint64_t seed = seedFirst +
static_cast<std::uint64_t
>(i);
133 m_seeder.run(X, m_k, seed, pool, m_centroids);
134 m_lloyd.run(X, m_centroids, m_k, maxIter, tol, pool, m_labels, m_inertia, m_nIter,
136 if (!anyImprovement || m_inertia < m_bestInertia) {
137 anyImprovement =
true;
138 m_bestInertia = m_inertia;
139 m_bestNIter = m_nIter;
140 m_bestConverged = m_converged;
141 std::memcpy(m_bestCentroids.data(), m_centroids.data(), m_k * d *
sizeof(T));
142 std::memcpy(m_bestLabels.data(), m_labels.data(), n *
sizeof(std::int32_t));
149 std::swap(m_centroids, m_bestCentroids);
150 std::swap(m_labels, m_bestLabels);
151 m_inertia = m_bestInertia;
152 m_nIter = m_bestNIter;
153 m_converged = m_bestConverged;
197 m_bestCentroids = NDArray<T, 2, Layout::Contig>({m_k, d});
215 NDArray<T, 2, Layout::Contig> m_bestCentroids{NDArray<T, 2, Layout::Contig>({0, 0})};