From 229ddd39d8da1b2fb12dc7cdddebcce5c69272c5 Mon Sep 17 00:00:00 2001 From: Bradley Dice Date: Mon, 5 Oct 2026 01:05:52 -0500 Subject: [PATCH] Run the NN-descent per-iteration host update on the calling thread GNND::build started and joined a std::thread on every iteration to run update_graph and sample_graph while the GPU ran add_reverse_edges and local_join. Each such thread brought up a second OpenMP team that competed with the calling thread's team for the same cores. Enqueue the iteration's (asynchronous) GPU work first and then run the host update on the calling thread; the two halves touch disjoint buffers, so the overlap and all synchronization points are unchanged, and exceptions from local_join now propagate instead of terminating the process. --- cpp/src/neighbors/detail/nn_descent.cuh | 9 ++++----- 1 file changed, 4 insertions(+), 5 deletions(-) diff --git a/cpp/src/neighbors/detail/nn_descent.cuh b/cpp/src/neighbors/detail/nn_descent.cuh index ad9ea2e573..c400bec3a2 100644 --- a/cpp/src/neighbors/detail/nn_descent.cuh +++ b/cpp/src/neighbors/detail/nn_descent.cuh @@ -2751,8 +2751,6 @@ void GNND::build(Data_t* data, raft::copy(res, d_list_sizes_old_.view(), graph_.h_list_sizes_old.view()); raft::resource::sync_stream(res); - std::thread update_and_sample_thread(update_and_sample, it); - RAFT_LOG_DEBUG("# GNND iteration: %lu / %lu", it + 1, build_config_.max_iterations); // Reuse dists_buffer_ to save GPU memory. graph_buffer_ cannot be reused, because it @@ -2786,7 +2784,8 @@ void GNND::build(Data_t* data, THROW("NN_DESCENT cannot be run for __CUDA_ARCH__ < 700"); } - update_and_sample_thread.join(); + // Overlaps with the device work enqueued above, which uses disjoint buffers. + update_and_sample(it > 0); if (update_counter_ == -1) { break; } raft::copy(res, graph_host_buffer_.view(), graph_buffer_.view()); @@ -2897,7 +2896,6 @@ void GNND::build( raft::copy(res, d_list_sizes_old_.view(), graph_.h_list_sizes_old.view()); raft::resource::sync_stream(res); - std::thread update_and_sample_thread(update_and_sample, it); RAFT_LOG_DEBUG("# GNND iteration: %lu / %lu", it + 1, build_config_.max_iterations); static_assert(DEGREE_ON_DEVICE * sizeof(*(dists_buffer_.data_handle())) >= @@ -2914,7 +2912,8 @@ void GNND::build( stream); local_join(stream, dataset, dist_epilogue); - update_and_sample_thread.join(); + // Overlaps with the device work enqueued above, which uses disjoint buffers. + update_and_sample(it > 0); if (update_counter_ == -1) { break; } raft::copy(res, graph_host_buffer_.view(), graph_buffer_.view()); raft::copy(res, dists_host_buffer_.view(), dists_buffer_.view());