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());