diff --git a/cpp/src/cluster/detail/kmeans.cuh b/cpp/src/cluster/detail/kmeans.cuh index e3ffb4a439..6787f9a421 100644 --- a/cpp/src/cluster/detail/kmeans.cuh +++ b/cpp/src/cluster/detail/kmeans.cuh @@ -394,34 +394,14 @@ void initScalableKMeansPlusPlus(raft::resources const& handle, int niter = std::min(8, (int)ceil(log(psi))); RAFT_LOG_DEBUG("KMeans||: psi = %g, log(psi) = %g, niter = %d ", psi, log(psi), niter); + auto newMinClusterDistanceVec = raft::make_device_vector(handle, n_samples); + // <<<< Step-3 >>> : for O( log(psi) ) times do for (int iter = 0; iter < niter; ++iter) { RAFT_LOG_DEBUG("KMeans|| - Iteration %d: # potential centroids sampled - %d", iter, potentialCentroids.extent(0)); - cuvs::cluster::kmeans::detail::minClusterDistanceCompute( - handle, - X, - potentialCentroids, - minClusterDistanceVec.view(), - L2NormX.view(), - L2NormBuf_OR_DistBuf, - params.metric, - params.batch_samples, - params.batch_centroids, - workspace); - - cuvs::cluster::kmeans::detail::computeClusterCost( - handle, - minClusterDistanceVec.view(), - workspace, - raft::make_device_scalar_view(clusterCost.data()), - raft::identity_op{}, - raft::add_op{}); - - psi = clusterCost.value(stream); - // <<<< Step-4 >>> : Sample each point x in X independently and identify new // potentialCentroids raft::random::uniform( @@ -458,6 +438,37 @@ void initScalableKMeansPlusPlus(raft::resources const& handle, potentialCentroids = raft::make_device_matrix_view(centroidsBuf.data(), tot_centroids, n_features); /// <<<< End of Step-5 >>> + + // Update d(x, C) using only the newly sampled candidates. + if (Cp.extent(0) > 0 && iter + 1 < niter) { + cuvs::cluster::kmeans::detail::minClusterDistanceCompute( + handle, + X, + Cp, + newMinClusterDistanceVec.view(), + L2NormX.view(), + L2NormBuf_OR_DistBuf, + params.metric, + params.batch_samples, + params.batch_centroids, + workspace); + + raft::linalg::map(handle, + minClusterDistanceVec.view(), + raft::min_op{}, + raft::make_const_mdspan(minClusterDistanceVec.view()), + raft::make_const_mdspan(newMinClusterDistanceVec.view())); + + cuvs::cluster::kmeans::detail::computeClusterCost( + handle, + minClusterDistanceVec.view(), + workspace, + raft::make_device_scalar_view(clusterCost.data()), + raft::identity_op{}, + raft::add_op{}); + + psi = clusterCost.value(stream); + } } /// <<<< Step-6 >>> RAFT_LOG_DEBUG("KMeans||: total # potential centroids sampled - %d",