diff --git a/cpp/src/decisiontree/batched-levelalgo/kernels/builder_kernels.cuh b/cpp/src/decisiontree/batched-levelalgo/kernels/builder_kernels.cuh index 3e25a312f3..6f56228d35 100644 --- a/cpp/src/decisiontree/batched-levelalgo/kernels/builder_kernels.cuh +++ b/cpp/src/decisiontree/batched-levelalgo/kernels/builder_kernels.cuh @@ -8,6 +8,7 @@ #include "../bins.cuh" #include "../objectives.cuh" #include "../quantiles.h" +#include "../random_utils.cuh" #include @@ -83,24 +84,9 @@ void launchLeafKernel(ObjectiveT objective, int batch_size, size_t smem_size, cudaStream_t builder_stream); -// 32-bit FNV1a hash -// Reference: http://www.isthe.com/chongo/tech/comp/fnv/index.html -const uint32_t fnv1a32_prime = uint32_t(16777619); -const uint32_t fnv1a32_basis = uint32_t(2166136261); -HDI uint32_t fnv1a32(uint32_t hash, uint32_t txt) -{ - hash ^= (txt >> 0) & 0xFF; - hash *= fnv1a32_prime; - hash ^= (txt >> 8) & 0xFF; - hash *= fnv1a32_prime; - hash ^= (txt >> 16) & 0xFF; - hash *= fnv1a32_prime; - hash ^= (txt >> 24) & 0xFF; - hash *= fnv1a32_prime; - return hash; -} - -// returns the lowest index in `array` whose value is greater or equal to `element` +// Returns the lowest index in `array` whose value is greater or equal to `element`. +// Values outside the quantile range are clamped to the edge bins: values below the +// first quantile return 0, and values above the last quantile return len - 1. template HDI IdxT lower_bound(DataT* array, IdxT len, DataT element) { diff --git a/cpp/src/decisiontree/batched-levelalgo/quantiles.cuh b/cpp/src/decisiontree/batched-levelalgo/quantiles.cuh index 493d164637..362c4004be 100644 --- a/cpp/src/decisiontree/batched-levelalgo/quantiles.cuh +++ b/cpp/src/decisiontree/batched-levelalgo/quantiles.cuh @@ -6,11 +6,14 @@ #pragma once #include "quantiles.h" +#include "random_utils.cuh" #include +#include #include #include +#include #include #include @@ -20,12 +23,42 @@ #include #include +#include #include #include namespace ML { namespace DT { +namespace detail { + +template +static __global__ void gatherUniformSampledColumnKernel( + T* out, const T* data, int sample_count, int n_rows, int col, uint64_t seed) +{ + int tid = blockIdx.x * blockDim.x + threadIdx.x; + auto col_seed = fnv1a32_basis; + col_seed = fnv1a32(col_seed, static_cast(seed)); + col_seed = fnv1a32(col_seed, static_cast(seed >> 32)); + col_seed = fnv1a32(col_seed, static_cast(col)); + // Sampling is with replacement. Duplicate values from sample collisions are + // removed later when quantile candidates are compacted with thrust::unique. + for (int sample_idx = tid; sample_idx < sample_count; sample_idx += blockDim.x * gridDim.x) { + // Use sample_idx as the generator subsequence so each output position is + // deterministic and independent of the CUDA block/thread layout. + raft::random::PCGenerator gen(col_seed, static_cast(sample_idx), uint64_t(0)); + raft::random::UniformIntDistParams uniform_int_dist_params; + uniform_int_dist_params.start = 0; + uniform_int_dist_params.end = n_rows; + uniform_int_dist_params.diff = static_cast(n_rows); + int row; + raft::random::custom_next(gen, &row, uniform_int_dist_params, int(0), int(0)); + out[sample_idx] = data[static_cast(col) * n_rows + row]; + } +} + +} // namespace detail + template static __global__ void computeQuantilesKernel( T* quantiles, int* n_bins, const T* sorted_data, const int max_n_bins, const int n_rows) @@ -58,57 +91,113 @@ using QuantileReturnValue = std::tuple, std::shared_ptr>, std::shared_ptr>>; +/** + * @brief Compute per-feature quantile split candidates from uniformly sampled rows. + * + * Each feature column is sampled independently with replacement using a deterministic + * seed derived from `seed`, the feature index, and the output sample index. When the + * requested sample budget is at least the local row count, the full column is used. + * + * @tparam T Floating-point input type. + * @param handle RAFT handle used for stream and resource access. + * @param data Column-major input matrix with shape `[n_cols, n_rows]`. + * @param max_n_bins Maximum number of quantile candidates to retain per feature. + * @param n_rows Number of rows in `data`. + * @param n_cols Number of columns in `data`. + * @param oversampling_factor Multiplier applied to `max_n_bins` to choose the + * sampled row budget per feature before sorting and quantile extraction. The + * default of 4 is a conservative choice while still bounding memory; for fixed + * `max_n_bins`, rank error decreases like O(1 / sqrt(oversampling_factor)), so + * returns from increasing this are strongly diminishing. + * @param seed User seed for deterministic sampling. + * @return Quantile metadata and owning device buffers for quantile values and bin counts. + */ template -CUML_EXPORT QuantileReturnValue computeQuantiles( - const raft::handle_t& handle, const T* data, int max_n_bins, int n_rows, int n_cols) +CUML_EXPORT QuantileReturnValue computeQuantiles(const raft::handle_t& handle, + const T* data, + int max_n_bins, + int n_rows, + int n_cols, + int oversampling_factor = 4, + uint64_t seed = uint64_t{0}) { raft::common::nvtx::push_range("computeQuantiles"); - auto stream = handle.get_stream(); - size_t temp_storage_bytes = 0; // for device radix sort - rmm::device_uvector sorted_column(n_rows, stream); - // acquire device vectors to store the quantiles + offsets + RAFT_EXPECTS(data != nullptr, "data pointer must not be null"); + RAFT_EXPECTS(max_n_bins > 0, "max_n_bins must be positive"); + RAFT_EXPECTS(n_rows > 0, "n_rows must be positive"); + RAFT_EXPECTS(n_cols > 0, "n_cols must be positive"); + RAFT_EXPECTS(oversampling_factor > 0, "oversampling_factor must be positive"); + + auto stream = handle.get_stream(); + int64_t size = static_cast(max_n_bins) * oversampling_factor; + int sample_count = + static_cast(std::min(static_cast(n_rows), std::max(1, size))); + + rmm::device_uvector sampled_column(sample_count, stream); + rmm::device_uvector sorted_sample(sample_count, stream); auto quantiles_array = std::make_shared>(n_cols * max_n_bins, stream); auto n_bins_array = std::make_shared>(n_cols, stream); - // get temp_storage_bytes for sorting - RAFT_CUDA_TRY(cub::DeviceRadixSort::SortKeys( - nullptr, temp_storage_bytes, data, sorted_column.data(), n_rows, 0, 8 * sizeof(T), stream)); - // allocate total memory needed for parallelized sorting + size_t temp_storage_bytes = 0; + RAFT_CUDA_TRY(cub::DeviceRadixSort::SortKeys(nullptr, + temp_storage_bytes, + sampled_column.data(), + sorted_sample.data(), + sample_count, + 0, + 8 * sizeof(T), + stream)); rmm::device_uvector d_temp_storage(temp_storage_bytes, stream); + + int n_threads = 256; + int n_blocks = raft::ceildiv(sample_count, n_threads); + n_blocks = std::min(n_blocks, 1024); + for (int col = 0; col < n_cols; col++) { - raft::common::nvtx::push_range("sorting columns"); - int col_offset = col * n_rows; + raft::common::nvtx::push_range("sample quantile column"); + if (sample_count == n_rows) { + RAFT_CUDA_TRY(cudaMemcpyAsync(sampled_column.data(), + data + static_cast(col) * n_rows, + sizeof(T) * n_rows, + cudaMemcpyDeviceToDevice, + stream)); + } else { + detail::gatherUniformSampledColumnKernel<<>>( + sampled_column.data(), data, sample_count, n_rows, col, seed); + RAFT_CUDA_TRY(cudaGetLastError()); + } + raft::common::nvtx::pop_range(); + + raft::common::nvtx::push_range("sort sampled quantile column"); RAFT_CUDA_TRY(cub::DeviceRadixSort::SortKeys((void*)(d_temp_storage.data()), temp_storage_bytes, - data + col_offset, - sorted_column.data(), - n_rows, + sampled_column.data(), + sorted_sample.data(), + sample_count, 0, 8 * sizeof(T), stream)); - RAFT_CUDA_TRY(cudaStreamSynchronize(stream)); - raft::common::nvtx::pop_range(); // sorting columns + raft::common::nvtx::pop_range(); - int n_blocks = 1; - int n_threads = min(1024, max_n_bins); int quantile_offset = col * max_n_bins; int bins_offset = col; - raft::common::nvtx::push_range("computeQuantilesKernel @quantile.cuh"); - computeQuantilesKernel<<>>( + raft::common::nvtx::push_range("computeQuantilesKernel @quantiles.cuh"); + computeQuantilesKernel<<<1, std::min(1024, max_n_bins), 0, stream>>>( quantiles_array->data() + quantile_offset, n_bins_array->data() + bins_offset, - sorted_column.data(), + sorted_sample.data(), max_n_bins, - n_rows); - RAFT_CUDA_TRY(cudaStreamSynchronize(handle.get_stream())); + sample_count); RAFT_CUDA_TRY(cudaGetLastError()); - raft::common::nvtx::pop_range(); // computeQuatilesKernel + raft::common::nvtx::pop_range(); } - // encapsulate the device pointers under a Quantiles struct + + handle.sync_stream(stream); + Quantiles quantiles; quantiles.quantiles_array = quantiles_array->data(); quantiles.n_bins_array = n_bins_array->data(); - raft::common::nvtx::pop_range(); // computeQuantiles + raft::common::nvtx::pop_range(); return std::make_tuple(quantiles, quantiles_array, n_bins_array); } diff --git a/cpp/src/decisiontree/batched-levelalgo/random_utils.cuh b/cpp/src/decisiontree/batched-levelalgo/random_utils.cuh new file mode 100644 index 0000000000..1e5a7e3a99 --- /dev/null +++ b/cpp/src/decisiontree/batched-levelalgo/random_utils.cuh @@ -0,0 +1,34 @@ +/* + * SPDX-FileCopyrightText: Copyright (c) 2026, NVIDIA CORPORATION. + * SPDX-License-Identifier: Apache-2.0 + */ + +#pragma once + +#include + +#include + +namespace ML { +namespace DT { + +// 32-bit FNV1a hash +// Reference: http://www.isthe.com/chongo/tech/comp/fnv/index.html +constexpr uint32_t fnv1a32_prime = uint32_t(16777619); +constexpr uint32_t fnv1a32_basis = uint32_t(2166136261); + +HDI uint32_t fnv1a32(uint32_t hash, uint32_t txt) +{ + hash ^= (txt >> 0) & 0xFF; + hash *= fnv1a32_prime; + hash ^= (txt >> 8) & 0xFF; + hash *= fnv1a32_prime; + hash ^= (txt >> 16) & 0xFF; + hash *= fnv1a32_prime; + hash ^= (txt >> 24) & 0xFF; + hash *= fnv1a32_prime; + return hash; +} + +} // namespace DT +} // namespace ML diff --git a/cpp/src/randomforest/randomforest.cuh b/cpp/src/randomforest/randomforest.cuh index 356bddef4f..f0115402aa 100644 --- a/cpp/src/randomforest/randomforest.cuh +++ b/cpp/src/randomforest/randomforest.cuh @@ -142,8 +142,8 @@ class RandomForest { // computing the quantiles: last two return values are shared pointers to device memory // encapsulated by quantiles struct - auto [quantiles, quantiles_array, n_bins_array] = - DT::computeQuantiles(handle, input, this->rf_params.tree_params.max_n_bins, n_rows, n_cols); + auto [quantiles, quantiles_array, n_bins_array] = DT::computeQuantiles( + handle, input, this->rf_params.tree_params.max_n_bins, n_rows, n_cols, 4, rf_params.seed); // n_streams should not be less than n_trees if (this->rf_params.n_trees < n_streams) n_streams = this->rf_params.n_trees; diff --git a/cpp/tests/sg/rf_test.cu b/cpp/tests/sg/rf_test.cu index 8e769d2f14..3484ad9bb8 100644 --- a/cpp/tests/sg/rf_test.cu +++ b/cpp/tests/sg/rf_test.cu @@ -27,6 +27,7 @@ #include #include #include +#include #include #include @@ -40,8 +41,10 @@ #include #include +#include #include #include +#include #include #include #include @@ -49,26 +52,6 @@ namespace ML { -namespace DT { - -template -using ReturnValue = std::tuple, - std::shared_ptr>, - std::shared_ptr>>; - -template -ReturnValue computeQuantiles( - const raft::handle_t& handle, const T* data, int max_n_bins, int n_rows, int n_cols); - -template <> -ReturnValue computeQuantiles( - const raft::handle_t& handle, const float* data, int max_n_bins, int n_rows, int n_cols); - -template <> -ReturnValue computeQuantiles( - const raft::handle_t& handle, const double* data, int max_n_bins, int n_rows, int n_cols); -} // namespace DT - // Utils for changing tuple into struct namespace detail { template @@ -751,6 +734,12 @@ class RFQuantileBinsLowerBoundTest : public ::testing::TestWithParam { auto params = ::testing::TestWithParam::GetParam(); thrust::device_vector data(params.n_rows); - thrust::device_vector histogram(params.max_n_bins); - thrust::host_vector h_histogram(params.max_n_bins); raft::random::Rng r(8); r.normal(data.data().get(), data.size(), T(0.0), T(2.0), nullptr); @@ -779,35 +766,16 @@ class RFQuantileTest : public ::testing::TestWithParam { int n_unique_bins; raft::copy(&n_unique_bins, quantiles.n_bins_array, 1, handle.get_stream()); - if (n_unique_bins < params.max_n_bins) { - return; // almost impossible that this happens, skip if so - } + if (n_unique_bins < params.max_n_bins) { ASSERT_GT(n_unique_bins, 1); } + ASSERT_LE(n_unique_bins, params.max_n_bins); - auto d_quantiles = quantiles.quantiles_array; - auto d_histogram = histogram.data().get(); - - thrust::for_each(data.begin(), data.end(), [=] __device__(T x) { - for (int j = 0; j < params.max_n_bins; j++) { - if (x <= d_quantiles[j]) { - atomicAdd(&d_histogram[j], 1); - break; - } - } - }); - - h_histogram = histogram; - int max_items_per_bin = raft::ceildiv(params.n_rows, params.max_n_bins); - int min_items_per_bin = max_items_per_bin - 1; - int total_items = 0; - for (int b = 0; b < params.max_n_bins; b++) { - ASSERT_TRUE(h_histogram[b] == max_items_per_bin or h_histogram[b] == min_items_per_bin) - << "No. samples in bin[" << b << "] = " << h_histogram[b] << " Expected " - << max_items_per_bin << " or " << min_items_per_bin << std::endl; - total_items += h_histogram[b]; + thrust::host_vector h_quantiles(params.max_n_bins); + raft::update_host( + h_quantiles.data(), quantiles.quantiles_array, params.max_n_bins, handle.get_stream()); + handle.sync_stream(); + for (int b = 1; b < n_unique_bins; b++) { + ASSERT_LT(h_quantiles[b - 1], h_quantiles[b]); } - ASSERT_EQ(params.n_rows, total_items) - << "Some samples from dataset are either missed of double counted in quantile bins" - << std::endl; } }; @@ -838,9 +806,9 @@ class RFQuantileVariableBinsTest : public ::testing::TestWithParamdata(), 1, handle.get_stream()); @@ -882,12 +850,137 @@ class RFQuantileVariableBinsTest : public ::testing::TestWithParam +class RFSampledQuantileExactFallbackTest : public ::testing::TestWithParam { + public: + void SetUp() override + { + auto params = ::testing::TestWithParam::GetParam(); + + auto stream_pool = std::make_shared(1); + raft::handle_t handle(rmm::cuda_stream_per_thread, stream_pool); + thrust::device_vector data(params.n_rows); + thrust::sequence(data.begin(), data.end(), T(0)); + + auto [sampled_quantiles, sampled_quantiles_array, sampled_n_bins_array] = DT::computeQuantiles( + handle, data.data().get(), params.max_n_bins, params.n_rows, 1, params.n_rows, params.seed); + + int sampled_n_bins; + raft::copy(&sampled_n_bins, sampled_n_bins_array->data(), 1, handle.get_stream()); + handle.sync_stream(); + + ASSERT_EQ(sampled_n_bins, params.max_n_bins); + + thrust::host_vector h_sampled(params.max_n_bins); + raft::update_host( + h_sampled.data(), sampled_quantiles.quantiles_array, params.max_n_bins, handle.get_stream()); + handle.sync_stream(); + + double bin_width = static_cast(params.n_rows) / params.max_n_bins; + for (int bin = 0; bin < sampled_n_bins; ++bin) { + int idx = int(round((bin + 1) * bin_width)) - 1; + idx = std::min(std::max(0, idx), params.n_rows - 1); + ASSERT_EQ(h_sampled[bin], T(idx)); + } + } +}; + +template +class RFSampledQuantileDeterminismTest : public ::testing::TestWithParam { + public: + void SetUp() override + { + auto params = ::testing::TestWithParam::GetParam(); + + auto stream_pool = std::make_shared(1); + raft::handle_t handle(rmm::cuda_stream_per_thread, stream_pool); + thrust::device_vector data(params.n_rows); + raft::random::Rng r(params.seed); + r.normal(data.data().get(), data.size(), T(0.0), T(2.0), nullptr); + + auto [quantiles_a, quantiles_array_a, n_bins_array_a] = DT::computeQuantiles( + handle, data.data().get(), params.max_n_bins, params.n_rows, 1, 4, params.seed); + auto [quantiles_b, quantiles_array_b, n_bins_array_b] = DT::computeQuantiles( + handle, data.data().get(), params.max_n_bins, params.n_rows, 1, 4, params.seed); + + int n_bins_a; + int n_bins_b; + raft::copy(&n_bins_a, n_bins_array_a->data(), 1, handle.get_stream()); + raft::copy(&n_bins_b, n_bins_array_b->data(), 1, handle.get_stream()); + handle.sync_stream(); + + ASSERT_EQ(n_bins_a, n_bins_b); + ASSERT_GT(n_bins_a, 1); + ASSERT_LE(n_bins_a, params.max_n_bins); + + thrust::host_vector h_quantiles_a(params.max_n_bins); + thrust::host_vector h_quantiles_b(params.max_n_bins); + raft::update_host( + h_quantiles_a.data(), quantiles_a.quantiles_array, params.max_n_bins, handle.get_stream()); + raft::update_host( + h_quantiles_b.data(), quantiles_b.quantiles_array, params.max_n_bins, handle.get_stream()); + handle.sync_stream(); + + for (int i = 0; i < n_bins_a; ++i) { + ASSERT_EQ(h_quantiles_a[i], h_quantiles_b[i]); + if (i > 0) { ASSERT_LT(h_quantiles_a[i - 1], h_quantiles_a[i]); } + } + } +}; + +template +class RFSampledQuantileRankErrorTest : public ::testing::TestWithParam { + public: + void SetUp() override + { + auto params = ::testing::TestWithParam::GetParam(); + + auto stream_pool = std::make_shared(1); + raft::handle_t handle(rmm::cuda_stream_per_thread, stream_pool); + thrust::device_vector data(params.n_rows); + thrust::sequence(data.begin(), data.end(), T(0)); + + auto [quantiles, quantiles_array, n_bins_array] = DT::computeQuantiles( + handle, data.data().get(), params.max_n_bins, params.n_rows, 1, 4, params.seed); + + int n_bins; + raft::copy(&n_bins, n_bins_array->data(), 1, handle.get_stream()); + handle.sync_stream(); + + ASSERT_EQ(n_bins, params.max_n_bins); + + thrust::host_vector h_quantiles(params.max_n_bins); + raft::update_host( + h_quantiles.data(), quantiles.quantiles_array, params.max_n_bins, handle.get_stream()); + handle.sync_stream(); + + double total_abs_rank_error = 0.0; + double max_abs_rank_error = 0.0; + for (int bin = 0; bin < n_bins; ++bin) { + double expected_rank = static_cast(bin + 1) / params.max_n_bins; + double actual_rank = (static_cast(h_quantiles[bin]) + 1.0) / params.n_rows; + double rank_error = std::abs(actual_rank - expected_rank); + total_abs_rank_error += rank_error; + max_abs_rank_error = std::max(max_abs_rank_error, rank_error); + } + + double mean_abs_rank_error = total_abs_rank_error / n_bins; + double sample_count = static_cast(params.max_n_bins * 4); + double rank_error_scale = 1.0 / std::sqrt(sample_count); + EXPECT_LT(mean_abs_rank_error, rank_error_scale); + EXPECT_LT(max_abs_rank_error, 3.0 * rank_error_scale); + } +}; + const std::vector inputs = {{1000, 16, 6078587519764079670LLU}, {1130, 32, 4884670006177930266LLU}, {1752, 67, 9175325892580481371LLU}, {2307, 99, 9507819643927052255LLU}, {5000, 128, 9507819643927052255LLU}}; +const std::vector rank_error_inputs = { + {10000, 128, 9507819643927052255LLU}}; + // float type quantile test typedef RFQuantileTest RFQuantileTestF; TEST_P(RFQuantileTestF, test) {} @@ -918,6 +1011,34 @@ typedef RFQuantileVariableBinsTest RFQuantileVariableBinsTestD; TEST_P(RFQuantileVariableBinsTestD, test) {} INSTANTIATE_TEST_CASE_P(RfTests, RFQuantileVariableBinsTestD, ::testing::ValuesIn(inputs)); +typedef RFSampledQuantileExactFallbackTest RFSampledQuantileExactFallbackTestF; +TEST_P(RFSampledQuantileExactFallbackTestF, test) {} +INSTANTIATE_TEST_CASE_P(RfTests, RFSampledQuantileExactFallbackTestF, ::testing::ValuesIn(inputs)); + +typedef RFSampledQuantileExactFallbackTest RFSampledQuantileExactFallbackTestD; +TEST_P(RFSampledQuantileExactFallbackTestD, test) {} +INSTANTIATE_TEST_CASE_P(RfTests, RFSampledQuantileExactFallbackTestD, ::testing::ValuesIn(inputs)); + +typedef RFSampledQuantileDeterminismTest RFSampledQuantileDeterminismTestF; +TEST_P(RFSampledQuantileDeterminismTestF, test) {} +INSTANTIATE_TEST_CASE_P(RfTests, RFSampledQuantileDeterminismTestF, ::testing::ValuesIn(inputs)); + +typedef RFSampledQuantileDeterminismTest RFSampledQuantileDeterminismTestD; +TEST_P(RFSampledQuantileDeterminismTestD, test) {} +INSTANTIATE_TEST_CASE_P(RfTests, RFSampledQuantileDeterminismTestD, ::testing::ValuesIn(inputs)); + +typedef RFSampledQuantileRankErrorTest RFSampledQuantileRankErrorTestF; +TEST_P(RFSampledQuantileRankErrorTestF, test) {} +INSTANTIATE_TEST_CASE_P(RfTests, + RFSampledQuantileRankErrorTestF, + ::testing::ValuesIn(rank_error_inputs)); + +typedef RFSampledQuantileRankErrorTest RFSampledQuantileRankErrorTestD; +TEST_P(RFSampledQuantileRankErrorTestD, test) {} +INSTANTIATE_TEST_CASE_P(RfTests, + RFSampledQuantileRankErrorTestD, + ::testing::ValuesIn(rank_error_inputs)); + //------------------------------------------------------------------------------------------------------ TEST(RfTest, TextDump)