Skip to content
Closed
Show file tree
Hide file tree
Changes from 1 commit
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
8 changes: 7 additions & 1 deletion c/include/cuvs/cluster/kmeans.h
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
/*
* SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION.
* SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: Apache-2.0
*/

Expand Down Expand Up @@ -199,6 +199,12 @@ struct cuvsKMeansParams {
* or n_samples for device data.
*/
int64_t init_size;

/**
* Whether host-resident multi-GPU KMeans should prefetch the next streaming
Comment thread
viclafargue marked this conversation as resolved.
* batch using a second device buffer. Ignored by other KMeans paths.
*/
bool streaming_batch_prefetch;
};

typedef struct cuvsKMeansParams* cuvsKMeansParams_t;
Expand Down
5 changes: 3 additions & 2 deletions c/src/cluster/kmeans.cpp
Original file line number Diff line number Diff line change
@@ -1,5 +1,5 @@
/*
* SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION.
* SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
* SPDX-License-Identifier: Apache-2.0
*/

Expand Down Expand Up @@ -316,7 +316,8 @@ extern "C" cuvsError_t cuvsKMeansParamsCreate_v2(cuvsKMeansParams_v2_t* params)
.hierarchical = false,
.hierarchical_n_iters = static_cast<int>(cpp_balanced_params.n_iters),
.streaming_batch_size = cpp_params.streaming_batch_size,
.init_size = cpp_params.init_size};
.init_size = cpp_params.init_size,
.streaming_batch_prefetch = cpp_params.streaming_batch_prefetch};
});
}

Expand Down
1 change: 1 addition & 0 deletions c/src/cluster/mg_kmeans.cpp
Original file line number Diff line number Diff line change
Expand Up @@ -38,6 +38,7 @@ cuvs::cluster::kmeans::params convert_params(const ParamsT& params)
kmeans_params.batch_centroids = params.batch_centroids;
kmeans_params.init_size = params.init_size;
kmeans_params.streaming_batch_size = params.streaming_batch_size;
kmeans_params.streaming_batch_prefetch = params.streaming_batch_prefetch;

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I really don't like the use of the term "streaming" here. It carries altogether too much baggage with it. Can we please just remove it here? I think we need to consider removing it from "streaming_batch_size" as well. Streaming to me means a k-means that runs through a streaming mechanism, meaning it can only go forward as an iterator and can restart from where it left off when more data is received. There are variants of k-means that satisfy this, but the batched variant we've implemented does not.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

The problem with batch_size is that we have two other parameters -- batch_centroids and batch_samples. So batch_size can be confusing. Perhaps we should reconsider the names of those two parameters then? Or maybe think of some other prefix for batch_size?

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

What are the meanings of batch samples vs batch size?

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

the tile sizes for the distance computation (local reductions are done on this tile) -- i.e. batch_samples * batch_centroids distances are computed at once. So I was saying that we are already doing what the first part of the flash-kmeans paper mentions (at least that was my understanding)

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

  /**
   * batch_samples and batch_centroids are used to tile 1NN computation which is
   * useful to optimize/control the memory footprint
   * Default tile is [batch_samples x n_clusters] i.e. when batch_centroids is 0
   * then don't tile the centroids
   *
   * NB: These parameters are unrelated to streaming_batch_size, which controls how many
   * samples to transfer from host to device per batch when processing out-of-core
   * data.
   */
  int batch_samples = 1 << 15;

  /**
   * if 0 then batch_centroids = n_clusters
   */
  int batch_centroids = 0;

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

BTW I'm not necessarily say we should or shouldn't make it the default, but I am saying we should look into the impact and provide real evidence either way (concrete numbers / formula always helps but benchmarks should ultimately guide us).

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

If we are renaming the streaming_batch_size param, we must do it quick because its in the C layer. If we dont change it this release, it'll be locked user experience for the next 6 months due to ABI stability.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Agree we should try to get this out in time for the release, but just in case it doesn't, we can always add the new parameter and handle it accordingly underneath until we are able to remove the old. Think about it like a deprecation.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Just in case, I suppose we could always open up a smaller PR to do just the rename. Let's see how things go next week.

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

I suppose we could always open up a smaller PR to do just the rename

I created this: #2328

return kmeans_params;
}

Expand Down
2 changes: 2 additions & 0 deletions c/tests/cluster/kmeans_mg_c.cu
Original file line number Diff line number Diff line change
Expand Up @@ -73,11 +73,13 @@ void test_mg_fit_host()

typename Api::params_t params;
ASSERT_EQ(Api::params_create(&params), CUVS_SUCCESS);
EXPECT_FALSE(params->streaming_batch_prefetch);
params->n_clusters = kNClusters;
params->max_iter = 100;
params->tol = 1e-6;
params->init = Array;
params->streaming_batch_size = 4; // force at least 2 streamed batches
params->streaming_batch_prefetch = true;

DLManagedTensor dataset_t{};
cuvs::core::to_dlpack(raft::make_host_matrix_view<float, int64_t>(
Expand Down
10 changes: 10 additions & 0 deletions cpp/include/cuvs/cluster/kmeans.hpp
Original file line number Diff line number Diff line change
Expand Up @@ -155,6 +155,16 @@ struct params : base_params {
* Default: 0 (process all data at once).
*/
int64_t streaming_batch_size = 0;

/**
* Whether host-resident multi-GPU KMeans should prefetch the next streaming
Comment thread
viclafargue marked this conversation as resolved.
* batch on a separate CUDA stream. Enabling this overlaps H2D transfer with
* computation by allocating a second device batch buffer on each rank.
*
* This option is ignored by single-GPU and device-resident fits.
* Default: false.
*/
bool streaming_batch_prefetch = false;
};

/**
Expand Down
39 changes: 32 additions & 7 deletions cpp/src/cluster/detail/kmeans_mg.cuh
Original file line number Diff line number Diff line change
Expand Up @@ -37,13 +37,15 @@
#include <raft/random/rng.cuh>
#include <raft/util/cudart_utils.hpp>

#include <rmm/cuda_stream.hpp>
#include <rmm/device_scalar.hpp>
#include <rmm/device_uvector.hpp>

#include <algorithm>
#include <cmath>
#include <cstdint>
#include <limits>
#include <memory>
#include <numeric>
#include <optional>
#include <random>
Expand Down Expand Up @@ -140,7 +142,18 @@ void mnmg_fit(
use_nccl ? raft::resource::set_current_device_to_rank(handle, rank) : handle;
mnmg_comms comms{dev_res, use_nccl, nccl_comm};

auto stream = comms.stream();
auto stream = comms.stream();
std::unique_ptr<rmm::cuda_stream> data_copy_stream_owner;
rmm::cuda_stream_view data_copy_stream{stream};
bool enable_data_prefetch = false;
if constexpr (!data_on_device) {
if (params.streaming_batch_prefetch) {
data_copy_stream_owner =
std::make_unique<rmm::cuda_stream>(rmm::cuda_stream::flags::non_blocking);

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Rather than creating the prefetch stream, lets just get it from the handle:
cuvs::spatial::knn::detail::utils::get_prefetch_stream(handle);

Copy link
Copy Markdown
Contributor

Choose a reason for hiding this comment

The reason will be displayed to describe this comment to others. Learn more.

Also i dont think we need an explicit argument for opt-in prefetch. If an extra stream is supplied in the pool, we simply use it for prefetch, else we dont prefetch. This has been the pattern in IVF-PQ, IVF-Flat and others.

data_copy_stream = data_copy_stream_owner->view();
enable_data_prefetch = true;
}
}
auto n_features = centroids.extent(1);
auto n_clusters = static_cast<IndexT>(params.n_clusters);
auto metric = params.metric;
Expand Down Expand Up @@ -405,11 +418,15 @@ void mnmg_fit(
data_batch_iterator_t data_batches(dev_res,
X_part,
static_cast<size_t>(streaming_batch_size),
stream,
data_copy_stream,
rmm::mr::get_current_device_resource_ref(),
true);
enable_data_prefetch);
auto data_it = data_batches.begin();
auto data_end = data_batches.end();
data_batches.prefetch_next_batch();

for (auto const& data_batch : data_batches) {
for (; data_it != data_end; ++data_it) {
auto const& data_batch = *data_it;
IndexT current_batch_size = static_cast<IndexT>(data_batch.size());
auto batch_offset = static_cast<IndexT>(data_batch.offset());

Expand Down Expand Up @@ -468,7 +485,9 @@ void mnmg_fit(
weight_per_cluster.view(),
raft::make_device_scalar_view(clustering_cost.data_handle()),
batch_workspace);
data_batches.prefetch_next_batch();
}
if (enable_data_prefetch) { raft::resource::sync_stream(dev_res); }
}
norms_cached = true;

Expand Down Expand Up @@ -533,11 +552,15 @@ void mnmg_fit(
data_batch_iterator_t data_batches(dev_res,
X_part,
static_cast<size_t>(streaming_batch_size),
stream,
data_copy_stream,
rmm::mr::get_current_device_resource_ref(),
true);
enable_data_prefetch);
auto data_it = data_batches.begin();
auto data_end = data_batches.end();
data_batches.prefetch_next_batch();

for (auto const& data_batch : data_batches) {
for (; data_it != data_end; ++data_it) {
auto const& data_batch = *data_it;
IndexT current_batch_size = static_cast<IndexT>(data_batch.size());
auto batch_offset = static_cast<IndexT>(data_batch.offset());

Expand All @@ -561,7 +584,9 @@ void mnmg_fit(
raft::make_const_mdspan(clustering_cost.view()),
raft::make_const_mdspan(batch_clustering_cost.view()),
clustering_cost.view());
data_batches.prefetch_next_batch();
}
if (enable_data_prefetch) { raft::resource::sync_stream(dev_res); }
}
comms.allreduce(clustering_cost.data_handle(), clustering_cost.data_handle(), 1);
raft::copy(&local_inertia, clustering_cost.data_handle(), 1, stream);
Expand Down
26 changes: 19 additions & 7 deletions cpp/tests/cluster/kmeans_mg.cu
Original file line number Diff line number Diff line change
Expand Up @@ -42,6 +42,7 @@ struct KmeansSNMGInputs {
int n_init;
cuvs::cluster::kmeans::params::InitMethod init = cuvs::cluster::kmeans::params::Array;
int max_iter = 20;
bool streaming_batch_prefetch = false;
};

template <typename T>
Expand Down Expand Up @@ -105,13 +106,14 @@ class KmeansSNMGTest : public ::testing::TestWithParam<KmeansSNMGInputs<T>> {
}

cuvs::cluster::kmeans::params snmg_params;
snmg_params.n_clusters = n_clusters;
snmg_params.tol = testparams_.tol;
snmg_params.max_iter = testparams_.max_iter;
snmg_params.n_init = testparams_.n_init;
snmg_params.rng_state.seed = 42;
snmg_params.init = testparams_.init;
snmg_params.streaming_batch_size = testparams_.streaming_batch_size;
snmg_params.n_clusters = n_clusters;
snmg_params.tol = testparams_.tol;
snmg_params.max_iter = testparams_.max_iter;
snmg_params.n_init = testparams_.n_init;
snmg_params.rng_state.seed = 42;
snmg_params.init = testparams_.init;
snmg_params.streaming_batch_size = testparams_.streaming_batch_size;
snmg_params.streaming_batch_prefetch = testparams_.streaming_batch_prefetch;

T snmg_inertia = T{0};
int64_t snmg_n_iter = 0;
Expand Down Expand Up @@ -281,6 +283,16 @@ const std::vector<KmeansSNMGInputs<float>> snmg_inputsf = {
{1000, 32, 5, 0.0001f, kmeans_weight_mode::none, 1000, 1},
{1000, 32, 5, 0.0001f, kmeans_weight_mode::uniform, 1000, 1},
{1000, 32, 5, 0.0001f, kmeans_weight_mode::none, 128, 1},
{1000,
32,
5,
0.0001f,
kmeans_weight_mode::none,
128,
1,
cuvs::cluster::kmeans::params::Array,
20,
true},
Comment thread
viclafargue marked this conversation as resolved.
{10000, 16, 10, 0.0001f, kmeans_weight_mode::none, 2000, 1},
{10000, 16, 10, 0.0001f, kmeans_weight_mode::uniform, 2000, 1},
{10000, 16, 10, 0.0001f, kmeans_weight_mode::none, 500, 1},
Expand Down
4 changes: 3 additions & 1 deletion python/cuvs/cuvs/cluster/kmeans/kmeans.pxd
Original file line number Diff line number Diff line change
Expand Up @@ -56,7 +56,8 @@ cdef extern from "cuvs/cluster/kmeans.h" nogil:
bool hierarchical,
int hierarchical_n_iters,
int64_t streaming_batch_size,
int64_t init_size
int64_t init_size,
bool streaming_batch_prefetch

ctypedef cuvsKMeansParams* cuvsKMeansParams_t
ctypedef cuvsKMeansParams_v2* cuvsKMeansParams_v2_t
Expand Down Expand Up @@ -90,3 +91,4 @@ cdef extern from "cuvs/cluster/kmeans.h" nogil:

cdef class KMeansParams:
cdef cuvsKMeansParams* params
cdef bool _streaming_batch_prefetch
15 changes: 15 additions & 0 deletions python/cuvs/cuvs/cluster/kmeans/kmeans.pyx
Original file line number Diff line number Diff line change
Expand Up @@ -94,13 +94,21 @@ cdef class KMeansParams:
increases.

Default: 0 (process all data at once).
streaming_batch_prefetch : bool
Whether host-resident multi-GPU KMeans should use a second device
Comment thread
viclafargue marked this conversation as resolved.
buffer to prefetch the next streaming batch and overlap H2D transfer
with computation. This can improve throughput at the cost of one
additional device batch buffer per GPU. Ignored by other KMeans paths.

Default: False.
hierarchical : bool
Whether to use hierarchical (balanced) kmeans or not
hierarchical_n_iters : int
For hierarchical k-means , defines the number of training iterations
"""

def __cinit__(self):
self._streaming_batch_prefetch = False
cuvsKMeansParamsCreate(&self.params)

def __dealloc__(self):
Expand All @@ -119,6 +127,7 @@ cdef class KMeansParams:
inertia_check=None,
init_size=None,
streaming_batch_size=None,
streaming_batch_prefetch=None,
hierarchical=None,
hierarchical_n_iters=None):
if metric is not None:
Expand Down Expand Up @@ -150,6 +159,8 @@ cdef class KMeansParams:
self.params.init_size = init_size
if streaming_batch_size is not None:
self.params.streaming_batch_size = streaming_batch_size
if streaming_batch_prefetch is not None:
self._streaming_batch_prefetch = streaming_batch_prefetch
if hierarchical is not None:
self.params.hierarchical = hierarchical
if hierarchical_n_iters is not None:
Expand Down Expand Up @@ -202,6 +213,10 @@ cdef class KMeansParams:
def streaming_batch_size(self):
return self.params.streaming_batch_size

@property
def streaming_batch_prefetch(self):
return self._streaming_batch_prefetch

@property
def hierarchical(self):
return self.params.hierarchical
Expand Down
12 changes: 12 additions & 0 deletions python/cuvs/cuvs/cluster/mg/kmeans/kmeans.pyx
Original file line number Diff line number Diff line change
Expand Up @@ -76,6 +76,17 @@ def fit(
``centroids`` is a host NumPy array containing the computed centroids,
``inertia`` is the final objective value, and ``n_iter`` is the number
of iterations run.

Notes
-----
For small streaming batches, call ``resources.set_memory_pool(...)``
before ``fit`` to reduce allocator contention between GPU worker threads.
Memory pools are not enabled automatically because they replace the
process-wide RMM resource on each managed device.

Set ``params.streaming_batch_prefetch=True`` to overlap H2D transfer with
computation. This allocates a second device batch buffer on every rank;
the default single-buffered path minimizes device-memory usage.
"""

if params.hierarchical:
Expand Down Expand Up @@ -153,6 +164,7 @@ def fit(
params_v2.hierarchical_n_iters = params.params.hierarchical_n_iters
params_v2.streaming_batch_size = params.params.streaming_batch_size
params_v2.init_size = params.params.init_size
params_v2.streaming_batch_prefetch = params._streaming_batch_prefetch

with cuda_interruptible():
check_cuvs(cuvsMultiGpuKMeansFit(
Expand Down
7 changes: 6 additions & 1 deletion python/cuvs/cuvs/tests/test_mg_kmeans.py
Original file line number Diff line number Diff line change
Expand Up @@ -93,7 +93,10 @@ def assert_inertia_matches_centroids(out, X, sample_weights):
@pytest.mark.parametrize("dtype", [np.float32, np.float64])
@pytest.mark.parametrize("init_method", ["Array", "KMeansPlusPlus", "Random"])
@pytest.mark.parametrize("weighted", [False, True])
def test_mg_kmeans_fit_options(dtype, init_method, weighted):
@pytest.mark.parametrize("streaming_batch_prefetch", [False, True])
def test_mg_kmeans_fit_options(
dtype, init_method, weighted, streaming_batch_prefetch
):
n_clusters = 4
X, initial_centroids = make_inputs(dtype, n_clusters=n_clusters)
resources = MultiGpuResources()
Expand All @@ -111,7 +114,9 @@ def test_mg_kmeans_fit_options(dtype, init_method, weighted):
n_init=3 if init_method == "Random" else 1,
init_size=X.shape[0],
streaming_batch_size=37,
streaming_batch_prefetch=streaming_batch_prefetch,
)
assert params.streaming_batch_prefetch == streaming_batch_prefetch
centroids = initial_centroids.copy() if init_method == "Array" else None

mg_out = mg_kmeans.fit(
Expand Down
Loading