Skip to content

Commit 4415405

Browse files
committed
Migrate stream view APIs to cuda::stream_ref
1 parent 438e660 commit 4415405

166 files changed

Lines changed: 1195 additions & 1101 deletions

File tree

Some content is hidden

Large Commits have some content hidden by default. Use the searchbox below for content that may be hidden.

‎c/src/core/c_api.cpp‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -136,15 +136,15 @@ extern "C" cuvsError_t cuvsStreamSet(cuvsResources_t res, cudaStream_t stream)
136136
{
137137
return cuvs::core::translate_exceptions([=] {
138138
auto res_ptr = reinterpret_cast<raft::resources*>(res);
139-
raft::resource::set_cuda_stream(*res_ptr, static_cast<rmm::cuda_stream_view>(stream));
139+
raft::resource::set_cuda_stream(*res_ptr, static_cast<cuda::stream_ref>(stream));
140140
});
141141
}
142142

143143
extern "C" cuvsError_t cuvsStreamGet(cuvsResources_t res, cudaStream_t* stream)
144144
{
145145
return cuvs::core::translate_exceptions([=] {
146146
auto res_ptr = reinterpret_cast<raft::resources*>(res);
147-
*stream = raft::resource::get_cuda_stream(*res_ptr);
147+
*stream = raft::resource::get_cuda_stream(*res_ptr).get();
148148
});
149149
}
150150

‎c/src/neighbors/nn_descent.cpp‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -107,7 +107,7 @@ void _get_distances(cuvsResources_t res, cuvsNNDescentIndex_t index, DLManagedTe
107107
src->data_handle(),
108108
dst.extent(0) * dst.extent(1) * sizeof(float),
109109
cudaMemcpyDefault,
110-
raft::resource::get_cuda_stream(*res_ptr));
110+
raft::resource::get_cuda_stream(*res_ptr).get());
111111

112112
} else {
113113
RAFT_FAIL("Unsupported nn-descent index dtype: %d and bits: %d", dtype.code, dtype.bits);

‎c/tests/neighbors/ann_ivf_sq_c.cu‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -108,7 +108,7 @@ TEST(IvfSqC, BuildSearch)
108108

109109
cuvsResources_t res;
110110
cuvsResourcesCreate(&res);
111-
cuvsStreamSet(res, stream);
111+
cuvsStreamSet(res, stream.get());
112112

113113
run_ivf_sq(res,
114114
n_rows,

‎cpp/bench/ann/src/cuvs/cuvs_ann_bench_utils.h‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -122,8 +122,8 @@ class configured_raft_resources {
122122
*/
123123
explicit configured_raft_resources(const std::shared_ptr<shared_raft_resources>& shared_res)
124124
: shared_res_{shared_res},
125-
res_{std::make_unique<raft::device_resources>(
126-
rmm::cuda_stream_view(get_stream_from_global_pool()))}
125+
res_{
126+
std::make_unique<raft::device_resources>(cuda::stream_ref(get_stream_from_global_pool()))}
127127
{
128128
raft::resource::set_large_workspace_resource(
129129
*res_, raft::mr::device_resource{shared_res_->get_large_memory_resource()});

‎cpp/include/cuvs/neighbors/common.hpp‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1168,7 +1168,7 @@ auto make_device_dense_row_major_dataset_from_src(raft::resources const& res,
11681168
RAFT_CUDA_TRY(cudaMemsetAsync(out_array.data_handle(),
11691169
0,
11701170
out_array.size() * sizeof(ValueT),
1171-
raft::resource::get_cuda_stream(res)));
1171+
raft::resource::get_cuda_stream(res).get()));
11721172
raft::copy_matrix(out_array.data_handle(),
11731173
target_stride,
11741174
src.data_handle(),

‎cpp/internal/cuvs_internal/neighbors/refine_helper.cuh‎

Lines changed: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -127,7 +127,7 @@ class RefineHelper {
127127
public:
128128
RefineInputs<IdxT> p;
129129
const raft::resources& handle_;
130-
rmm::cuda_stream_view stream_;
130+
cuda::stream_ref stream_;
131131

132132
raft::device_matrix<DataT, IdxT, row_major> dataset;
133133
raft::device_matrix<DataT, IdxT, row_major> queries;

‎cpp/src/cluster/detail/agglomerative.cuh‎

Lines changed: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -280,7 +280,7 @@ void extract_flattened_clusters(raft::resources const& handle,
280280
rmm::device_uvector<value_idx> levels(n_vertices, stream);
281281

282282
value_idx n_blocks = raft::ceildiv(n_vertices, (value_idx)tpb);
283-
write_levels_kernel<<<n_blocks, tpb, 0, stream>>>(children, levels.data(), n_vertices);
283+
write_levels_kernel<<<n_blocks, tpb, 0, stream.get()>>>(children, levels.data(), n_vertices);
284284
/**
285285
* Step 1: Find label roots:
286286
*
@@ -323,7 +323,7 @@ void extract_flattened_clusters(raft::resources const& handle,
323323
*/
324324
value_idx cut_level = (n_edges / 2) - (n_clusters - 1);
325325

326-
inherit_labels<<<n_blocks, tpb, 0, stream>>>(
326+
inherit_labels<<<n_blocks, tpb, 0, stream.get()>>>(
327327
children, levels.data(), n_leaves, tmp_labels.data(), cut_level, n_vertices);
328328

329329
// copy tmp labels to actual labels

‎cpp/src/cluster/detail/connectivities.cuh‎

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -61,7 +61,7 @@ struct distance_graph_impl<Linkage::KNN_GRAPH, value_idx, value_t> {
6161
auto thrust_policy = raft::resource::get_thrust_policy(handle);
6262

6363
// Need to symmetrize knn into undirected graph
64-
raft::sparse::COO<value_t, value_idx> knn_graph_coo(stream);
64+
raft::sparse::COO<value_t, value_idx> knn_graph_coo(stream.get());
6565

6666
auto X_view = raft::make_device_matrix_view<const value_t, value_idx, raft::row_major>(X, m, n);
6767
cuvs::neighbors::detail::knn_graph<value_idx, value_t, size_t>(
@@ -92,7 +92,7 @@ struct distance_graph_impl<Linkage::KNN_GRAPH, value_idx, value_t> {
9292
raft::make_const_mdspan(vals_in_view));
9393

9494
raft::sparse::convert::sorted_coo_to_csr(
95-
knn_graph_coo.rows(), knn_graph_coo.nnz, indptr.data(), m + 1, stream);
95+
knn_graph_coo.rows(), knn_graph_coo.nnz, indptr.data(), m + 1, stream.get());
9696

9797
// TODO: Wouldn't need to copy here if we could compute knn
9898
// graph directly on the device uvectors
@@ -140,7 +140,7 @@ void pairwise_distances(const raft::resources& handle,
140140
value_idx nnz = m * m;
141141

142142
value_idx blocks = raft::ceildiv(nnz, (value_idx)256);
143-
fill_indices2<value_idx><<<blocks, 256, 0, stream>>>(indices, m, nnz);
143+
fill_indices2<value_idx><<<blocks, 256, 0, stream.get()>>>(indices, m, nnz);
144144

145145
raft::linalg::map_offset(handle,
146146
raft::make_device_vector_view<value_idx, value_idx>(indptr, m),

‎cpp/src/cluster/detail/kmeans.cuh‎

Lines changed: 5 additions & 5 deletions
Original file line numberDiff line numberDiff line change
@@ -95,7 +95,7 @@ void kmeansPlusPlus(raft::resources const& handle,
9595
rmm::device_uvector<char>& workspace)
9696
{
9797
raft::common::nvtx::range<cuvs::common::nvtx::domain::cuvs> fun_scope("kmeansPlusPlus");
98-
cudaStream_t stream = raft::resource::get_cuda_stream(handle);
98+
cudaStream_t stream = raft::resource::get_cuda_stream(handle).get();
9999
auto n_samples = X.extent(0);
100100
auto n_features = X.extent(1);
101101
auto n_clusters = params.n_clusters;
@@ -309,7 +309,7 @@ void initScalableKMeansPlusPlus(raft::resources const& handle,
309309
{
310310
raft::common::nvtx::range<cuvs::common::nvtx::domain::cuvs> fun_scope(
311311
"initScalableKMeansPlusPlus");
312-
cudaStream_t stream = raft::resource::get_cuda_stream(handle);
312+
cudaStream_t stream = raft::resource::get_cuda_stream(handle).get();
313313
auto n_samples = X.extent(0);
314314
auto n_features = X.extent(1);
315315
auto n_clusters = params.n_clusters;
@@ -573,7 +573,7 @@ void kmeans_fit(
573573
auto n_features = X.extent(1);
574574
auto n_clusters = pams.n_clusters;
575575
auto metric = pams.metric;
576-
cudaStream_t stream = raft::resource::get_cuda_stream(handle);
576+
cudaStream_t stream = raft::resource::get_cuda_stream(handle).get();
577577

578578
if (sample_weight.has_value())
579579
RAFT_EXPECTS(sample_weight.value().extent(0) == n_samples,
@@ -1035,7 +1035,7 @@ void kmeans_predict(raft::resources const& handle,
10351035
raft::common::nvtx::range<cuvs::common::nvtx::domain::cuvs> fun_scope("kmeans_predict");
10361036
auto n_samples = X.extent(0);
10371037
auto n_features = X.extent(1);
1038-
cudaStream_t stream = raft::resource::get_cuda_stream(handle);
1038+
cudaStream_t stream = raft::resource::get_cuda_stream(handle).get();
10391039
// Check that parameters are valid
10401040
if (sample_weight.has_value())
10411041
RAFT_EXPECTS(sample_weight.value().extent(0) == n_samples,
@@ -1186,7 +1186,7 @@ void kmeans_transform(raft::resources const& handle,
11861186
"kmeans only supports L2Expanded or L2SqrtExpanded distance metrics.");
11871187
raft::common::nvtx::range<cuvs::common::nvtx::domain::cuvs> fun_scope("kmeans_transform");
11881188
raft::default_logger().set_level(pams.verbosity);
1189-
cudaStream_t stream = raft::resource::get_cuda_stream(handle);
1189+
cudaStream_t stream = raft::resource::get_cuda_stream(handle).get();
11901190
auto n_samples = X.extent(0);
11911191
auto n_features = X.extent(1);
11921192
auto n_clusters = pams.n_clusters;

‎cpp/src/cluster/detail/kmeans_balanced.cuh‎

Lines changed: 49 additions & 32 deletions
Original file line numberDiff line numberDiff line change
@@ -142,7 +142,7 @@ inline std::enable_if_t<std::is_floating_point_v<MathT>> predict_core(
142142
&beta,
143143
distances.data(),
144144
n_clusters,
145-
stream);
145+
stream.get());
146146

147147
auto distances_const_view = raft::make_device_matrix_view<const MathT, IdxT, raft::row_major>(
148148
distances.data(), n_rows, n_clusters);
@@ -286,13 +286,29 @@ void calc_centers_and_sizes(const raft::resources& handle,
286286

287287
// Apply mapping only when the data and math types are different.
288288
if constexpr (std::is_same_v<T, MathT>) {
289-
raft::linalg::reduce_rows_by_key(
290-
dataset, dim, labels, nullptr, n_rows, dim, n_clusters, centers, stream, reset_counters);
289+
raft::linalg::reduce_rows_by_key(dataset,
290+
dim,
291+
labels,
292+
nullptr,
293+
n_rows,
294+
dim,
295+
n_clusters,
296+
centers,
297+
stream.get(),
298+
reset_counters);
291299
} else {
292300
// todo(lsugy): use iterator from KV output of fusedL2NN
293301
thrust::transform_iterator<MappingOpT, const T*> mapping_itr(dataset, mapping_op);
294-
raft::linalg::reduce_rows_by_key(
295-
mapping_itr, dim, labels, nullptr, n_rows, dim, n_clusters, centers, stream, reset_counters);
302+
raft::linalg::reduce_rows_by_key(mapping_itr,
303+
dim,
304+
labels,
305+
nullptr,
306+
n_rows,
307+
dim,
308+
n_clusters,
309+
centers,
310+
stream.get(),
311+
reset_counters);
296312
}
297313

298314
// Compute weight of each cluster
@@ -689,38 +705,39 @@ auto adjust_centers(const raft::resources& handle,
689705
search_count.set_value_to_zero_async(stream);
690706
const dim3 grid_dim(raft::ceildiv(n_clusters, static_cast<IdxT>(kBlockDimY)), 1, 1);
691707
adjust_centers_random_donor_kernel<kBlockDimY>
692-
<<<grid_dim, block_dim, 0, stream>>>(centers,
693-
n_clusters,
694-
dim,
695-
dataset,
696-
n_rows,
697-
labels,
698-
cluster_sizes,
699-
lower_threshold,
700-
static_cast<IdxT>(n_rows / n_clusters),
701-
centroid_offset,
702-
ofst,
703-
search_count.data(),
704-
update_count.data(),
705-
mapping_op);
708+
<<<grid_dim, block_dim, 0, stream.get()>>>(centers,
709+
n_clusters,
710+
dim,
711+
dataset,
712+
n_rows,
713+
labels,
714+
cluster_sizes,
715+
lower_threshold,
716+
static_cast<IdxT>(n_rows / n_clusters),
717+
centroid_offset,
718+
ofst,
719+
search_count.data(),
720+
update_count.data(),
721+
mapping_op);
706722
return update_count.value(stream) > 0; // NB: rmm scalar performs the sync
707723
}
708724

709725
raft::update_device(receiver_clusters.data(), host_receiver_clusters.data(), n_pairs, stream);
710726
raft::update_device(donor_clusters.data(), host_donor_clusters.data(), n_pairs, stream);
711727
const dim3 grid_dim(raft::ceildiv(n_pairs, static_cast<IdxT>(kBlockDimY)), 1, 1);
712-
adjust_centers_kernel<kBlockDimY><<<grid_dim, block_dim, 0, stream>>>(centers,
713-
n_pairs,
714-
dim,
715-
dataset,
716-
n_rows,
717-
labels,
718-
receiver_clusters.data(),
719-
donor_clusters.data(),
720-
centroid_offset,
721-
ofst,
722-
update_count.data(),
723-
mapping_op);
728+
adjust_centers_kernel<kBlockDimY>
729+
<<<grid_dim, block_dim, 0, stream.get()>>>(centers,
730+
n_pairs,
731+
dim,
732+
dataset,
733+
n_rows,
734+
labels,
735+
receiver_clusters.data(),
736+
donor_clusters.data(),
737+
centroid_offset,
738+
ofst,
739+
update_count.data(),
740+
mapping_op);
724741
auto n_updates = update_count.value(stream); // NB: rmm scalar performs the sync
725742
RAFT_EXPECTS(n_updates == n_pairs, "Balanced k-means failed to update all adjusted centers");
726743
return n_updates > 0;
@@ -1068,7 +1085,7 @@ auto build_fine_clusters(const raft::resources& handle,
10681085
}
10691086

10701087
thrust::transform_iterator<MappingOpT, const T*> mapping_itr(dataset_mptr, mapping_op);
1071-
raft::matrix::gather(mapping_itr, dim, n_rows, mc_trainset_ids, k, mc_trainset, stream);
1088+
raft::matrix::gather(mapping_itr, dim, n_rows, mc_trainset_ids, k, mc_trainset, stream.get());
10721089
if (params.metric == cuvs::distance::DistanceType::L2Expanded ||
10731090
params.metric == cuvs::distance::DistanceType::L2SqrtExpanded ||
10741091
params.metric == cuvs::distance::DistanceType::CosineExpanded) {

0 commit comments

Comments
 (0)