Skip to content

Commit 73f9d5b

Browse files
authored
Adopt CUDA stream compatibility accessors (#2557)
## Summary Use the `get()` and `sync()` compatibility aliases added in [RMM #2537](rapidsai/rmm#2537). These spellings are shared by `rmm::cuda_stream_view` and `cuda::stream_ref`. This preserves existing stream types and public APIs while extracting mechanical accessor updates from the broader [stream migration](rapidsai/build-planning#318). It is independently buildable without [RMM #2372](rapidsai/rmm#2372) and leaves the migration PR focused on actual type and signature changes. This updates raw CUDA, library, kernel-launch, and legacy API boundaries throughout cuVS while preserving current stream types. Changes that require RAFT to return `cuda::stream_ref` remain in [cuVS #2521](#2521). Authors: - Bradley Dice (https://github.com/bdice) Approvers: - Divye Gala (https://github.com/divyegala) URL: #2557
1 parent 4a45125 commit 73f9d5b

150 files changed

Lines changed: 1139 additions & 1056 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: 1 addition & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -144,7 +144,7 @@ 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: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
11
/*
2-
* SPDX-FileCopyrightText: Copyright (c) 2025, NVIDIA CORPORATION.
2+
* SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
33
* SPDX-License-Identifier: Apache-2.0
44
*/
55

@@ -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: 2 additions & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
11
/*
2-
* SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION.
2+
* SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
33
* SPDX-License-Identifier: Apache-2.0
44
*/
55

@@ -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,

c/tests/neighbors/ann_mg_c.cu

Lines changed: 9 additions & 7 deletions
Original file line numberDiff line numberDiff line change
@@ -1,10 +1,12 @@
11
/*
2-
* SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION.
2+
* SPDX-FileCopyrightText: Copyright (c) 2025-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
33
* SPDX-License-Identifier: Apache-2.0
44
*/
55

66
#include <cuda.h>
77
#include <gtest/gtest.h>
8+
#include <cuda/stream>
9+
810
#include <raft/core/device_mdarray.hpp>
911
#include <raft/core/handle.hpp>
1012
#include <raft/random/rng.cuh>
@@ -158,12 +160,12 @@ class MgCTest : public ::testing::TestWithParam<mg_test_params> {
158160

159161
protected:
160162
mg_test_params params;
161-
rmm::device_uvector<float> index_data{0, rmm::cuda_stream_default};
162-
rmm::device_uvector<float> query_data{0, rmm::cuda_stream_default};
163-
rmm::device_uvector<int64_t> neighbors_data{0, rmm::cuda_stream_default};
164-
rmm::device_uvector<float> distances_data{0, rmm::cuda_stream_default};
165-
rmm::device_uvector<int64_t> ref_neighbors_data{0, rmm::cuda_stream_default};
166-
rmm::device_uvector<float> ref_distances_data{0, rmm::cuda_stream_default};
163+
rmm::device_uvector<float> index_data{0, cuda::stream_ref{cudaStream_t{cudaStreamDefault}}};
164+
rmm::device_uvector<float> query_data{0, cuda::stream_ref{cudaStream_t{cudaStreamDefault}}};
165+
rmm::device_uvector<int64_t> neighbors_data{0, cuda::stream_ref{cudaStream_t{cudaStreamDefault}}};
166+
rmm::device_uvector<float> distances_data{0, cuda::stream_ref{cudaStream_t{cudaStreamDefault}}};
167+
rmm::device_uvector<int64_t> ref_neighbors_data{0, cuda::stream_ref{cudaStream_t{cudaStreamDefault}}};
168+
rmm::device_uvector<float> ref_distances_data{0, cuda::stream_ref{cudaStream_t{cudaStreamDefault}}};
167169

168170
// Host memory for multi-GPU tests
169171
std::vector<float> index_data_host;

cpp/bench/ann/src/common/cuda_huge_page_resource.hpp

Lines changed: 2 additions & 1 deletion
Original file line numberDiff line numberDiff line change
@@ -1,9 +1,10 @@
11
/*
2-
* SPDX-FileCopyrightText: Copyright (c) 2023-2026, NVIDIA CORPORATION.
2+
* SPDX-FileCopyrightText: Copyright (c) 2023-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
33
* SPDX-License-Identifier: Apache-2.0
44
*/
55
#pragma once
66

7+
#include <cuda/stream>
78
#include <raft/core/error.hpp>
89
#include <raft/core/logger_macros.hpp>
910

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/naive_knn.cuh

Lines changed: 1 addition & 2 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
11
/*
2-
* SPDX-FileCopyrightText: Copyright (c) 2023-2026, NVIDIA CORPORATION.
2+
* SPDX-FileCopyrightText: Copyright (c) 2023-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
33
* SPDX-License-Identifier: Apache-2.0
44
*/
55

@@ -11,7 +11,6 @@
1111
#include <raft/util/cuda_utils.cuh>
1212

1313
#include <raft/core/resource/cuda_stream.hpp>
14-
#include <rmm/cuda_stream_view.hpp>
1514
#include <rmm/device_uvector.hpp>
1615
#include <rmm/mr/per_device_resource.hpp>
1716
#include <rmm/resource_ref.hpp>

cpp/src/cluster/detail/agglomerative.cuh

Lines changed: 3 additions & 3 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
11
/*
2-
* SPDX-FileCopyrightText: Copyright (c) 2021-2026, NVIDIA CORPORATION.
2+
* SPDX-FileCopyrightText: Copyright (c) 2021-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
33
* SPDX-License-Identifier: Apache-2.0
44
*/
55

@@ -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: 4 additions & 4 deletions
Original file line numberDiff line numberDiff line change
@@ -1,5 +1,5 @@
11
/*
2-
* SPDX-FileCopyrightText: Copyright (c) 2021-2026, NVIDIA CORPORATION.
2+
* SPDX-FileCopyrightText: Copyright (c) 2021-2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved.
33
* SPDX-License-Identifier: Apache-2.0
44
*/
55

@@ -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;

0 commit comments

Comments
 (0)