-
Notifications
You must be signed in to change notification settings - Fork 243
Exposing opt-in double buffering for MG batched KMeans #2323
New issue
Have a question about this project? Sign up for a free GitHub account to open an issue and contact its maintainers and the community.
By clicking “Sign up for GitHub”, you agree to our terms of service and privacy statement. We’ll occasionally send you account related emails.
Already on GitHub? Sign in to your account
Changes from 1 commit
07e7dfc
26121fc
e9f0cb2
3d60e58
66bbad0
9a92778
File filter
Filter by extension
Conversations
Jump to
Diff view
Diff view
There are no files selected for viewing
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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; | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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.
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. The problem with
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more. What are the meanings of batch samples vs batch size?
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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.
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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).
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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.
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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.
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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.
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe reason will be displayed to describe this comment to others. Learn more.
I created this: #2328 |
||
| return kmeans_params; | ||
| } | ||
|
|
||
|
|
||
| Original file line number | Diff line number | Diff line change |
|---|---|---|
|
|
@@ -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> | ||
|
|
@@ -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); | ||
|
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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:
Contributor
There was a problem hiding this comment. Choose a reason for hiding this commentThe 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; | ||
|
|
@@ -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()); | ||
|
|
||
|
|
@@ -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; | ||
|
|
||
|
|
@@ -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()); | ||
|
|
||
|
|
@@ -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); | ||
|
|
||
Uh oh!
There was an error while loading. Please reload this page.