Skip to content

Migrate stream APIs from rmm::cuda_stream_view to cuda::stream_ref - #5639

Open
bdice wants to merge 4 commits into
rapidsai:mainfrom
bdice:cuda-stream-ref
Open

Migrate stream APIs from rmm::cuda_stream_view to cuda::stream_ref#5639
bdice wants to merge 4 commits into
rapidsai:mainfrom
bdice:cuda-stream-ref

Conversation

@bdice

@bdice bdice commented Aug 28, 2026

Copy link
Copy Markdown
Contributor

Summary

Track the coordinated migration of stream APIs and call sites from rmm::cuda_stream_view to CCCL's cuda::stream_ref. This propagates cuda::stream_ref through RMM containers and memory resources, RAFT resource and handle APIs, downstream C++ interfaces, Python/Cython bindings, benchmarks, tests, and documentation.

This migrates affected cuGraph API signatures and call sites, including the MTMG handle stream accessor, while preserving cuda::stream_ref until a raw stream is required.

Depends on rapidsai/rmm#2372 and NVIDIA/raft#3129.

Tracked in rapidsai/build-planning#318.

Migrations

  • Pass cuda::stream_ref through stream pools, resource accessors, conditionals, and downstream APIs without converting to rmm::cuda_stream_view
  • Use cuda::stream_ref constructions for default/legacy/per-thread streams
    • rmm::cuda_stream_default ➡️ cuda::stream_ref{cudaStream_t{cudaStreamDefault}}
    • rmm::cuda_stream_legacy ➡️ cuda::stream_ref{cudaStreamLegacy}
    • rmm::cuda_stream_per_thread ➡️ cuda::stream_ref{cudaStreamPerThread}
  • Use .get() when calling an API that requires a raw cudaStream_t, including CUDA runtime, library, CUB, and legacy API boundaries (previously rmm::cuda_stream_view used value())
  • Use .sync() when synchronizing a cuda::stream_ref (previously rmm::cuda_stream_view used synchronize())
  • Update Cython declarations and call sites to pass stream references directly where supported

@copy-pr-bot

copy-pr-bot Bot commented Aug 28, 2026

Copy link
Copy Markdown

Auto-sync is disabled for draft pull requests in this repository. Workflows must be run manually.

Contributors can view more details about this message here.

Signed-off-by: Bradley Dice <bdice@bradleydice.com>
@bdice bdice changed the title Use cuda::stream_ref for pooled streams Migrate stream APIs from rmm::cuda_stream_view to cuda::stream_ref Sep 2, 2026
@bdice
bdice marked this pull request as ready for review September 2, 2026 22:55
@bdice
bdice requested a review from a team as a code owner September 2, 2026 22:55
@bdice bdice added breaking Breaking change improvement Improvement / enhancement to an existing function labels Sep 3, 2026
rapids-bot Bot pushed a commit to rapidsai/cugraph-gnn that referenced this pull request Sep 3, 2026
)

## Summary

Track the coordinated migration of stream APIs and call sites from `rmm::cuda_stream_view` to CCCL's `cuda::stream_ref`. This propagates `cuda::stream_ref` through RMM containers and memory resources, RAFT resource and handle APIs, downstream C++ interfaces, Python/Cython bindings, benchmarks, tests, and documentation.

This updates the cuGraph-GNN developer guide to document `cuda::stream_ref` as the stream type used by stream-ordered operations.

Depends on rapidsai/cugraph#5639.

Tracked in rapidsai/build-planning#318.

## Migrations

- Pass `cuda::stream_ref` through stream pools, resource accessors, conditionals, and downstream APIs without converting to `rmm::cuda_stream_view`
- Use `cuda::stream_ref` constructions for default/legacy/per-thread streams
  - `rmm::cuda_stream_default` ➡️ `cuda::stream_ref{cudaStream_t{cudaStreamDefault}}`
  - `rmm::cuda_stream_legacy` ➡️ `cuda::stream_ref{cudaStreamLegacy}`
  - `rmm::cuda_stream_per_thread` ➡️ `cuda::stream_ref{cudaStreamPerThread}`
- Use `.get()` when calling an API that requires a raw `cudaStream_t`, including CUDA runtime, library, CUB, and legacy API boundaries (previously `rmm::cuda_stream_view` used `value()`)
- Use `.sync()` when synchronizing a `cuda::stream_ref` (previously `rmm::cuda_stream_view` used `synchronize()`)
- Update Cython declarations and call sites to pass stream references directly where supported

Authors:
  - Bradley Dice (https://github.com/bdice)

Approvers:
  - Alex Barghi (https://github.com/alexbarghi-nv)

URL: #529

@ChuckHastings ChuckHastings left a comment

Copy link
Copy Markdown
Collaborator

Choose a reason for hiding this comment

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

LGTM

return raft_handle_.is_stream_pool_initialized()
? raft_handle_.get_stream_from_stream_pool(thread_rank_)
: raft_handle_.get_stream();
: static_cast<cuda::stream_ref>(raft_handle_.get_stream());

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.

This is a bit confusing to me.

raft_handle_.get_stream_from_stream_pool() returns cuda::stream_ref but raft_handle_.get_stream() returns something else and needs to be type-casted? Better return the same type?

host_scalar_allgather(major_comm, v_list_range[0], handle.get_stream().get());
auto local_v_list_range_lasts =
host_scalar_allgather(major_comm, v_list_range[1], handle.get_stream());
host_scalar_allgather(major_comm, v_list_range[1], handle.get_stream().get());

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.

To be consistent with the rest of the API (use cuda::stream_ref as much as possible till raw cuda stream is absolutely necessary), should we better update host_scalar_allgather and other host scalar utility functions to take cuda::stream_ref and get the raw stream right before calling the NCCL functions in raft?

size_t mem_frugal_threshold, // take the memory frugal approach (instead of thrust::sort) if #
// elements to groupby is no smaller than this value
rmm::cuda_stream_view stream_view,
cuda::stream_ref stream_view,

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.

We have been mixing stream_view and stream. stream_view now sounds like a misnomer. Better use stream consistently?

raft::update_host(rx_counts.data(), d_rx_value_counts.data(), comm_size, stream_view.value());
stream_view.synchronize();
raft::update_host(tx_counts.data(), d_tx_value_counts.data(), comm_size, stream_view.get());
raft::update_host(rx_counts.data(), d_rx_value_counts.data(), comm_size, stream_view.get());

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.

Should we better update raft::update_host to take cuda::stream_ref?

std::vector<size_t> h_offsets(d_offsets.size() + 2);
raft::update_host(h_offsets.data() + 1, d_offsets.data(), d_offsets.size(), stream_view);
RAFT_CUDA_TRY(cudaStreamSynchronize(stream_view));
RAFT_CUDA_TRY(cudaStreamSynchronize(stream_view.get()));

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.

Instead of directly calling cudaStreamSynchronize(), should we better call stream_view.sync()? I assume the philosophy here is to avoid using raw CUDA stream unless it is absolutely necessary?

raft::update_host(
rx_aligned_counts.data(), d_rx_aligned_counts.data(), d_rx_aligned_counts.size(), stream_view);
RAFT_CUDA_TRY(cudaStreamSynchronize(stream_view));
RAFT_CUDA_TRY(cudaStreamSynchronize(stream_view.get()));

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.

Won't stream_view.sync() be better?

Sign up for free to join this conversation on GitHub. Already have an account? Sign in to comment

Labels

breaking Breaking change improvement Improvement / enhancement to an existing function

Projects

None yet

Development

Successfully merging this pull request may close these issues.

4 participants