From b1775fafa16a079de77a2a1899116dcbbf67ab39 Mon Sep 17 00:00:00 2001 From: Hyunsu Cho Date: Fri, 18 Sep 2026 23:33:08 +0000 Subject: [PATCH 1/4] Fix build for mg C++ tests --- cpp/tests/mg/knn.cu | 10 +++++----- cpp/tests/mg/knn_test_helper.cuh | 6 +++--- cpp/tests/mg/pca.cu | 19 ++++++++++--------- 3 files changed, 18 insertions(+), 17 deletions(-) diff --git a/cpp/tests/mg/knn.cu b/cpp/tests/mg/knn.cu index b38dc40cdc..5167eeb13f 100644 --- a/cpp/tests/mg/knn.cu +++ b/cpp/tests/mg/knn.cu @@ -51,7 +51,7 @@ class BruteForceKNNTest : public ::testing::TestWithParam { raft::comms::initialize_mpi_comms(&handle, MPI_COMM_WORLD); const auto& comm = handle.get_comms(); - cudaStream_t stream = handle.get_stream(); + auto stream = handle.get_stream(); int my_rank = comm.get_rank(); int size = comm.get_size(); @@ -124,7 +124,7 @@ class BruteForceKNNTest : public ::testing::TestWithParam { out_d_parts.push_back(out_d); out_i_parts.push_back(out_i); - generate_partition(query_d, params.min_rows, params.n_cols, 5, stream); + generate_partition(query_d, params.min_rows, params.n_cols, 5, stream.get()); } std::vector index_parts; @@ -140,7 +140,7 @@ class BruteForceKNNTest : public ::testing::TestWithParam { index_parts.push_back(i_d); - generate_partition(i_d, params.min_rows, params.n_cols, 5, stream); + generate_partition(i_d, params.min_rows, params.n_cols, 5, stream.get()); } Matrix::PartDescriptor idx_desc( @@ -169,8 +169,8 @@ class BruteForceKNNTest : public ::testing::TestWithParam { handle.sync_stream(stream); - std::cout << raft::arr2Str(out_i_parts[0]->ptr, 10, "final_out_I", stream) << std::endl; - std::cout << raft::arr2Str(out_d_parts[0]->ptr, 10, "final_out_D", stream) << std::endl; + std::cout << raft::arr2Str(out_i_parts[0]->ptr, 10, "final_out_I", stream.get()) << std::endl; + std::cout << raft::arr2Str(out_d_parts[0]->ptr, 10, "final_out_D", stream.get()) << std::endl; /** * Verify expected results diff --git a/cpp/tests/mg/knn_test_helper.cuh b/cpp/tests/mg/knn_test_helper.cuh index 4438104b7d..bf8c0f458a 100644 --- a/cpp/tests/mg/knn_test_helper.cuh +++ b/cpp/tests/mg/knn_test_helper.cuh @@ -118,7 +118,7 @@ class KNNTestHelper { params.n_cols, params.n_classes, my_rank, - this->stream); + this->stream.get()); y.resize(this->index_parts_per_rank); for (int i = 0; i < this->index_parts_per_rank; i++) { @@ -173,7 +173,7 @@ class KNNTestHelper { std::cout << "Finished!" << std::endl; - std::cout << raft::arr2Str(out_parts[0]->ptr, 10, "final_out", stream) << std::endl; + std::cout << raft::arr2Str(out_parts[0]->ptr, 10, "final_out", stream.get()) << std::endl; } void release_ressources(const KNNParams& params) @@ -237,7 +237,7 @@ class KNNTestHelper { Matrix::PartDescriptor* query_desc = nullptr; std::vector> y; - cudaStream_t stream = 0; + cuda::stream_ref stream; private: int index_parts_per_rank; diff --git a/cpp/tests/mg/pca.cu b/cpp/tests/mg/pca.cu index cc115f1a73..a4b5a88e61 100644 --- a/cpp/tests/mg/pca.cu +++ b/cpp/tests/mg/pca.cu @@ -50,7 +50,7 @@ class PCAOpgTest : public testing::TestWithParam { totalRanks = comm.get_size(); raft::random::Rng r(params.seed + myRank); - RAFT_CUBLAS_TRY(cublasSetStream(cublasHandle, stream)); + RAFT_CUBLAS_TRY(cublasSetStream(cublasHandle, stream.get())); if (myRank == 0) { std::cout << "Testing PCA of " << params.M << " x " << params.N << " matrix" << std::endl; @@ -66,8 +66,8 @@ class PCAOpgTest : public testing::TestWithParam { Matrix::PartDescriptor desc( params.M, params.N, totalPartsToRanks, comm.get_rank(), params.layout); std::vector*> inParts; - Matrix::opg::allocate(handle, inParts, desc, myRank, stream); - Matrix::opg::randomize(handle, r, inParts, desc, myRank, stream, T(10.0), T(20.0)); + Matrix::opg::allocate(handle, inParts, desc, myRank, stream.get()); + Matrix::opg::randomize(handle, r, inParts, desc, myRank, stream.get(), T(10.0), T(20.0)); handle.sync_stream(); prmsPCA.n_rows = params.M; @@ -103,28 +103,29 @@ class PCAOpgTest : public testing::TestWithParam { false); CUML_LOG_DEBUG( - raft::arr2Str(singular_vals.data(), params.N_components, "Singular Vals", stream).c_str()); + raft::arr2Str(singular_vals.data(), params.N_components, "Singular Vals", stream.get()) + .c_str()); CUML_LOG_DEBUG( - raft::arr2Str(explained_var.data(), params.N_components, "Explained Variance", stream) + raft::arr2Str(explained_var.data(), params.N_components, "Explained Variance", stream.get()) .c_str()); CUML_LOG_DEBUG( raft::arr2Str( - explained_var_ratio.data(), params.N_components, "Explained Variance Ratio", stream) + explained_var_ratio.data(), params.N_components, "Explained Variance Ratio", stream.get()) .c_str()); CUML_LOG_DEBUG( - raft::arr2Str(components.data(), params.N_components * params.N, "Components", stream) + raft::arr2Str(components.data(), params.N_components * params.N, "Components", stream.get()) .c_str()); - Matrix::opg::deallocate(handle, inParts, desc, myRank, stream); + Matrix::opg::deallocate(handle, inParts, desc, myRank, stream.get()); } protected: PCAOpgParams params; raft::handle_t handle; - cudaStream_t stream = 0; + cuda::stream_ref stream; int myRank; int totalRanks; ML::paramsPCAMG prmsPCA; From a4bc83d72406c701032558c41a748a37fa3df49a Mon Sep 17 00:00:00 2001 From: Hyunsu Cho Date: Fri, 18 Sep 2026 23:48:56 +0000 Subject: [PATCH 2/4] Explicitly construct default stream --- cpp/tests/mg/knn_test_helper.cuh | 2 +- cpp/tests/mg/pca.cu | 2 +- 2 files changed, 2 insertions(+), 2 deletions(-) diff --git a/cpp/tests/mg/knn_test_helper.cuh b/cpp/tests/mg/knn_test_helper.cuh index bf8c0f458a..f04fd15350 100644 --- a/cpp/tests/mg/knn_test_helper.cuh +++ b/cpp/tests/mg/knn_test_helper.cuh @@ -237,7 +237,7 @@ class KNNTestHelper { Matrix::PartDescriptor* query_desc = nullptr; std::vector> y; - cuda::stream_ref stream; + cuda::stream_ref stream{cudaStream_t{nullptr}}; private: int index_parts_per_rank; diff --git a/cpp/tests/mg/pca.cu b/cpp/tests/mg/pca.cu index a4b5a88e61..5f01a558c2 100644 --- a/cpp/tests/mg/pca.cu +++ b/cpp/tests/mg/pca.cu @@ -125,7 +125,7 @@ class PCAOpgTest : public testing::TestWithParam { protected: PCAOpgParams params; raft::handle_t handle; - cuda::stream_ref stream; + cuda::stream_ref stream{cudaStream_t{nullptr}}; int myRank; int totalRanks; ML::paramsPCAMG prmsPCA; From bd759ed1ce915bddb988e9c51c784747c3c36b3d Mon Sep 17 00:00:00 2001 From: Hyunsu Cho Date: Wed, 23 Sep 2026 04:14:59 +0000 Subject: [PATCH 3/4] [Dask RF] Implement weighted row sampling using global CDF --- cpp/src/randomforest/randomforest.cuh | 245 +++++++++++++++++++++++--- cpp/tests/CMakeLists.txt | 3 + cpp/tests/mg/rf_row_sampler.cu | 214 ++++++++++++++++++++++ cpp/tests/mg/rf_test.cu | 40 +---- cpp/tests/mg/rf_test_utils.hpp | 57 ++++++ 5 files changed, 496 insertions(+), 63 deletions(-) create mode 100644 cpp/tests/mg/rf_row_sampler.cu create mode 100644 cpp/tests/mg/rf_test_utils.hpp diff --git a/cpp/src/randomforest/randomforest.cuh b/cpp/src/randomforest/randomforest.cuh index 4fdfb730ee..2a2a588e64 100644 --- a/cpp/src/randomforest/randomforest.cuh +++ b/cpp/src/randomforest/randomforest.cuh @@ -26,6 +26,7 @@ #include #include #include +#include #include #include #include @@ -46,6 +47,7 @@ #include #include #include +#include namespace ML { @@ -60,6 +62,20 @@ struct NonzeroSampleWeight { __device__ bool operator()(T weight) const { return weight != T(0); } }; +template +struct SubtractOffset { + T offset; + + __device__ T operator()(T value) const { return value - offset; } +}; + +template +struct IsPositiveAndLessThan { + T upper; + + __device__ bool operator()(T value) const { return value >= T{0} && value < upper; } +}; + // Matches estimator behavior: when bootstrapping is enabled and sample weights exist, // those weights are materialized by drawing bootstrap rows according to them. class RowSampler { @@ -77,12 +93,21 @@ class RowSampler { n_sampled_rows_(n_sampled_rows), bootstrap_masks_(bootstrap_masks), sample_weight_(sample_weight), + distributed_(raft::resource::comms_initialized(handle) && handle.get_comms().get_size() > 1), + rank_(distributed_ ? handle.get_comms().get_rank() : 0), + comm_size_(distributed_ ? handle.get_comms().get_size() : 1), + global_n_rows_(n_rows), + rank_row_offset_(0), + local_sample_weight_sum_(0.0), sample_weight_sum_(0.0), + rank_weight_offset_(0.0), sample_weight_cdf_(0, handle.get_stream()) { ASSERT(bootstrap_masks_ == nullptr || DT::is_dev_ptr(bootstrap_masks_), "bootstrap_masks must be a GPU pointer"); - validate_sample_weight(handle, sample_weight_, n_rows_); + validate_distributed_inputs(handle); + compute_global_row_counts(handle); + validate_sample_weight(handle, sample_weight_, n_rows_, distributed_); if (use_weighted_bootstrap()) { sample_weight_cdf_.resize(ML::narrow_cast(n_rows_), handle.get_stream()); thrust::inclusive_scan(rmm::exec_policy(handle.get_stream()), @@ -93,8 +118,8 @@ class RowSampler { // Empty distributed partitions still pass a non-null pointer so every rank selects the same // weighted objective type, but they have no local weight sum to validate. - if (sample_weight_ != nullptr && n_rows_ > 0) { - sample_weight_sum_ = compute_sample_weight_sum(handle); + if (sample_weight_ != nullptr) { + compute_global_sample_weights(handle); ASSERT(sample_weight_sum_ > 0.0, "sample_weight values must contain at least one positive value"); } @@ -105,6 +130,11 @@ class RowSampler { selected_rows_.emplace_back(n_sampled_rows_size, stream); if (use_weighted_bootstrap()) { weighted_draw_scratch_.emplace_back(n_sampled_rows_size, stream); + if (distributed_) { + rank_local_weighted_draw_scratch_.emplace_back(n_sampled_rows_size, stream); + } + } else if (distributed_ && bootstrap_) { + uniform_draw_scratch_.emplace_back(n_sampled_rows_size, stream); } } } @@ -117,7 +147,10 @@ class RowSampler { raft::common::nvtx::range fun_scope("bootstrapping row IDs @randomforest.cuh"); auto& selected_rows = selected_rows_[stream_id]; - if (n_rows_ == 0) { return selected_rows; } + if (n_rows_ == 0) { + selected_rows.resize(0, stream); + return selected_rows; + } raft::resources stream_resources; raft::resource::set_cuda_stream(stream_resources, stream); @@ -137,16 +170,62 @@ class RowSampler { weighted_draw_scratch.size(), 0.0, sample_weight_sum_); - thrust::upper_bound(rmm::exec_policy(stream), - sample_weight_cdf_.data(), - sample_weight_cdf_.data() + n_rows_, - weighted_draw_scratch.begin(), - weighted_draw_scratch.end(), - selected_rows.begin()); + if (distributed_) { + // Each rank filters weighted_draw_scratch and only keeps the element in the range + // [rank_weight_offset_, rank_weight_offset_ + local_sample_weight_sum_). + // This ensures that the rank selects only the samples that are local to the rank. + selected_rows.resize(ML::narrow_cast(n_sampled_rows_), stream); + auto local_draws_begin = thrust::make_transform_iterator( + weighted_draw_scratch.begin(), SubtractOffset{rank_weight_offset_}); + auto& rank_local_draws = rank_local_weighted_draw_scratch_[stream_id]; + auto rank_local_draw_end = + thrust::copy_if(rmm::exec_policy(stream), + local_draws_begin, + local_draws_begin + weighted_draw_scratch.size(), + rank_local_draws.begin(), + IsPositiveAndLessThan{local_sample_weight_sum_}); + auto n_rank_local_draws = rank_local_draw_end - rank_local_draws.begin(); + selected_rows.resize(n_rank_local_draws, stream); + thrust::upper_bound(rmm::exec_policy(stream), + sample_weight_cdf_.data(), + sample_weight_cdf_.data() + n_rows_, + rank_local_draws.begin(), + rank_local_draw_end, + selected_rows.begin()); + } else { + thrust::upper_bound(rmm::exec_policy(stream), + sample_weight_cdf_.data(), + sample_weight_cdf_.data() + n_rows_, + weighted_draw_scratch.begin(), + weighted_draw_scratch.end(), + selected_rows.begin()); + } } else if (bootstrap_) { // Draw bootstrap rows uniformly when there are no sample weights. - raft::random::uniformInt( - stream_resources, rng_state, selected_rows.data(), selected_rows.size(), 0, n_rows_); + if (distributed_) { + // Each rank filters uniform_draw_scratch and only keeps the element in the range + // [rank_row_offset_, rank_row_offset_ + n_rows_). + // This ensures that the rank selects only the samples that are local to the rank. + auto& uniform_draw_scratch = uniform_draw_scratch_[stream_id]; + raft::random::uniformInt(stream_resources, + rng_state, + uniform_draw_scratch.data(), + uniform_draw_scratch.size(), + 0, + global_n_rows_); + selected_rows.resize(ML::narrow_cast(n_sampled_rows_), stream); + auto local_rows_begin = thrust::make_transform_iterator( + uniform_draw_scratch.begin(), SubtractOffset{rank_row_offset_}); + auto selected_rows_end = thrust::copy_if(rmm::exec_policy(stream), + local_rows_begin, + local_rows_begin + uniform_draw_scratch.size(), + selected_rows.begin(), + IsPositiveAndLessThan{n_rows_}); + selected_rows.resize(selected_rows_end - selected_rows.begin(), stream); + } else { + raft::random::uniformInt( + stream_resources, rng_state, selected_rows.data(), selected_rows.size(), 0, n_rows_); + } } else if (sample_weight_ != nullptr) { // Remove zero-weight rows from the non-bootstrap row set. selected_rows.resize(ML::narrow_cast(n_sampled_rows_), stream); @@ -188,23 +267,101 @@ class RowSampler { tree_mask); } - double compute_sample_weight_sum(const raft::handle_t& handle) const + void validate_distributed_inputs(const raft::handle_t& handle) const { - if (use_weighted_bootstrap()) { - double weight_sum = 0.0; - raft::update_host( - &weight_sum, sample_weight_cdf_.data() + n_rows_ - 1, 1, handle.get_stream()); - handle.sync_stream(); - return weight_sum; + ASSERT(n_rows_ >= 0, "n_rows must be non-negative"); + ASSERT(n_sampled_rows_ >= 0, "n_sampled_rows must be non-negative"); + if (!distributed_) { return; } + + auto stream = handle.get_stream().get(); + rmm::device_uvector local_values(2, stream); + rmm::device_uvector gathered_values(ML::checked_mul(2, comm_size_), + stream); + std::int64_t h_local_values[2] = {sample_weight_ == nullptr ? 0 : 1, + bootstrap_ ? n_sampled_rows_ : 0}; + raft::update_device(local_values.data(), h_local_values, 2, stream); + handle.get_comms().allgather(local_values.data(), gathered_values.data(), 2, stream); + ASSERT(handle.get_comms().sync_stream(stream) == raft::comms::status_t::SUCCESS, + "An error occurred while validating distributed RF row-sampler inputs."); + + std::vector h_gathered_values(gathered_values.size()); + raft::update_host( + h_gathered_values.data(), gathered_values.data(), gathered_values.size(), stream); + handle.sync_stream(stream); + for (int i = 0; i < comm_size_; ++i) { + ASSERT(h_gathered_values[2 * i] == h_gathered_values[0], + "sample_weight must be supplied consistently on every rank"); + if (bootstrap_) { + ASSERT(h_gathered_values[2 * i + 1] == h_gathered_values[1], + "n_sampled_rows must be identical on every rank when bootstrapping"); + } + } + } + + void compute_global_row_counts(const raft::handle_t& handle) + { + if (!distributed_) { return; } + + auto stream = handle.get_stream().get(); + rmm::device_uvector local_row_count(1, stream); + rmm::device_uvector rank_row_counts(comm_size_, stream); + raft::update_device(local_row_count.data(), &n_rows_, 1, stream); + handle.get_comms().allgather(local_row_count.data(), rank_row_counts.data(), 1, stream); + ASSERT(handle.get_comms().sync_stream(stream) == raft::comms::status_t::SUCCESS, + "An error occurred in the distributed RF row-count all-gather."); + + // Compute the total row count over all ranks + global_n_rows_ = thrust::reduce( + rmm::exec_policy(stream), rank_row_counts.begin(), rank_row_counts.end(), std::int64_t{0}); + + // Compute the sum of row counts in ranks 0, 1, ..., (rank_ - 1). + rank_row_offset_ = thrust::reduce(rmm::exec_policy(stream), + rank_row_counts.begin(), + rank_row_counts.begin() + rank_, + std::int64_t{0}); + ASSERT(global_n_rows_ > 0, "global row count must be positive"); + } + + void compute_global_sample_weights(const raft::handle_t& handle) + { + if (n_rows_ > 0) { + if (use_weighted_bootstrap()) { + raft::update_host(&local_sample_weight_sum_, + sample_weight_cdf_.data() + n_rows_ - 1, + 1, + handle.get_stream()); + handle.sync_stream(); + } else { + local_sample_weight_sum_ = thrust::reduce( + rmm::exec_policy(handle.get_stream()), sample_weight_, sample_weight_ + n_rows_, 0.0); + } + } + + if (!distributed_) { + sample_weight_sum_ = local_sample_weight_sum_; + return; } - return thrust::reduce( - rmm::exec_policy(handle.get_stream()), sample_weight_, sample_weight_ + n_rows_, 0.0); + auto stream = handle.get_stream().get(); + rmm::device_uvector local_weight_sum(1, stream); + rmm::device_uvector rank_weight_sums(comm_size_, stream); + raft::update_device(local_weight_sum.data(), &local_sample_weight_sum_, 1, stream); + handle.get_comms().allgather(local_weight_sum.data(), rank_weight_sums.data(), 1, stream); + ASSERT(handle.get_comms().sync_stream(stream) == raft::comms::status_t::SUCCESS, + "An error occurred in the distributed RF weight-sum all-gather."); + + // Compute the sum of sample weights in all ranks + sample_weight_sum_ = thrust::reduce( + rmm::exec_policy(stream), rank_weight_sums.begin(), rank_weight_sums.end(), 0.0); + // Compute the sum of sample weights in ranks 0, 1, ..., (rank_ - 1). + rank_weight_offset_ = thrust::reduce( + rmm::exec_policy(stream), rank_weight_sums.begin(), rank_weight_sums.begin() + rank_, 0.0); } static void validate_sample_weight(const raft::handle_t& handle, const double* sample_weight, - std::int64_t n_rows) + std::int64_t n_rows, + bool distributed) { ASSERT(sample_weight == nullptr || DT::is_dev_ptr(sample_weight), "sample_weight must be a GPU pointer"); @@ -214,6 +371,22 @@ class RowSampler { sample_weight, sample_weight + n_rows, InvalidSampleWeight{}); + if (distributed) { + int invalid_status = has_invalid ? 1 : 0; + rmm::device_uvector d_invalid_status(1, handle.get_stream()); + raft::update_device(d_invalid_status.data(), &invalid_status, 1, handle.get_stream()); + handle.get_comms().allreduce(d_invalid_status.data(), + d_invalid_status.data(), + 1, + raft::comms::op_t::MAX, + handle.get_stream().get()); + ASSERT( + handle.get_comms().sync_stream(handle.get_stream().get()) == raft::comms::status_t::SUCCESS, + "An error occurred while validating distributed RF sample weights."); + raft::update_host(&invalid_status, d_invalid_status.data(), 1, handle.get_stream()); + handle.sync_stream(); + has_invalid = invalid_status != 0; + } ASSERT(!has_invalid, "sample_weight values must be finite and non-negative"); } @@ -225,10 +398,19 @@ class RowSampler { std::int64_t n_sampled_rows_; bool* bootstrap_masks_; const double* sample_weight_; + bool distributed_; + int rank_; + int comm_size_; + std::int64_t global_n_rows_; + std::int64_t rank_row_offset_; + double local_sample_weight_sum_; double sample_weight_sum_; + double rank_weight_offset_; rmm::device_uvector sample_weight_cdf_; std::deque> selected_rows_; std::deque> weighted_draw_scratch_; + std::deque> rank_local_weighted_draw_scratch_; + std::deque> uniform_draw_scratch_; }; } // namespace detail @@ -313,10 +495,25 @@ class RandomForest { raft::resource::comms_initialized(handle) && handle.get_comms().get_size() > 1; this->error_checking(input, labels, n_rows, n_cols, false, distributed); std::int64_t const n_rows_i64 = n_rows; - std::int64_t n_sampled_rows = 0; + std::int64_t global_n_rows = n_rows_i64; + if (distributed) { + rmm::device_uvector d_global_n_rows(1, handle.get_stream()); + raft::update_device(d_global_n_rows.data(), &global_n_rows, 1, handle.get_stream()); + handle.get_comms().allreduce(d_global_n_rows.data(), + d_global_n_rows.data(), + 1, + raft::comms::op_t::SUM, + handle.get_stream().get()); + ASSERT( + handle.get_comms().sync_stream(handle.get_stream().get()) == raft::comms::status_t::SUCCESS, + "An error occurred in the distributed RF global row-count all-reduce."); + raft::update_host(&global_n_rows, d_global_n_rows.data(), 1, handle.get_stream()); + handle.sync_stream(); + } + std::int64_t n_sampled_rows = 0; if (this->rf_params.bootstrap) { n_sampled_rows = - static_cast(std::round(this->rf_params.max_samples * n_rows_i64)); + static_cast(std::round(this->rf_params.max_samples * global_n_rows)); } else { if (this->rf_params.max_samples != 1.0) { CUML_LOG_WARN( diff --git a/cpp/tests/CMakeLists.txt b/cpp/tests/CMakeLists.txt index efc2c5e447..951367e656 100644 --- a/cpp/tests/CMakeLists.txt +++ b/cpp/tests/CMakeLists.txt @@ -229,6 +229,9 @@ if(BUILD_CUML_MG_TESTS) ConfigureTest( PREFIX MG NAME RF_QUANTILE_TEST mg/rf_quantile_test.cu MPI RAFT_DISTRIBUTED ML_INCLUDE ) + ConfigureTest( + PREFIX MG NAME RF_ROW_SAMPLER_TEST mg/rf_row_sampler.cu MPI RAFT_DISTRIBUTED ML_INCLUDE + ) ConfigureTest(PREFIX MG NAME RF_TEST mg/rf_test.cu MPI RAFT_DISTRIBUTED ML_INCLUDE) else(MPI_CXX_FOUND) message("OpenMPI not found. Skipping MultiGPU tests '${CUML_MG_TEST_TARGET}'") diff --git a/cpp/tests/mg/rf_row_sampler.cu b/cpp/tests/mg/rf_row_sampler.cu new file mode 100644 index 0000000000..bded2a2ffd --- /dev/null +++ b/cpp/tests/mg/rf_row_sampler.cu @@ -0,0 +1,214 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#include "rf_test_utils.hpp" +#include "test_opg_utils.h" + +#include +#include +#include + +#include +#include + +#include + +#include +#include +#include + +#include +#include +#include +#include +#include +#include +#include + +namespace ML { +namespace Test { +namespace opg { + +enum class WeightKind { ThreeRows, Uniform, WithZeros, Skewed }; + +struct RowSamplerTestParams { + PartitionKind partition_kind; + WeightKind weight_kind; +}; + +void initialize_mpi_once() +{ + int mpi_initialized = 0; + MPI_Initialized(&mpi_initialized); + if (!mpi_initialized) { MPI_Init(nullptr, nullptr); } +} + +void get_mpi_local_rank_size(int& local_rank, int& local_size) +{ + MPI_Comm local_comm{}; + MPI_Comm_split_type(MPI_COMM_WORLD, MPI_COMM_TYPE_SHARED, 0, MPI_INFO_NULL, &local_comm); + MPI_Comm_rank(local_comm, &local_rank); + MPI_Comm_size(local_comm, &local_size); + MPI_Comm_free(&local_comm); +} + +std::vector make_weights(WeightKind kind) +{ + switch (kind) { + case WeightKind::ThreeRows: return {0.1, 0.4, 0.5}; + case WeightKind::Uniform: return std::vector(23, 1.0); + case WeightKind::WithZeros: + return {0.0, 1.0, 0.0, 2.0, 4.0, 0.0, 3.0, 0.0, 5.0, 1.0, 0.0, 4.0, 0.0}; + case WeightKind::Skewed: + return {0.001, 0.002, 0.004, 0.008, 0.016, 0.032, 0.064, 0.128, 0.245, 0.5}; + } + return {}; +} + +class RfMgRowSamplerTest : public ::testing::TestWithParam {}; + +TEST_P(RfMgRowSamplerTest, SamplesGlobalWeightDistribution) +{ + initialize_mpi_once(); + int rank = 0; + int size = 1; + MPI_Comm_rank(MPI_COMM_WORLD, &rank); + MPI_Comm_size(MPI_COMM_WORLD, &size); + + int local_rank = 0; + int local_size = 1; + get_mpi_local_rank_size(local_rank, local_size); + int n_gpus = 0; + RAFT_CUDA_TRY(cudaGetDeviceCount(&n_gpus)); + ASSERT_GE(n_gpus, local_size); + RAFT_CUDA_TRY(cudaSetDevice(local_rank)); + + auto stream_pool = std::make_shared(1); + raft::handle_t handle(cuda::stream_ref{cudaStreamPerThread}, stream_pool); + raft::comms::initialize_mpi_comms(&handle, MPI_COMM_WORLD); + + auto global_weights = make_weights(GetParam().weight_kind); + auto local_rows = local_rows_for_rank( + static_cast(global_weights.size()), rank, size, GetParam().partition_kind); + std::vector local_weights(local_rows.size()); + std::transform(local_rows.begin(), local_rows.end(), local_weights.begin(), [&](int row) { + return global_weights[row]; + }); + + auto weight_buffer_size = std::max(std::size_t{1}, local_weights.size()); + rmm::device_uvector d_weights(weight_buffer_size, handle.get_stream()); + raft::update_device( + d_weights.data(), local_weights.data(), local_weights.size(), handle.get_stream()); + + RF_params rf_params{}; + rf_params.bootstrap = true; + rf_params.seed = 123456789ULL; + rf_params.n_streams = 1; + + constexpr std::int64_t sample_count = 200000; + detail::RowSampler sampler(handle, + rf_params, + static_cast(local_rows.size()), + sample_count, + 1, + nullptr, + d_weights.data()); + + auto stream = handle.get_stream_from_stream_pool(0); + auto& row_ids = sampler.sample(0, 0, stream.get()); + std::vector h_row_ids(row_ids.size()); + raft::update_host(h_row_ids.data(), row_ids.data(), row_ids.size(), stream); + handle.sync_stream(stream); + + std::vector local_counts(global_weights.size(), 0); + std::uint64_t local_invalid_row_count = 0; + for (auto local_row : h_row_ids) { + if (local_row < 0 || local_row >= static_cast(local_rows.size())) { + local_invalid_row_count++; + continue; + } + local_counts[local_rows[local_row]]++; + } + + std::uint64_t global_invalid_row_count = 0; + MPI_Allreduce( + &local_invalid_row_count, &global_invalid_row_count, 1, MPI_UINT64_T, MPI_SUM, MPI_COMM_WORLD); + std::vector global_counts(global_weights.size(), 0); + MPI_Allreduce(local_counts.data(), + global_counts.data(), + static_cast(global_counts.size()), + MPI_UINT64_T, + MPI_SUM, + MPI_COMM_WORLD); + + ASSERT_EQ(global_invalid_row_count, 0); + auto observed_total = + std::accumulate(global_counts.begin(), global_counts.end(), std::uint64_t{0}); + ASSERT_EQ(observed_total, static_cast(sample_count)); + + double weight_sum = std::accumulate(global_weights.begin(), global_weights.end(), 0.0); + for (std::size_t row = 0; row < global_weights.size(); ++row) { + double expected_probability = global_weights[row] / weight_sum; + double observed_probability = static_cast(global_counts[row]) / sample_count; + if (expected_probability == 0.0) { + EXPECT_EQ(global_counts[row], 0) << "global row " << row; + continue; + } + double standard_error = + std::sqrt(expected_probability * (1.0 - expected_probability) / sample_count); + double tolerance = std::max(0.002, 6.0 * standard_error); + EXPECT_NEAR(observed_probability, expected_probability, tolerance) << "global row " << row; + } +} + +std::string row_sampler_test_name(::testing::TestParamInfo const& test_info) +{ + char const* partition_name = nullptr; + switch (test_info.param.partition_kind) { + case PartitionKind::Contiguous: partition_name = "Contiguous"; break; + case PartitionKind::Strided: partition_name = "Strided"; break; + case PartitionKind::Imbalanced: partition_name = "Imbalanced"; break; + case PartitionKind::EmptyNonRootRanks: partition_name = "EmptyNonRootRanks"; break; + } + char const* weight_name = nullptr; + switch (test_info.param.weight_kind) { + case WeightKind::ThreeRows: weight_name = "ThreeRows"; break; + case WeightKind::Uniform: weight_name = "Uniform"; break; + case WeightKind::WithZeros: weight_name = "WithZeros"; break; + case WeightKind::Skewed: weight_name = "Skewed"; break; + } + return std::string{partition_name} + weight_name; +} + +std::vector make_row_sampler_inputs() +{ + std::vector inputs; + for (auto partition : {PartitionKind::Contiguous, + PartitionKind::Strided, + PartitionKind::Imbalanced, + PartitionKind::EmptyNonRootRanks}) { + for (auto weights : + {WeightKind::ThreeRows, WeightKind::Uniform, WeightKind::WithZeros, WeightKind::Skewed}) { + inputs.push_back({partition, weights}); + } + } + return inputs; +} + +INSTANTIATE_TEST_SUITE_P(RfRowSamplerTests, + RfMgRowSamplerTest, + ::testing::ValuesIn(make_row_sampler_inputs()), + row_sampler_test_name); + +} // namespace opg +} // namespace Test +} // namespace ML + +int main(int argc, char** argv) +{ + ::testing::InitGoogleTest(&argc, argv); + ::testing::AddGlobalTestEnvironment(new MLCommon::Test::opg::MPIEnvironment()); + return RUN_ALL_TESTS(); +} diff --git a/cpp/tests/mg/rf_test.cu b/cpp/tests/mg/rf_test.cu index 8adc1d8f44..e7f67be2d5 100644 --- a/cpp/tests/mg/rf_test.cu +++ b/cpp/tests/mg/rf_test.cu @@ -4,6 +4,7 @@ */ #include "../prims/test_utils.h" +#include "rf_test_utils.hpp" #include "test_opg_utils.h" #include @@ -33,8 +34,6 @@ namespace ML { namespace Test { namespace opg { -enum class PartitionKind { Contiguous, Strided, Imbalanced, EmptyNonRootRanks }; - struct RfMgTestParams { int n_rows; int n_cols; @@ -134,43 +133,6 @@ void expect_forests_equal(RandomForestMetaData const& distributed_forest, } } -std::vector local_rows_for_rank(int n_rows, int rank, int size, PartitionKind kind) -{ - std::vector rows; - if (kind == PartitionKind::Strided) { - for (int row = rank; row < n_rows; row += size) { - rows.push_back(row); - } - return rows; - } - - std::vector counts(size, n_rows / size); - for (int i = 0; i < n_rows % size; ++i) { - counts[i]++; - } - if (kind == PartitionKind::Imbalanced && size > 1) { - counts.assign(size, 0); - counts[0] = std::max(1, (n_rows * 3) / 4); - int remaining = n_rows - counts[0]; - for (int i = 1; i < size; ++i) { - counts[i] = remaining / (size - 1); - } - for (int i = 1; i <= remaining % (size - 1); ++i) { - counts[i]++; - } - } else if (kind == PartitionKind::EmptyNonRootRanks && size > 1) { - counts.assign(size, 0); - counts[0] = n_rows; - } - - int begin = std::accumulate(counts.begin(), counts.begin() + rank, 0); - rows.resize(counts[rank]); - std::iota(rows.begin(), rows.end(), begin); - rows.erase(std::remove_if(rows.begin(), rows.end(), [=](int row) { return row >= n_rows; }), - rows.end()); - return rows; -} - template void make_local_dataset(RfMgTestParams const& params, std::vector const& rows, diff --git a/cpp/tests/mg/rf_test_utils.hpp b/cpp/tests/mg/rf_test_utils.hpp new file mode 100644 index 0000000000..cf351ced1f --- /dev/null +++ b/cpp/tests/mg/rf_test_utils.hpp @@ -0,0 +1,57 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION & AFFILIATES. All rights reserved. + * SPDX-License-Identifier: Apache-2.0 + */ + +#pragma once + +#include +#include +#include + +namespace ML { +namespace Test { +namespace opg { + +enum class PartitionKind { Contiguous, Strided, Imbalanced, EmptyNonRootRanks }; + +inline std::vector local_rows_for_rank(int n_rows, int rank, int size, PartitionKind kind) +{ + std::vector rows; + if (kind == PartitionKind::Strided) { + for (int row = rank; row < n_rows; row += size) { + rows.push_back(row); + } + return rows; + } + + std::vector counts(size, n_rows / size); + for (int i = 0; i < n_rows % size; ++i) { + counts[i]++; + } + if (kind == PartitionKind::Imbalanced && size > 1) { + counts.assign(size, 0); + counts[0] = std::max(1, (n_rows * 3) / 4); + int remaining = n_rows - counts[0]; + for (int i = 1; i < size; ++i) { + counts[i] = remaining / (size - 1); + } + for (int i = 1; i <= remaining % (size - 1); ++i) { + counts[i]++; + } + } else if (kind == PartitionKind::EmptyNonRootRanks && size > 1) { + counts.assign(size, 0); + counts[0] = n_rows; + } + + int begin = std::accumulate(counts.begin(), counts.begin() + rank, 0); + rows.resize(counts[rank]); + std::iota(rows.begin(), rows.end(), begin); + rows.erase(std::remove_if(rows.begin(), rows.end(), [=](int row) { return row >= n_rows; }), + rows.end()); + return rows; +} + +} // namespace opg +} // namespace Test +} // namespace ML From f17b660294d5d36e12bb74c1808c66ac34c5e96d Mon Sep 17 00:00:00 2001 From: Hyunsu Cho Date: Wed, 30 Sep 2026 00:32:29 +0000 Subject: [PATCH 4/4] Reduce the number of sampled rows by 1 / comm_size_ --- cpp/src/randomforest/randomforest.cuh | 4 +++- 1 file changed, 3 insertions(+), 1 deletion(-) diff --git a/cpp/src/randomforest/randomforest.cuh b/cpp/src/randomforest/randomforest.cuh index 2a2a588e64..d055c077bd 100644 --- a/cpp/src/randomforest/randomforest.cuh +++ b/cpp/src/randomforest/randomforest.cuh @@ -493,6 +493,7 @@ class RandomForest { const raft::handle_t& handle = user_handle; bool distributed = raft::resource::comms_initialized(handle) && handle.get_comms().get_size() > 1; + auto const comm_size = distributed ? handle.get_comms().get_size() : 1; this->error_checking(input, labels, n_rows, n_cols, false, distributed); std::int64_t const n_rows_i64 = n_rows; std::int64_t global_n_rows = n_rows_i64; @@ -512,8 +513,9 @@ class RandomForest { } std::int64_t n_sampled_rows = 0; if (this->rf_params.bootstrap) { - n_sampled_rows = + auto const global_n_sampled_rows = static_cast(std::round(this->rf_params.max_samples * global_n_rows)); + n_sampled_rows = raft::ceildiv(global_n_sampled_rows, static_cast(comm_size)); } else { if (this->rf_params.max_samples != 1.0) { CUML_LOG_WARN(