From 71454fd89adeee5c4d3d77548f2494ea77907421 Mon Sep 17 00:00:00 2001 From: Rory Mitchell Date: Wed, 10 Jun 2026 03:09:54 -0700 Subject: [PATCH 1/3] Add random forest bin interfaces for weights support --- .../decisiontree/batched-levelalgo/bins.cuh | 141 +++++++++++++++--- .../kernels/classification-double.cu | 2 +- .../kernels/classification-float.cu | 2 +- .../kernels/regression-double.cu | 2 +- .../kernels/regression-float.cu | 2 +- .../batched-levelalgo/objectives.cuh | 140 ++++++++++------- cpp/tests/sg/rf_test.cu | 12 +- 7 files changed, 210 insertions(+), 91 deletions(-) diff --git a/cpp/src/decisiontree/batched-levelalgo/bins.cuh b/cpp/src/decisiontree/batched-levelalgo/bins.cuh index 6125cfda6e..2e370d7073 100644 --- a/cpp/src/decisiontree/batched-levelalgo/bins.cuh +++ b/cpp/src/decisiontree/batched-levelalgo/bins.cuh @@ -3,61 +3,156 @@ * SPDX-License-Identifier: Apache-2.0 */ #pragma once + #include namespace ML { namespace DT { -struct CountBin { - // double covers both the unweighted count path and the future weighted-count - // path with one bin type; 32-bit int would overflow on large weighted counts. - double x; - CountBin(CountBin const&) = default; - HDI CountBin(double x_) : x(x_) {} - HDI CountBin() : x(0.0) {} +using BinCountT = unsigned long long int; +static_assert(sizeof(BinCountT) == 8, "BinCountT must be 64 bits"); + +struct ClassificationBin { + BinCountT count; - DI static void IncrementHistogram(CountBin* hist, int n_bins, int b, int label) + ClassificationBin(ClassificationBin const&) = default; + HDI ClassificationBin(BinCountT count_) : count(count_) {} + HDI ClassificationBin() : count(0) {} + + DI static void IncrementHistogram(ClassificationBin* hist, int n_bins, int b, int label) { auto offset = label * n_bins + b; - CountBin::AtomicAdd(hist + offset, {1.0}); + ClassificationBin::AtomicAdd(hist + offset, {1}); } - DI static void AtomicAdd(CountBin* address, CountBin val) { atomicAdd(&address->x, val.x); } - HDI CountBin& operator+=(const CountBin& b) + DI static void AtomicAdd(ClassificationBin* address, ClassificationBin val) { - x += b.x; + atomicAdd(&address->count, val.count); + } + HDI BinCountT Count() const { return count; } + HDI double Weight() const { return static_cast(count); } + HDI ClassificationBin& operator+=(const ClassificationBin& b) + { + count += b.count; return *this; } - HDI CountBin operator+(CountBin b) const + HDI ClassificationBin operator+(ClassificationBin b) const { b += *this; return b; } }; -struct AggregateBin { +struct WeightedClassificationBin { + BinCountT count; + double weight; + + WeightedClassificationBin(WeightedClassificationBin const&) = default; + HDI WeightedClassificationBin(BinCountT count_, double weight_) : count(count_), weight(weight_) + { + } + HDI WeightedClassificationBin() : count(0), weight(0.0) {} + + DI static void IncrementHistogram(WeightedClassificationBin* hist, int n_bins, int b, int label) + { + WeightedClassificationBin::IncrementHistogram(hist, n_bins, b, label, 1.0); + } + DI static void IncrementHistogram( + WeightedClassificationBin* hist, int n_bins, int b, int label, double weight) + { + auto offset = label * n_bins + b; + WeightedClassificationBin::AtomicAdd(hist + offset, {1, weight}); + } + DI static void AtomicAdd(WeightedClassificationBin* address, WeightedClassificationBin val) + { + atomicAdd(&address->count, val.count); + atomicAdd(&address->weight, val.weight); + } + HDI BinCountT Count() const { return count; } + HDI double Weight() const { return weight; } + HDI WeightedClassificationBin& operator+=(const WeightedClassificationBin& b) + { + count += b.count; + weight += b.weight; + return *this; + } + HDI WeightedClassificationBin operator+(WeightedClassificationBin b) const + { + b += *this; + return b; + } +}; + +struct RegressionBin { + double label_sum; + BinCountT count; + + RegressionBin(RegressionBin const&) = default; + HDI RegressionBin() : label_sum(0.0), count(0) {} + HDI RegressionBin(double label_sum, BinCountT count) : label_sum(label_sum), count(count) {} + + DI static void IncrementHistogram(RegressionBin* hist, int n_bins, int b, double label) + { + RegressionBin::AtomicAdd(hist + b, {label, 1}); + } + DI static void AtomicAdd(RegressionBin* address, RegressionBin val) + { + atomicAdd(&address->label_sum, val.label_sum); + atomicAdd(&address->count, val.count); + } + HDI double LabelSum() const { return label_sum; } + HDI BinCountT Count() const { return count; } + HDI double Weight() const { return static_cast(count); } + HDI RegressionBin& operator+=(const RegressionBin& b) + { + label_sum += b.label_sum; + count += b.count; + return *this; + } + HDI RegressionBin operator+(RegressionBin b) const + { + b += *this; + return b; + } +}; + +struct WeightedRegressionBin { double label_sum; - int count; + BinCountT count; + double weight; - AggregateBin(AggregateBin const&) = default; - HDI AggregateBin() : label_sum(0.0), count(0) {} - HDI AggregateBin(double label_sum, int count) : label_sum(label_sum), count(count) {} + WeightedRegressionBin(WeightedRegressionBin const&) = default; + HDI WeightedRegressionBin() : label_sum(0.0), count(0), weight(0.0) {} + HDI WeightedRegressionBin(double label_sum, BinCountT count, double weight) + : label_sum(label_sum), count(count), weight(weight) + { + } - DI static void IncrementHistogram(AggregateBin* hist, int n_bins, int b, double label) + DI static void IncrementHistogram(WeightedRegressionBin* hist, int n_bins, int b, double label) + { + WeightedRegressionBin::IncrementHistogram(hist, n_bins, b, label, 1.0); + } + DI static void IncrementHistogram( + WeightedRegressionBin* hist, int n_bins, int b, double label, double weight) { - AggregateBin::AtomicAdd(hist + b, {label, 1}); + WeightedRegressionBin::AtomicAdd(hist + b, {label * weight, 1, weight}); } - DI static void AtomicAdd(AggregateBin* address, AggregateBin val) + DI static void AtomicAdd(WeightedRegressionBin* address, WeightedRegressionBin val) { atomicAdd(&address->label_sum, val.label_sum); atomicAdd(&address->count, val.count); + atomicAdd(&address->weight, val.weight); } - HDI AggregateBin& operator+=(const AggregateBin& b) + HDI double LabelSum() const { return label_sum; } + HDI BinCountT Count() const { return count; } + HDI double Weight() const { return weight; } + HDI WeightedRegressionBin& operator+=(const WeightedRegressionBin& b) { label_sum += b.label_sum; count += b.count; + weight += b.weight; return *this; } - HDI AggregateBin operator+(AggregateBin b) const + HDI WeightedRegressionBin operator+(WeightedRegressionBin b) const { b += *this; return b; diff --git a/cpp/src/decisiontree/batched-levelalgo/kernels/classification-double.cu b/cpp/src/decisiontree/batched-levelalgo/kernels/classification-double.cu index ed72a703e9..354576f3cd 100644 --- a/cpp/src/decisiontree/batched-levelalgo/kernels/classification-double.cu +++ b/cpp/src/decisiontree/batched-levelalgo/kernels/classification-double.cu @@ -14,7 +14,7 @@ using _DataT = double; using _LabelT = int; using _IdxT = int; using _ObjectiveT = ClassificationObjectiveFunction<_DataT, _LabelT, _IdxT>; -using _BinT = CountBin; +using _BinT = ClassificationBin; using _DatasetT = Dataset<_DataT, _LabelT, _IdxT>; using _NodeT = SparseTreeNode<_DataT, _LabelT, _IdxT>; } // namespace DT diff --git a/cpp/src/decisiontree/batched-levelalgo/kernels/classification-float.cu b/cpp/src/decisiontree/batched-levelalgo/kernels/classification-float.cu index a2257f2fd2..e199d214ca 100644 --- a/cpp/src/decisiontree/batched-levelalgo/kernels/classification-float.cu +++ b/cpp/src/decisiontree/batched-levelalgo/kernels/classification-float.cu @@ -14,7 +14,7 @@ using _DataT = float; using _LabelT = int; using _IdxT = int; using _ObjectiveT = ClassificationObjectiveFunction<_DataT, _LabelT, _IdxT>; -using _BinT = CountBin; +using _BinT = ClassificationBin; using _DatasetT = Dataset<_DataT, _LabelT, _IdxT>; using _NodeT = SparseTreeNode<_DataT, _LabelT, _IdxT>; } // namespace DT diff --git a/cpp/src/decisiontree/batched-levelalgo/kernels/regression-double.cu b/cpp/src/decisiontree/batched-levelalgo/kernels/regression-double.cu index 44bbeab917..d0f100ea46 100644 --- a/cpp/src/decisiontree/batched-levelalgo/kernels/regression-double.cu +++ b/cpp/src/decisiontree/batched-levelalgo/kernels/regression-double.cu @@ -14,7 +14,7 @@ using _DataT = double; using _LabelT = double; using _IdxT = int; using _ObjectiveT = RegressionObjectiveFunction<_DataT, _LabelT, _IdxT>; -using _BinT = AggregateBin; +using _BinT = RegressionBin; using _DatasetT = Dataset<_DataT, _LabelT, _IdxT>; using _NodeT = SparseTreeNode<_DataT, _LabelT, _IdxT>; } // namespace DT diff --git a/cpp/src/decisiontree/batched-levelalgo/kernels/regression-float.cu b/cpp/src/decisiontree/batched-levelalgo/kernels/regression-float.cu index 0008a39bca..f9453bad77 100644 --- a/cpp/src/decisiontree/batched-levelalgo/kernels/regression-float.cu +++ b/cpp/src/decisiontree/batched-levelalgo/kernels/regression-float.cu @@ -14,7 +14,7 @@ using _DataT = float; using _LabelT = float; using _IdxT = int; using _ObjectiveT = RegressionObjectiveFunction<_DataT, _LabelT, _IdxT>; -using _BinT = AggregateBin; +using _BinT = RegressionBin; using _DatasetT = Dataset<_DataT, _LabelT, _IdxT>; using _NodeT = SparseTreeNode<_DataT, _LabelT, _IdxT>; } // namespace DT diff --git a/cpp/src/decisiontree/batched-levelalgo/objectives.cuh b/cpp/src/decisiontree/batched-levelalgo/objectives.cuh index 3dbb1360b2..6e9ac16496 100644 --- a/cpp/src/decisiontree/batched-levelalgo/objectives.cuh +++ b/cpp/src/decisiontree/batched-levelalgo/objectives.cuh @@ -22,7 +22,7 @@ class ClassificationObjectiveFunction { using DataT = DataT_; using LabelT = LabelT_; using IdxT = IdxT_; - using BinT = CountBin; + using BinT = ClassificationBin; private: IdxT nclasses; @@ -31,29 +31,42 @@ class ClassificationObjectiveFunction { DI IdxT CountLeft(BinT const* hist, IdxT i, IdxT n_bins) const { - IdxT nLeft = 0; + BinCountT nLeft = 0; for (IdxT j = 0; j < nclasses; ++j) { - nLeft += hist[n_bins * j + i].x; + nLeft += hist[n_bins * j + i].Count(); } - return nLeft; + return static_cast(nLeft); } - HDI DataT GiniGain(BinT const* hist, IdxT i, IdxT n_bins, IdxT len, IdxT nLeft, IdxT nRight) const + HDI double WeightAt(BinT const* hist, IdxT i, IdxT n_bins) const + { + double weight = 0.0; + for (IdxT j = 0; j < nclasses; ++j) { + weight += hist[n_bins * j + i].Weight(); + } + return weight; + } + + HDI DataT GiniGain(BinT const* hist, IdxT i, IdxT n_bins, IdxT, IdxT, IdxT) const { constexpr DataT One = DataT(1.0); - auto invLen = One / len; - auto invLeft = One / nLeft; - auto invRight = One / nRight; - auto gain = DataT(0.0); + auto total_weight = WeightAt(hist, n_bins - 1, n_bins); + auto left_weight = WeightAt(hist, i, n_bins); + auto right_weight = total_weight - left_weight; + + auto invLen = One / DataT(total_weight); + auto invLeft = One / DataT(left_weight); + auto invRight = One / DataT(right_weight); + auto gain = DataT(0.0); for (IdxT j = 0; j < nclasses; ++j) { double val_i = 0.0; - auto lval_i = hist[n_bins * j + i].x; + auto lval_i = hist[n_bins * j + i].Weight(); auto lval = DataT(lval_i); gain += lval * invLeft * lval * invLen; val_i += lval_i; - auto total_sum = hist[n_bins * j + n_bins - 1].x; + auto total_sum = hist[n_bins * j + n_bins - 1].Weight(); auto rval_i = total_sum - lval_i; auto rval = DataT(rval_i); gain += rval * invRight * rval * invLen; @@ -66,23 +79,26 @@ class ClassificationObjectiveFunction { return gain; } - HDI DataT - EntropyGain(BinT const* hist, IdxT i, IdxT n_bins, IdxT len, IdxT nLeft, IdxT nRight) const + HDI DataT EntropyGain(BinT const* hist, IdxT i, IdxT n_bins, IdxT, IdxT, IdxT) const { + auto total_weight = WeightAt(hist, n_bins - 1, n_bins); + auto left_weight = WeightAt(hist, i, n_bins); + auto right_weight = total_weight - left_weight; + auto gain{DataT(0.0)}; - auto invLeft{DataT(1.0) / nLeft}; - auto invRight{DataT(1.0) / nRight}; - auto invLen{DataT(1.0) / len}; + auto invLeft{DataT(1.0) / DataT(left_weight)}; + auto invRight{DataT(1.0) / DataT(right_weight)}; + auto invLen{DataT(1.0) / DataT(total_weight)}; for (IdxT c = 0; c < nclasses; ++c) { double val_i = 0.0; - auto lval_i = hist[n_bins * c + i].x; + auto lval_i = hist[n_bins * c + i].Weight(); if (lval_i != 0) { auto lval = DataT(lval_i); gain += raft::log(lval * invLeft) / raft::log(DataT(2)) * lval * invLen; } val_i += lval_i; - auto total_sum = hist[n_bins * c + n_bins - 1].x; + auto total_sum = hist[n_bins * c + n_bins - 1].Weight(); auto rval_i = total_sum - lval_i; if (rval_i != 0) { auto rval = DataT(rval_i); @@ -138,10 +154,10 @@ class ClassificationObjectiveFunction { // Output probability double total = 0.0; for (int i = 0; i < nclasses; i++) { - total += shist[i].x; + total += shist[i].Weight(); } for (int i = 0; i < nclasses; i++) { - out[i] = DataT(shist[i].x) / total; + out[i] = DataT(shist[i].Weight()) / total; } } }; @@ -152,85 +168,99 @@ class RegressionObjectiveFunction { using DataT = DataT_; using LabelT = LabelT_; using IdxT = IdxT_; - using BinT = AggregateBin; + using BinT = RegressionBin; private: IdxT min_samples_leaf; CRITERION criterion; static constexpr auto eps_ = 10 * std::numeric_limits::epsilon(); - HDI DataT MSEGain(BinT const* hist, IdxT i, IdxT n_bins, IdxT len, IdxT nLeft, IdxT nRight) const + HDI DataT MSEGain(BinT const* hist, IdxT i, IdxT n_bins, IdxT, IdxT, IdxT) const { - auto invLen = DataT(1.0) / len; - auto label_sum = hist[n_bins - 1].label_sum; + auto parent_weight = hist[n_bins - 1].Weight(); + auto left_weight = hist[i].Weight(); + auto right_weight = parent_weight - left_weight; + + auto invLen = DataT(1.0) / DataT(parent_weight); + auto label_sum = DataT(hist[n_bins - 1].LabelSum()); + auto left_label_sum = DataT(hist[i].LabelSum()); DataT parent_obj = -label_sum * label_sum * invLen; - DataT left_obj = -(hist[i].label_sum * hist[i].label_sum) / nLeft; - DataT right_label_sum = hist[i].label_sum - label_sum; - DataT right_obj = -(right_label_sum * right_label_sum) / nRight; + DataT left_obj = -(left_label_sum * left_label_sum) / DataT(left_weight); + DataT right_label_sum = label_sum - left_label_sum; + DataT right_obj = -(right_label_sum * right_label_sum) / DataT(right_weight); DataT gain = parent_obj - (left_obj + right_obj); gain *= DataT(0.5) * invLen; return gain; } - HDI DataT - PoissonGain(BinT const* hist, IdxT i, IdxT n_bins, IdxT len, IdxT nLeft, IdxT nRight) const + HDI DataT PoissonGain(BinT const* hist, IdxT i, IdxT n_bins, IdxT, IdxT, IdxT) const { - auto invLen = DataT(1) / len; - auto label_sum = hist[n_bins - 1].label_sum; - auto left_label_sum = (hist[i].label_sum); - auto right_label_sum = (hist[n_bins - 1].label_sum - hist[i].label_sum); + auto parent_weight = hist[n_bins - 1].Weight(); + auto left_weight = hist[i].Weight(); + auto right_weight = parent_weight - left_weight; + + auto invLen = DataT(1) / DataT(parent_weight); + auto label_sum = DataT(hist[n_bins - 1].LabelSum()); + auto left_label_sum = DataT(hist[i].LabelSum()); + auto right_label_sum = DataT(hist[n_bins - 1].LabelSum() - hist[i].LabelSum()); // label sum cannot be non-positive if (label_sum < eps_ || left_label_sum < eps_ || right_label_sum < eps_) return -std::numeric_limits::max(); DataT parent_obj = -label_sum * raft::log(label_sum * invLen); - DataT left_obj = -left_label_sum * raft::log(left_label_sum / nLeft); - DataT right_obj = -right_label_sum * raft::log(right_label_sum / nRight); + DataT left_obj = -left_label_sum * raft::log(left_label_sum / DataT(left_weight)); + DataT right_obj = -right_label_sum * raft::log(right_label_sum / DataT(right_weight)); DataT gain = parent_obj - (left_obj + right_obj); gain = gain * invLen; return gain; } - HDI DataT - GammaGain(BinT const* hist, IdxT i, IdxT n_bins, IdxT len, IdxT nLeft, IdxT nRight) const + HDI DataT GammaGain(BinT const* hist, IdxT i, IdxT n_bins, IdxT, IdxT, IdxT) const { - auto invLen = DataT(1) / len; - auto label_sum = hist[n_bins - 1].label_sum; - auto left_label_sum = (hist[i].label_sum); - auto right_label_sum = (hist[n_bins - 1].label_sum - hist[i].label_sum); + auto parent_weight = hist[n_bins - 1].Weight(); + auto left_weight = hist[i].Weight(); + auto right_weight = parent_weight - left_weight; + + auto invLen = DataT(1) / DataT(parent_weight); + auto label_sum = DataT(hist[n_bins - 1].LabelSum()); + auto left_label_sum = DataT(hist[i].LabelSum()); + auto right_label_sum = DataT(hist[n_bins - 1].LabelSum() - hist[i].LabelSum()); // label sum cannot be non-positive if (label_sum < eps_ || left_label_sum < eps_ || right_label_sum < eps_) return -std::numeric_limits::max(); - DataT parent_obj = len * raft::log(label_sum * invLen); - DataT left_obj = nLeft * raft::log(left_label_sum / nLeft); - DataT right_obj = nRight * raft::log(right_label_sum / nRight); + DataT parent_obj = DataT(parent_weight) * raft::log(label_sum * invLen); + DataT left_obj = DataT(left_weight) * raft::log(left_label_sum / DataT(left_weight)); + DataT right_obj = DataT(right_weight) * raft::log(right_label_sum / DataT(right_weight)); DataT gain = parent_obj - (left_obj + right_obj); gain = gain * invLen; return gain; } - HDI DataT InverseGaussianGain( - BinT const* hist, IdxT i, IdxT n_bins, IdxT len, IdxT nLeft, IdxT nRight) const + HDI DataT InverseGaussianGain(BinT const* hist, IdxT i, IdxT n_bins, IdxT, IdxT, IdxT) const { - auto label_sum = hist[n_bins - 1].label_sum; - auto left_label_sum = (hist[i].label_sum); - auto right_label_sum = (hist[n_bins - 1].label_sum - hist[i].label_sum); + auto parent_weight = hist[n_bins - 1].Weight(); + auto left_weight = hist[i].Weight(); + auto right_weight = parent_weight - left_weight; + + auto label_sum = DataT(hist[n_bins - 1].LabelSum()); + auto left_label_sum = DataT(hist[i].LabelSum()); + auto right_label_sum = DataT(hist[n_bins - 1].LabelSum() - hist[i].LabelSum()); // label sum cannot be non-positive if (label_sum < eps_ || left_label_sum < eps_ || right_label_sum < eps_) return -std::numeric_limits::max(); - DataT parent_obj = -DataT(len) * DataT(len) / label_sum; - DataT left_obj = -DataT(nLeft) * DataT(nLeft) / left_label_sum; - DataT right_obj = -DataT(nRight) * DataT(nRight) / right_label_sum; + DataT parent_obj = -DataT(parent_weight) * DataT(parent_weight) / label_sum; + DataT left_obj = -DataT(left_weight) * DataT(left_weight) / left_label_sum; + DataT right_obj = -DataT(right_weight) * DataT(right_weight) / right_label_sum; DataT gain = parent_obj - (left_obj + right_obj); - gain = gain / (2 * len); + gain = gain / (2 * DataT(parent_weight)); return gain; } @@ -261,7 +291,7 @@ class RegressionObjectiveFunction { { Split sp; for (IdxT i = threadIdx.x; i < n_bins; i += blockDim.x) { - auto nLeft = shist[i].count; + auto nLeft = static_cast(shist[i].Count()); auto nRight = len - nLeft; auto gain = -std::numeric_limits::max(); if (nLeft >= min_samples_leaf && nRight >= min_samples_leaf) { @@ -275,7 +305,7 @@ class RegressionObjectiveFunction { static DI void SetLeafVector(BinT const* shist, int nclasses, DataT* out) { for (int i = 0; i < nclasses; i++) { - out[i] = shist[i].label_sum / shist[i].count; + out[i] = shist[i].LabelSum() / shist[i].Weight(); } } }; diff --git a/cpp/tests/sg/rf_test.cu b/cpp/tests/sg/rf_test.cu index 8dd1c7b54c..4accaa3c16 100644 --- a/cpp/tests/sg/rf_test.cu +++ b/cpp/tests/sg/rf_test.cu @@ -1153,7 +1153,7 @@ class ObjectiveTest : public ::testing::TestWithParam { { std::default_random_engine rng; std::vector data(params.n_rows); - if constexpr (std::is_same::value) // classification case + if constexpr (std::is_same::value) // classification case { for (auto& d : data) { d = RandUnder(params.n_classes); @@ -1181,7 +1181,7 @@ class ObjectiveTest : public ::testing::TestWithParam { IdxT bin_width = raft::ceildiv(params.n_rows, params.max_n_bins); auto data_begin = data.begin() + b * bin_width; auto data_end = data_begin + bin_width; - if constexpr (std::is_same::value) { // classification case + if constexpr (std::is_same::value) { // classification case auto count{IdxT(0)}; std::for_each(data_begin, data_end, [&](auto d) { if (d == c) ++count; @@ -1448,13 +1448,7 @@ class ObjectiveTest : public ::testing::TestWithParam { { auto count{IdxT(0)}; for (auto c = 0; c < params.n_classes; ++c) { - if constexpr (std::is_same::value) // countbin - { - count += static_cast(cdf_hist[params.max_n_bins * c + idx].x); - } else // aggregatebin - { - count += cdf_hist[params.max_n_bins * c + idx].count; - } + count += static_cast(cdf_hist[params.max_n_bins * c + idx].Count()); } return count; } From 5373dd34862a2a927aef66722e946b840dae825f Mon Sep 17 00:00:00 2001 From: Rory Mitchell Date: Wed, 10 Jun 2026 04:06:19 -0700 Subject: [PATCH 2/3] Add weighted random forest objective tests --- .../batched-levelalgo/objectives.cuh | 11 +- cpp/tests/sg/rf_test.cu | 494 +++++++++++++----- 2 files changed, 358 insertions(+), 147 deletions(-) diff --git a/cpp/src/decisiontree/batched-levelalgo/objectives.cuh b/cpp/src/decisiontree/batched-levelalgo/objectives.cuh index 6e9ac16496..f7456ba8b3 100644 --- a/cpp/src/decisiontree/batched-levelalgo/objectives.cuh +++ b/cpp/src/decisiontree/batched-levelalgo/objectives.cuh @@ -12,17 +12,19 @@ #include #include +#include namespace ML { namespace DT { -template +template class ClassificationObjectiveFunction { public: using DataT = DataT_; using LabelT = LabelT_; using IdxT = IdxT_; - using BinT = ClassificationBin; + using BinT = std::conditional_t; + static constexpr bool weighted = weighted_; private: IdxT nclasses; @@ -162,13 +164,14 @@ class ClassificationObjectiveFunction { } }; -template +template class RegressionObjectiveFunction { public: using DataT = DataT_; using LabelT = LabelT_; using IdxT = IdxT_; - using BinT = RegressionBin; + using BinT = std::conditional_t; + static constexpr bool weighted = weighted_; private: IdxT min_samples_leaf; diff --git a/cpp/tests/sg/rf_test.cu b/cpp/tests/sg/rf_test.cu index 4accaa3c16..b816c69ba6 100644 --- a/cpp/tests/sg/rf_test.cu +++ b/cpp/tests/sg/rf_test.cu @@ -1142,18 +1142,25 @@ class ObjectiveTest : public ::testing::TestWithParam { typedef typename ObjectiveT::IdxT IdxT; typedef typename ObjectiveT::BinT BinT; - static constexpr auto eps_ = 10 * std::numeric_limits::epsilon(); + static constexpr auto eps_ = 10 * std::numeric_limits::epsilon(); + static constexpr bool is_classification = std::is_same::value || + std::is_same::value; + static constexpr bool is_weighted = ObjectiveT::weighted; ObjectiveTestParameters params; + std::mt19937_64 rng; public: - auto RandUnder(int const end = 10000) { return rand() % end; } + auto RandUnder(int const end = 10000) + { + std::uniform_int_distribution dist(0, end - 1); + return dist(rng); + } auto GenRandomData() { - std::default_random_engine rng; std::vector data(params.n_rows); - if constexpr (std::is_same::value) // classification case + if constexpr (is_classification) // classification case { for (auto& d : data) { d = RandUnder(params.n_classes); @@ -1172,25 +1179,53 @@ class ObjectiveTest : public ::testing::TestWithParam { return data; } - auto GenHist(std::vector data) + auto GenSampleWeights() + { + std::vector sample_weights(params.n_rows, DataT(1)); + if constexpr (is_weighted) { + std::uniform_real_distribution weight_dist(DataT(0.2), DataT(3.0)); + for (auto& w : sample_weights) { + w = weight_dist(rng); + } + } + return sample_weights; + } + + auto GenHist(std::vector const& data, std::vector const& sample_weights) { std::vector cdf_hist, pdf_hist; + IdxT bin_width = raft::ceildiv(params.n_rows, params.max_n_bins); for (auto c = 0; c < params.n_classes; ++c) { for (auto b = 0; b < params.max_n_bins; ++b) { - IdxT bin_width = raft::ceildiv(params.n_rows, params.max_n_bins); - auto data_begin = data.begin() + b * bin_width; - auto data_end = data_begin + bin_width; - if constexpr (std::is_same::value) { // classification case - auto count{IdxT(0)}; - std::for_each(data_begin, data_end, [&](auto d) { - if (d == c) ++count; - }); - pdf_hist.emplace_back(count); + auto bin_begin = b * bin_width; + auto bin_end = bin_begin + bin_width; + if constexpr (is_classification) { + auto count{BinCountT(0)}; + auto weight{DataT(0)}; + for (auto i = bin_begin; i < bin_end; ++i) { + if (data[i] == DataT(c)) { + ++count; + weight += sample_weights[i]; + } + } + if constexpr (is_weighted) { + pdf_hist.emplace_back(count, weight); + } else { + pdf_hist.emplace_back(count); + } } else { // regression case auto label_sum{DataT(0)}; - label_sum = std::accumulate(data_begin, data_end, DataT(0)); - pdf_hist.emplace_back(label_sum, bin_width); + auto weight{DataT(0)}; + for (auto i = bin_begin; i < bin_end; ++i) { + label_sum += data[i] * sample_weights[i]; + weight += sample_weights[i]; + } + if constexpr (is_weighted) { + pdf_hist.emplace_back(label_sum, bin_width, weight); + } else { + pdf_hist.emplace_back(label_sum, bin_width); + } } auto cumulative = b > 0 ? cdf_hist.back() : BinT(); @@ -1202,33 +1237,54 @@ class ObjectiveTest : public ::testing::TestWithParam { return std::make_pair(cdf_hist, pdf_hist); } - auto MSE(std::vector const& data) // 1/n * 1/2 * sum((y - y_pred) * (y - y_pred)) + auto SplitOffset(std::size_t const split_bin_index) { - DataT sum = std::accumulate(data.begin(), data.end(), DataT(0)); - DataT const mean = sum / data.size(); + auto bin_width = raft::ceildiv(params.n_rows, params.max_n_bins); + return (split_bin_index + 1) * bin_width; + } + + auto MSE(std::vector const& data, + std::vector const& + sample_weights) // 1/w * 1/2 * sum(w_i * (y - y_pred) * (y - y_pred)) + { + DataT weight_sum = std::accumulate(sample_weights.begin(), sample_weights.end(), DataT(0)); + DataT sum{0}; + for (std::size_t i = 0; i < data.size(); ++i) { + sum += data[i] * sample_weights[i]; + } + DataT const mean = sum / weight_sum; auto mse{DataT(0.0)}; // mse: mean squared error - std::for_each(data.begin(), data.end(), [&](auto d) { - mse += (d - mean) * (d - mean); // unit deviance - }); + for (std::size_t i = 0; i < data.size(); ++i) { + auto d = data[i]; + mse += sample_weights[i] * (d - mean) * (d - mean); // unit deviance + } - mse /= 2 * data.size(); - return std::make_tuple(mse, sum, DataT(data.size())); + mse /= 2 * weight_sum; + return std::make_tuple(mse, sum, DataT(data.size()), weight_sum); } - auto MSEGroundTruthGain(std::vector const& data, std::size_t split_bin_index) + auto MSEGroundTruthGain(std::vector const& data, + std::vector const& sample_weights, + std::size_t split_bin_index) { - auto bin_width = raft::ceildiv(params.n_rows, params.max_n_bins); - std::vector left_data(data.begin(), data.begin() + (split_bin_index + 1) * bin_width); - std::vector right_data(data.begin() + (split_bin_index + 1) * bin_width, data.end()); - - auto [parent_mse, label_sum, n] = MSE(data); - auto [left_mse, label_sum_left, n_left] = MSE(left_data); - auto [right_mse, label_sum_right, n_right] = MSE(right_data); - - auto gain = - parent_mse - ((n_left / n) * left_mse + // the minimizing objective function is half deviance - (n_right / n) * right_mse); // gain in long form without proxy + auto split_offset = SplitOffset(split_bin_index); + std::vector left_data(data.begin(), data.begin() + split_offset); + std::vector right_data(data.begin() + split_offset, data.end()); + std::vector left_sample_weights(sample_weights.begin(), + sample_weights.begin() + split_offset); + std::vector right_sample_weights(sample_weights.begin() + split_offset, + sample_weights.end()); + + auto [parent_mse, label_sum, n, weight_sum] = MSE(data, sample_weights); + auto [left_mse, label_sum_left, n_left, weight_sum_left] = MSE(left_data, left_sample_weights); + auto [right_mse, label_sum_right, n_right, weight_sum_right] = + MSE(right_data, right_sample_weights); + + auto gain = parent_mse - + ((weight_sum_left / weight_sum) * left_mse + // the minimizing objective function + // is half deviance + (weight_sum_right / weight_sum) * right_mse); // gain in long form without proxy // edge cases if (n_left < params.min_samples_leaf or n_right < params.min_samples_leaf) @@ -1238,34 +1294,50 @@ class ObjectiveTest : public ::testing::TestWithParam { } auto InverseGaussianHalfDeviance( - std::vector const& - data) // 1/n * 2 * sum((y - y_pred) * (y - y_pred)/(y * (y_pred) * (y_pred))) + std::vector const& data, + std::vector const& sample_weights) // 1/w * 2 * sum(w_i * (y - y_pred) * (y - + // y_pred)/(y * (y_pred) * (y_pred))) { - DataT sum = std::accumulate(data.begin(), data.end(), DataT(0)); - DataT const mean = sum / data.size(); + DataT weight_sum = std::accumulate(sample_weights.begin(), sample_weights.end(), DataT(0)); + DataT sum{0}; + for (std::size_t i = 0; i < data.size(); ++i) { + sum += data[i] * sample_weights[i]; + } + DataT const mean = sum / weight_sum; auto ighd{DataT(0.0)}; // ighd: inverse gaussian half deviance - std::for_each(data.begin(), data.end(), [&](auto d) { - ighd += (d - mean) * (d - mean) / (d * mean * mean); // unit deviance - }); + for (std::size_t i = 0; i < data.size(); ++i) { + auto d = data[i]; + ighd += sample_weights[i] * (d - mean) * (d - mean) / (d * mean * mean); // unit deviance + } - ighd /= 2 * data.size(); - return std::make_tuple(ighd, sum, DataT(data.size())); + ighd /= 2 * weight_sum; + return std::make_tuple(ighd, sum, DataT(data.size()), weight_sum); } - auto InverseGaussianGroundTruthGain(std::vector const& data, std::size_t split_bin_index) + auto InverseGaussianGroundTruthGain(std::vector const& data, + std::vector const& sample_weights, + std::size_t split_bin_index) { - auto bin_width = raft::ceildiv(params.n_rows, params.max_n_bins); - std::vector left_data(data.begin(), data.begin() + (split_bin_index + 1) * bin_width); - std::vector right_data(data.begin() + (split_bin_index + 1) * bin_width, data.end()); - - auto [parent_ighd, label_sum, n] = InverseGaussianHalfDeviance(data); - auto [left_ighd, label_sum_left, n_left] = InverseGaussianHalfDeviance(left_data); - auto [right_ighd, label_sum_right, n_right] = InverseGaussianHalfDeviance(right_data); + auto split_offset = SplitOffset(split_bin_index); + std::vector left_data(data.begin(), data.begin() + split_offset); + std::vector right_data(data.begin() + split_offset, data.end()); + std::vector left_sample_weights(sample_weights.begin(), + sample_weights.begin() + split_offset); + std::vector right_sample_weights(sample_weights.begin() + split_offset, + sample_weights.end()); + + auto [parent_ighd, label_sum, n, weight_sum] = + InverseGaussianHalfDeviance(data, sample_weights); + auto [left_ighd, label_sum_left, n_left, weight_sum_left] = + InverseGaussianHalfDeviance(left_data, left_sample_weights); + auto [right_ighd, label_sum_right, n_right, weight_sum_right] = + InverseGaussianHalfDeviance(right_data, right_sample_weights); auto gain = parent_ighd - - ((n_left / n) * left_ighd + // the minimizing objective function is half deviance - (n_right / n) * right_ighd); // gain in long form without proxy + ((weight_sum_left / weight_sum) * + left_ighd + // the minimizing objective function is half deviance + (weight_sum_right / weight_sum) * right_ighd); // gain in long form without proxy // edge cases if (n_left < params.min_samples_leaf or n_right < params.min_samples_leaf or label_sum < eps_ or @@ -1275,36 +1347,49 @@ class ObjectiveTest : public ::testing::TestWithParam { return gain; } - auto GammaHalfDeviance( - std::vector const& data) // 1/n * 2 * sum(log(y_pred/y_true) + y_true/y_pred - 1) + auto GammaHalfDeviance(std::vector const& data, std::vector const& sample_weights) + // 1/w * 2 * sum(w_i * (log(y_pred/y_true) + y_true/y_pred - 1)) { + DataT weight_sum = std::accumulate(sample_weights.begin(), sample_weights.end(), DataT(0)); DataT sum(0); - sum = std::accumulate(data.begin(), data.end(), DataT(0)); - DataT const mean = sum / data.size(); + for (std::size_t i = 0; i < data.size(); ++i) { + sum += data[i] * sample_weights[i]; + } + DataT const mean = sum / weight_sum; DataT ghd(0); // gamma half deviance - std::for_each(data.begin(), data.end(), [&](auto& element) { - auto log_y = raft::log(element ? element : DataT(1.0)); - ghd += raft::log(mean) - log_y + element / mean - 1; - }); + for (std::size_t i = 0; i < data.size(); ++i) { + auto& element = data[i]; + auto log_y = raft::log(element ? element : DataT(1.0)); + ghd += sample_weights[i] * (raft::log(mean) - log_y + element / mean - 1); + } - ghd /= data.size(); - return std::make_tuple(ghd, sum, DataT(data.size())); + ghd /= weight_sum; + return std::make_tuple(ghd, sum, DataT(data.size()), weight_sum); } - auto GammaGroundTruthGain(std::vector const& data, std::size_t split_bin_index) + auto GammaGroundTruthGain(std::vector const& data, + std::vector const& sample_weights, + std::size_t split_bin_index) { - auto bin_width = raft::ceildiv(params.n_rows, params.max_n_bins); - std::vector left_data(data.begin(), data.begin() + (split_bin_index + 1) * bin_width); - std::vector right_data(data.begin() + (split_bin_index + 1) * bin_width, data.end()); - - auto [parent_ghd, label_sum, n] = GammaHalfDeviance(data); - auto [left_ghd, label_sum_left, n_left] = GammaHalfDeviance(left_data); - auto [right_ghd, label_sum_right, n_right] = GammaHalfDeviance(right_data); - - auto gain = - parent_ghd - ((n_left / n) * left_ghd + // the minimizing objective function is half deviance - (n_right / n) * right_ghd); // gain in long form without proxy + auto split_offset = SplitOffset(split_bin_index); + std::vector left_data(data.begin(), data.begin() + split_offset); + std::vector right_data(data.begin() + split_offset, data.end()); + std::vector left_sample_weights(sample_weights.begin(), + sample_weights.begin() + split_offset); + std::vector right_sample_weights(sample_weights.begin() + split_offset, + sample_weights.end()); + + auto [parent_ghd, label_sum, n, weight_sum] = GammaHalfDeviance(data, sample_weights); + auto [left_ghd, label_sum_left, n_left, weight_sum_left] = + GammaHalfDeviance(left_data, left_sample_weights); + auto [right_ghd, label_sum_right, n_right, weight_sum_right] = + GammaHalfDeviance(right_data, right_sample_weights); + + auto gain = parent_ghd - + ((weight_sum_left / weight_sum) * left_ghd + // the minimizing objective function + // is half deviance + (weight_sum_right / weight_sum) * right_ghd); // gain in long form without proxy // edge cases if (n_left < params.min_samples_leaf or n_right < params.min_samples_leaf or label_sum < eps_ or @@ -1315,33 +1400,49 @@ class ObjectiveTest : public ::testing::TestWithParam { } auto PoissonHalfDeviance( - std::vector const& data) // 1/n * sum(y_true * log(y_true/y_pred) + y_pred - y_true) + std::vector const& data, + std::vector const& + sample_weights) // 1/w * sum(w_i * (y_true * log(y_true/y_pred) + y_pred - y_true)) { - DataT sum = std::accumulate(data.begin(), data.end(), DataT(0)); - auto const mean = sum / data.size(); + DataT weight_sum = std::accumulate(sample_weights.begin(), sample_weights.end(), DataT(0)); + DataT sum{0}; + for (std::size_t i = 0; i < data.size(); ++i) { + sum += data[i] * sample_weights[i]; + } + auto const mean = sum / weight_sum; auto poisson_half_deviance{DataT(0.0)}; - std::for_each(data.begin(), data.end(), [&](auto d) { + for (std::size_t i = 0; i < data.size(); ++i) { + auto d = data[i]; auto log_y = raft::log(d ? d : DataT(1.0)); // we don't want nans - poisson_half_deviance += d * (log_y - raft::log(mean)) + mean - d; - }); + poisson_half_deviance += sample_weights[i] * (d * (log_y - raft::log(mean)) + mean - d); + } - poisson_half_deviance /= data.size(); - return std::make_tuple(poisson_half_deviance, sum, DataT(data.size())); + poisson_half_deviance /= weight_sum; + return std::make_tuple(poisson_half_deviance, sum, DataT(data.size()), weight_sum); } - auto PoissonGroundTruthGain(std::vector const& data, std::size_t split_bin_index) + auto PoissonGroundTruthGain(std::vector const& data, + std::vector const& sample_weights, + std::size_t split_bin_index) { - auto bin_width = raft::ceildiv(params.n_rows, params.max_n_bins); - std::vector left_data(data.begin(), data.begin() + (split_bin_index + 1) * bin_width); - std::vector right_data(data.begin() + (split_bin_index + 1) * bin_width, data.end()); - - auto [parent_phd, label_sum, n] = PoissonHalfDeviance(data); - auto [left_phd, label_sum_left, n_left] = PoissonHalfDeviance(left_data); - auto [right_phd, label_sum_right, n_right] = PoissonHalfDeviance(right_data); - - auto gain = parent_phd - ((n_left / n) * left_phd + - (n_right / n) * right_phd); // gain in long form without proxy + auto split_offset = SplitOffset(split_bin_index); + std::vector left_data(data.begin(), data.begin() + split_offset); + std::vector right_data(data.begin() + split_offset, data.end()); + std::vector left_sample_weights(sample_weights.begin(), + sample_weights.begin() + split_offset); + std::vector right_sample_weights(sample_weights.begin() + split_offset, + sample_weights.end()); + + auto [parent_phd, label_sum, n, weight_sum] = PoissonHalfDeviance(data, sample_weights); + auto [left_phd, label_sum_left, n_left, weight_sum_left] = + PoissonHalfDeviance(left_data, left_sample_weights); + auto [right_phd, label_sum_right, n_right, weight_sum_right] = + PoissonHalfDeviance(right_data, right_sample_weights); + + auto gain = parent_phd - + ((weight_sum_left / weight_sum) * left_phd + + (weight_sum_right / weight_sum) * right_phd); // gain in long form without proxy // edge cases if (n_left < params.min_samples_leaf or n_right < params.min_samples_leaf or label_sum < eps_ or @@ -1351,95 +1452,117 @@ class ObjectiveTest : public ::testing::TestWithParam { return gain; } - auto Entropy(std::vector const& data) + auto Entropy(std::vector const& data, std::vector const& sample_weights) { // sum((n_c/n_total)*(log(n_c/n_total))) + DataT weight_sum = std::accumulate(sample_weights.begin(), sample_weights.end(), DataT(0)); DataT entropy(0); for (auto c = 0; c < params.n_classes; ++c) { - IdxT sum(0); - std::for_each(data.begin(), data.end(), [&](auto d) { - if (d == DataT(c)) ++sum; - }); - DataT class_proba = DataT(sum) / data.size(); + DataT sum(0); + for (std::size_t i = 0; i < data.size(); ++i) { + if (data[i] == DataT(c)) { sum += sample_weights[i]; } + } + DataT class_proba = sum / weight_sum; entropy += -class_proba * raft::log(class_proba ? class_proba : DataT(1)) / raft::log(DataT(2)); // adding gain } return entropy; } - auto EntropyGroundTruthGain(std::vector const& data, std::size_t const split_bin_index) + auto EntropyGroundTruthGain(std::vector const& data, + std::vector const& sample_weights, + std::size_t const split_bin_index) { - auto bin_width = raft::ceildiv(params.n_rows, params.max_n_bins); - std::vector left_data(data.begin(), data.begin() + (split_bin_index + 1) * bin_width); - std::vector right_data(data.begin() + (split_bin_index + 1) * bin_width, data.end()); - - auto parent_entropy = Entropy(data); - auto left_entropy = Entropy(left_data); - auto right_entropy = Entropy(right_data); - DataT n = data.size(); - DataT left_n = left_data.size(); - DataT right_n = right_data.size(); + auto split_offset = SplitOffset(split_bin_index); + std::vector left_data(data.begin(), data.begin() + split_offset); + std::vector right_data(data.begin() + split_offset, data.end()); + std::vector left_sample_weights(sample_weights.begin(), + sample_weights.begin() + split_offset); + std::vector right_sample_weights(sample_weights.begin() + split_offset, + sample_weights.end()); + + auto parent_entropy = Entropy(data, sample_weights); + auto left_entropy = Entropy(left_data, left_sample_weights); + auto right_entropy = Entropy(right_data, right_sample_weights); + DataT n = std::accumulate(sample_weights.begin(), sample_weights.end(), DataT(0)); + DataT left_n = + std::accumulate(left_sample_weights.begin(), left_sample_weights.end(), DataT(0)); + DataT right_n = + std::accumulate(right_sample_weights.begin(), right_sample_weights.end(), DataT(0)); auto gain = parent_entropy - ((left_n / n) * left_entropy + (right_n / n) * right_entropy); // edge cases - if (left_n < params.min_samples_leaf or right_n < params.min_samples_leaf) { + if (left_data.size() < std::size_t(params.min_samples_leaf) or + right_data.size() < std::size_t(params.min_samples_leaf)) { return -std::numeric_limits::max(); } else { return gain; } } - auto GiniImpurity(std::vector const& data) + auto GiniImpurity(std::vector const& data, std::vector const& sample_weights) { // sum((n_c/n_total)(1-(n_c/n_total))) + DataT weight_sum = std::accumulate(sample_weights.begin(), sample_weights.end(), DataT(0)); DataT gini(0); for (auto c = 0; c < params.n_classes; ++c) { - IdxT sum(0); - std::for_each(data.begin(), data.end(), [&](auto d) { - if (d == DataT(c)) ++sum; - }); - DataT class_proba = DataT(sum) / data.size(); + DataT sum(0); + for (std::size_t i = 0; i < data.size(); ++i) { + if (data[i] == DataT(c)) { sum += sample_weights[i]; } + } + DataT class_proba = sum / weight_sum; gini += class_proba * (1 - class_proba); // adding gain } return gini; } - auto GiniGroundTruthGain(std::vector const& data, std::size_t const split_bin_index) + auto GiniGroundTruthGain(std::vector const& data, + std::vector const& sample_weights, + std::size_t const split_bin_index) { - auto bin_width = raft::ceildiv(params.n_rows, params.max_n_bins); - std::vector left_data(data.begin(), data.begin() + (split_bin_index + 1) * bin_width); - std::vector right_data(data.begin() + (split_bin_index + 1) * bin_width, data.end()); - - auto parent_gini = GiniImpurity(data); - auto left_gini = GiniImpurity(left_data); - auto right_gini = GiniImpurity(right_data); - DataT n = data.size(); - DataT left_n = left_data.size(); - DataT right_n = right_data.size(); + auto split_offset = SplitOffset(split_bin_index); + std::vector left_data(data.begin(), data.begin() + split_offset); + std::vector right_data(data.begin() + split_offset, data.end()); + std::vector left_sample_weights(sample_weights.begin(), + sample_weights.begin() + split_offset); + std::vector right_sample_weights(sample_weights.begin() + split_offset, + sample_weights.end()); + + auto parent_gini = GiniImpurity(data, sample_weights); + auto left_gini = GiniImpurity(left_data, left_sample_weights); + auto right_gini = GiniImpurity(right_data, right_sample_weights); + DataT n = std::accumulate(sample_weights.begin(), sample_weights.end(), DataT(0)); + DataT left_n = + std::accumulate(left_sample_weights.begin(), left_sample_weights.end(), DataT(0)); + DataT right_n = + std::accumulate(right_sample_weights.begin(), right_sample_weights.end(), DataT(0)); auto gain = parent_gini - ((left_n / n) * left_gini + (right_n / n) * right_gini); // edge cases - if (left_n < params.min_samples_leaf or right_n < params.min_samples_leaf) { + if (left_data.size() < std::size_t(params.min_samples_leaf) or + right_data.size() < std::size_t(params.min_samples_leaf)) { return -std::numeric_limits::max(); } else { return gain; } } - auto GroundTruthGain(std::vector const& data, std::size_t const split_bin_index) + auto GroundTruthGain(std::vector const& data, + std::vector const& sample_weights, + std::size_t const split_bin_index) { if constexpr (ObjectiveConfig::splitCriteria == CRITERION::MSE) { - return MSEGroundTruthGain(data, split_bin_index); + return MSEGroundTruthGain(data, sample_weights, split_bin_index); } else if constexpr (ObjectiveConfig::splitCriteria == CRITERION::POISSON) { - return PoissonGroundTruthGain(data, split_bin_index); + return PoissonGroundTruthGain(data, sample_weights, split_bin_index); } else if constexpr (ObjectiveConfig::splitCriteria == CRITERION::GAMMA) { - return GammaGroundTruthGain(data, split_bin_index); + return GammaGroundTruthGain(data, sample_weights, split_bin_index); } else if constexpr (ObjectiveConfig::splitCriteria == CRITERION::INVERSE_GAUSSIAN) { - return InverseGaussianGroundTruthGain(data, split_bin_index); + return InverseGaussianGroundTruthGain(data, sample_weights, split_bin_index); } else if constexpr (ObjectiveConfig::splitCriteria == CRITERION::ENTROPY) { - return EntropyGroundTruthGain(data, split_bin_index); + return EntropyGroundTruthGain(data, sample_weights, split_bin_index); } else if constexpr (ObjectiveConfig::splitCriteria == CRITERION::GINI) { - return GiniGroundTruthGain(data, split_bin_index); + return GiniGroundTruthGain(data, sample_weights, split_bin_index); } return DataT(0.0); } @@ -1456,13 +1579,14 @@ class ObjectiveTest : public ::testing::TestWithParam { void SetUp() override { params = ::testing::TestWithParam::GetParam(); - srand(params.seed); + rng.seed(params.seed); ObjectiveT objective(params.n_classes, params.min_samples_leaf, ObjectiveConfig::splitCriteria); auto data = GenRandomData(); - auto [cdf_hist, pdf_hist] = GenHist(data); - auto split_bin_index = RandUnder(params.max_n_bins); - auto ground_truth_gain = GroundTruthGain(data, split_bin_index); + auto sample_weights = GenSampleWeights(); + auto [cdf_hist, pdf_hist] = GenHist(data, sample_weights); + auto split_bin_index = RandUnder(params.max_n_bins - 1); + auto ground_truth_gain = GroundTruthGain(data, sample_weights, split_bin_index); auto len = NumLeftOfBin(cdf_hist, params.max_n_bins - 1); auto nLeft = NumLeftOfBin(cdf_hist, split_bin_index); @@ -1534,6 +1658,20 @@ TEST_P(MSEObjectiveTestF, MSEObjectiveTest) {} INSTANTIATE_TEST_CASE_P(RfTests, MSEObjectiveTestF, ::testing::ValuesIn(mse_objective_test_parameters)); +typedef ObjectiveTest< + ObjectiveTestConfig, CRITERION::MSE>> + WeightedMSEObjectiveTestD; +TEST_P(WeightedMSEObjectiveTestD, MSEObjectiveTest) {} +INSTANTIATE_TEST_CASE_P(RfTests, + WeightedMSEObjectiveTestD, + ::testing::ValuesIn(mse_objective_test_parameters)); +typedef ObjectiveTest< + ObjectiveTestConfig, CRITERION::MSE>> + WeightedMSEObjectiveTestF; +TEST_P(WeightedMSEObjectiveTestF, MSEObjectiveTest) {} +INSTANTIATE_TEST_CASE_P(RfTests, + WeightedMSEObjectiveTestF, + ::testing::ValuesIn(mse_objective_test_parameters)); // poisson objective test typedef ObjectiveTest< @@ -1550,6 +1688,20 @@ TEST_P(PoissonObjectiveTestF, poissonObjectiveTest) {} INSTANTIATE_TEST_CASE_P(RfTests, PoissonObjectiveTestF, ::testing::ValuesIn(poisson_objective_test_parameters)); +typedef ObjectiveTest< + ObjectiveTestConfig, CRITERION::POISSON>> + WeightedPoissonObjectiveTestD; +TEST_P(WeightedPoissonObjectiveTestD, poissonObjectiveTest) {} +INSTANTIATE_TEST_CASE_P(RfTests, + WeightedPoissonObjectiveTestD, + ::testing::ValuesIn(poisson_objective_test_parameters)); +typedef ObjectiveTest< + ObjectiveTestConfig, CRITERION::POISSON>> + WeightedPoissonObjectiveTestF; +TEST_P(WeightedPoissonObjectiveTestF, poissonObjectiveTest) {} +INSTANTIATE_TEST_CASE_P(RfTests, + WeightedPoissonObjectiveTestF, + ::testing::ValuesIn(poisson_objective_test_parameters)); // gamma objective test typedef ObjectiveTest< @@ -1566,6 +1718,20 @@ TEST_P(GammaObjectiveTestF, GammaObjectiveTest) {} INSTANTIATE_TEST_CASE_P(RfTests, GammaObjectiveTestF, ::testing::ValuesIn(gamma_objective_test_parameters)); +typedef ObjectiveTest< + ObjectiveTestConfig, CRITERION::GAMMA>> + WeightedGammaObjectiveTestD; +TEST_P(WeightedGammaObjectiveTestD, GammaObjectiveTest) {} +INSTANTIATE_TEST_CASE_P(RfTests, + WeightedGammaObjectiveTestD, + ::testing::ValuesIn(gamma_objective_test_parameters)); +typedef ObjectiveTest< + ObjectiveTestConfig, CRITERION::GAMMA>> + WeightedGammaObjectiveTestF; +TEST_P(WeightedGammaObjectiveTestF, GammaObjectiveTest) {} +INSTANTIATE_TEST_CASE_P(RfTests, + WeightedGammaObjectiveTestF, + ::testing::ValuesIn(gamma_objective_test_parameters)); // InvGauss objective test typedef ObjectiveTest, @@ -1582,6 +1748,20 @@ TEST_P(InverseGaussianObjectiveTestF, InverseGaussianObjectiveTest) {} INSTANTIATE_TEST_CASE_P(RfTests, InverseGaussianObjectiveTestF, ::testing::ValuesIn(invgauss_objective_test_parameters)); +typedef ObjectiveTest, + CRITERION::INVERSE_GAUSSIAN>> + WeightedInverseGaussianObjectiveTestD; +TEST_P(WeightedInverseGaussianObjectiveTestD, InverseGaussianObjectiveTest) {} +INSTANTIATE_TEST_CASE_P(RfTests, + WeightedInverseGaussianObjectiveTestD, + ::testing::ValuesIn(invgauss_objective_test_parameters)); +typedef ObjectiveTest, + CRITERION::INVERSE_GAUSSIAN>> + WeightedInverseGaussianObjectiveTestF; +TEST_P(WeightedInverseGaussianObjectiveTestF, InverseGaussianObjectiveTest) {} +INSTANTIATE_TEST_CASE_P(RfTests, + WeightedInverseGaussianObjectiveTestF, + ::testing::ValuesIn(invgauss_objective_test_parameters)); // entropy objective test typedef ObjectiveTest< @@ -1598,6 +1778,20 @@ TEST_P(EntropyObjectiveTestF, entropyObjectiveTest) {} INSTANTIATE_TEST_CASE_P(RfTests, EntropyObjectiveTestF, ::testing::ValuesIn(entropy_objective_test_parameters)); +typedef ObjectiveTest< + ObjectiveTestConfig, CRITERION::ENTROPY>> + WeightedEntropyObjectiveTestD; +TEST_P(WeightedEntropyObjectiveTestD, entropyObjectiveTest) {} +INSTANTIATE_TEST_CASE_P(RfTests, + WeightedEntropyObjectiveTestD, + ::testing::ValuesIn(entropy_objective_test_parameters)); +typedef ObjectiveTest< + ObjectiveTestConfig, CRITERION::ENTROPY>> + WeightedEntropyObjectiveTestF; +TEST_P(WeightedEntropyObjectiveTestF, entropyObjectiveTest) {} +INSTANTIATE_TEST_CASE_P(RfTests, + WeightedEntropyObjectiveTestF, + ::testing::ValuesIn(entropy_objective_test_parameters)); // gini objective test typedef ObjectiveTest< @@ -1614,6 +1808,20 @@ TEST_P(GiniObjectiveTestF, giniObjectiveTest) {} INSTANTIATE_TEST_CASE_P(RfTests, GiniObjectiveTestF, ::testing::ValuesIn(gini_objective_test_parameters)); +typedef ObjectiveTest< + ObjectiveTestConfig, CRITERION::GINI>> + WeightedGiniObjectiveTestD; +TEST_P(WeightedGiniObjectiveTestD, giniObjectiveTest) {} +INSTANTIATE_TEST_CASE_P(RfTests, + WeightedGiniObjectiveTestD, + ::testing::ValuesIn(gini_objective_test_parameters)); +typedef ObjectiveTest< + ObjectiveTestConfig, CRITERION::GINI>> + WeightedGiniObjectiveTestF; +TEST_P(WeightedGiniObjectiveTestF, giniObjectiveTest) {} +INSTANTIATE_TEST_CASE_P(RfTests, + WeightedGiniObjectiveTestF, + ::testing::ValuesIn(gini_objective_test_parameters)); #ifndef NDEBUG // Feature sampling bias test From 41270ceff308f98e49a2d503230a577aa421ed82 Mon Sep 17 00:00:00 2001 From: Rory Mitchell Date: Wed, 10 Jun 2026 04:36:48 -0700 Subject: [PATCH 3/3] Handle RF weighted objective edge cases --- .../batched-levelalgo/objectives.cuh | 39 ++++++++++-- cpp/tests/sg/rf_test.cu | 59 ++++++++++++++++--- 2 files changed, 86 insertions(+), 12 deletions(-) diff --git a/cpp/src/decisiontree/batched-levelalgo/objectives.cuh b/cpp/src/decisiontree/batched-levelalgo/objectives.cuh index f7456ba8b3..5ec6d42cc7 100644 --- a/cpp/src/decisiontree/batched-levelalgo/objectives.cuh +++ b/cpp/src/decisiontree/batched-levelalgo/objectives.cuh @@ -56,6 +56,9 @@ class ClassificationObjectiveFunction { auto left_weight = WeightAt(hist, i, n_bins); auto right_weight = total_weight - left_weight; + if (total_weight <= 0.0 || left_weight <= 0.0 || right_weight <= 0.0) + return -std::numeric_limits::max(); + auto invLen = One / DataT(total_weight); auto invLeft = One / DataT(left_weight); auto invRight = One / DataT(right_weight); @@ -87,6 +90,9 @@ class ClassificationObjectiveFunction { auto left_weight = WeightAt(hist, i, n_bins); auto right_weight = total_weight - left_weight; + if (total_weight <= 0.0 || left_weight <= 0.0 || right_weight <= 0.0) + return -std::numeric_limits::max(); + auto gain{DataT(0.0)}; auto invLeft{DataT(1.0) / DataT(left_weight)}; auto invRight{DataT(1.0) / DataT(right_weight)}; @@ -121,6 +127,9 @@ class ClassificationObjectiveFunction { HDI DataT GainPerSplit(BinT const* hist, IdxT i, IdxT n_bins, IdxT len, IdxT nLeft, IdxT nRight) const { + if (nLeft < min_samples_leaf || nRight < min_samples_leaf) + return -std::numeric_limits::max(); + switch (criterion) { case CRITERION::GINI: return GiniGain(hist, i, n_bins, len, nLeft, nRight); case CRITERION::ENTROPY: return EntropyGain(hist, i, n_bins, len, nLeft, nRight); @@ -158,6 +167,12 @@ class ClassificationObjectiveFunction { for (int i = 0; i < nclasses; i++) { total += shist[i].Weight(); } + if (total <= 0.0) { + for (int i = 0; i < nclasses; i++) { + out[i] = DataT(0); + } + return; + } for (int i = 0; i < nclasses; i++) { out[i] = DataT(shist[i].Weight()) / total; } @@ -184,6 +199,9 @@ class RegressionObjectiveFunction { auto left_weight = hist[i].Weight(); auto right_weight = parent_weight - left_weight; + if (parent_weight <= 0.0 || left_weight <= 0.0 || right_weight <= 0.0) + return -std::numeric_limits::max(); + auto invLen = DataT(1.0) / DataT(parent_weight); auto label_sum = DataT(hist[n_bins - 1].LabelSum()); auto left_label_sum = DataT(hist[i].LabelSum()); @@ -203,13 +221,16 @@ class RegressionObjectiveFunction { auto left_weight = hist[i].Weight(); auto right_weight = parent_weight - left_weight; + if (parent_weight <= 0.0 || left_weight <= 0.0 || right_weight <= 0.0) + return -std::numeric_limits::max(); + auto invLen = DataT(1) / DataT(parent_weight); auto label_sum = DataT(hist[n_bins - 1].LabelSum()); auto left_label_sum = DataT(hist[i].LabelSum()); auto right_label_sum = DataT(hist[n_bins - 1].LabelSum() - hist[i].LabelSum()); // label sum cannot be non-positive - if (label_sum < eps_ || left_label_sum < eps_ || right_label_sum < eps_) + if (label_sum <= eps_ || left_label_sum <= eps_ || right_label_sum <= eps_) return -std::numeric_limits::max(); DataT parent_obj = -label_sum * raft::log(label_sum * invLen); @@ -227,13 +248,16 @@ class RegressionObjectiveFunction { auto left_weight = hist[i].Weight(); auto right_weight = parent_weight - left_weight; + if (parent_weight <= 0.0 || left_weight <= 0.0 || right_weight <= 0.0) + return -std::numeric_limits::max(); + auto invLen = DataT(1) / DataT(parent_weight); auto label_sum = DataT(hist[n_bins - 1].LabelSum()); auto left_label_sum = DataT(hist[i].LabelSum()); auto right_label_sum = DataT(hist[n_bins - 1].LabelSum() - hist[i].LabelSum()); // label sum cannot be non-positive - if (label_sum < eps_ || left_label_sum < eps_ || right_label_sum < eps_) + if (label_sum <= eps_ || left_label_sum <= eps_ || right_label_sum <= eps_) return -std::numeric_limits::max(); DataT parent_obj = DataT(parent_weight) * raft::log(label_sum * invLen); @@ -251,12 +275,15 @@ class RegressionObjectiveFunction { auto left_weight = hist[i].Weight(); auto right_weight = parent_weight - left_weight; + if (parent_weight <= 0.0 || left_weight <= 0.0 || right_weight <= 0.0) + return -std::numeric_limits::max(); + auto label_sum = DataT(hist[n_bins - 1].LabelSum()); auto left_label_sum = DataT(hist[i].LabelSum()); auto right_label_sum = DataT(hist[n_bins - 1].LabelSum() - hist[i].LabelSum()); // label sum cannot be non-positive - if (label_sum < eps_ || left_label_sum < eps_ || right_label_sum < eps_) + if (label_sum <= eps_ || left_label_sum <= eps_ || right_label_sum <= eps_) return -std::numeric_limits::max(); DataT parent_obj = -DataT(parent_weight) * DataT(parent_weight) / label_sum; @@ -272,6 +299,9 @@ class RegressionObjectiveFunction { HDI DataT GainPerSplit(BinT const* hist, IdxT i, IdxT n_bins, IdxT len, IdxT nLeft, IdxT nRight) const { + if (nLeft < min_samples_leaf || nRight < min_samples_leaf) + return -std::numeric_limits::max(); + switch (criterion) { case CRITERION::MSE: return MSEGain(hist, i, n_bins, len, nLeft, nRight); case CRITERION::POISSON: return PoissonGain(hist, i, n_bins, len, nLeft, nRight); @@ -308,7 +338,8 @@ class RegressionObjectiveFunction { static DI void SetLeafVector(BinT const* shist, int nclasses, DataT* out) { for (int i = 0; i < nclasses; i++) { - out[i] = shist[i].LabelSum() / shist[i].Weight(); + auto weight = shist[i].Weight(); + out[i] = weight > 0.0 ? shist[i].LabelSum() / weight : DataT(0); } } }; diff --git a/cpp/tests/sg/rf_test.cu b/cpp/tests/sg/rf_test.cu index b816c69ba6..2051d7c683 100644 --- a/cpp/tests/sg/rf_test.cu +++ b/cpp/tests/sg/rf_test.cu @@ -1194,12 +1194,14 @@ class ObjectiveTest : public ::testing::TestWithParam { auto GenHist(std::vector const& data, std::vector const& sample_weights) { std::vector cdf_hist, pdf_hist; - IdxT bin_width = raft::ceildiv(params.n_rows, params.max_n_bins); + auto bin_width = static_cast(raft::ceildiv(params.n_rows, params.max_n_bins)); for (auto c = 0; c < params.n_classes; ++c) { for (auto b = 0; b < params.max_n_bins; ++b) { - auto bin_begin = b * bin_width; - auto bin_end = bin_begin + bin_width; + auto bin_begin = + std::min(static_cast(b) * bin_width, data.size()); + auto bin_end = std::min(bin_begin + bin_width, data.size()); + auto bin_count = static_cast(bin_end - bin_begin); if constexpr (is_classification) { auto count{BinCountT(0)}; auto weight{DataT(0)}; @@ -1222,9 +1224,9 @@ class ObjectiveTest : public ::testing::TestWithParam { weight += sample_weights[i]; } if constexpr (is_weighted) { - pdf_hist.emplace_back(label_sum, bin_width, weight); + pdf_hist.emplace_back(label_sum, bin_count, weight); } else { - pdf_hist.emplace_back(label_sum, bin_width); + pdf_hist.emplace_back(label_sum, bin_count); } } @@ -1239,8 +1241,16 @@ class ObjectiveTest : public ::testing::TestWithParam { auto SplitOffset(std::size_t const split_bin_index) { - auto bin_width = raft::ceildiv(params.n_rows, params.max_n_bins); - return (split_bin_index + 1) * bin_width; + auto bin_width = static_cast(raft::ceildiv(params.n_rows, params.max_n_bins)); + return std::min((split_bin_index + 1) * bin_width, + static_cast(params.n_rows)); + } + + auto SplitBinIndexUpperBound() + { + auto bin_width = raft::ceildiv(params.n_rows, params.max_n_bins); + auto non_empty_n_bins = raft::ceildiv(params.n_rows, bin_width); + return std::max(1, non_empty_n_bins - 1); } auto MSE(std::vector const& data, @@ -1585,7 +1595,7 @@ class ObjectiveTest : public ::testing::TestWithParam { auto data = GenRandomData(); auto sample_weights = GenSampleWeights(); auto [cdf_hist, pdf_hist] = GenHist(data, sample_weights); - auto split_bin_index = RandUnder(params.max_n_bins - 1); + auto split_bin_index = RandUnder(SplitBinIndexUpperBound()); auto ground_truth_gain = GroundTruthGain(data, sample_weights, split_bin_index); auto len = NumLeftOfBin(cdf_hist, params.max_n_bins - 1); auto nLeft = NumLeftOfBin(cdf_hist, split_bin_index); @@ -1601,11 +1611,39 @@ class ObjectiveTest : public ::testing::TestWithParam { } }; +TEST(WeightedObjectiveEdgeCases, ClassificationRejectsZeroWeightChild) +{ + using ObjectiveT = ClassificationObjectiveFunction; + WeightedClassificationBin hist[]{{1, 0.0}, {1, 0.0}, {0, 0.0}, {1, 1.0}}; + CRITERION criteria[] = {CRITERION::GINI, CRITERION::ENTROPY}; + + for (auto criterion : criteria) { + ObjectiveT objective(2, 1, criterion); + auto gain = objective.GainPerSplit(hist, 0, 2, 2, 1, 1); + EXPECT_EQ(gain, -std::numeric_limits::max()); + } +} + +TEST(WeightedObjectiveEdgeCases, RegressionRejectsZeroWeightChild) +{ + using ObjectiveT = RegressionObjectiveFunction; + WeightedRegressionBin hist[]{{0.0, 1, 0.0}, {2.0, 2, 1.0}}; + CRITERION criteria[] = { + CRITERION::MSE, CRITERION::POISSON, CRITERION::GAMMA, CRITERION::INVERSE_GAUSSIAN}; + + for (auto criterion : criteria) { + ObjectiveT objective(1, 1, criterion); + auto gain = objective.GainPerSplit(hist, 0, 2, 2, 1, 1); + EXPECT_EQ(gain, -std::numeric_limits::max()); + } +} + const std::vector mse_objective_test_parameters = { {9507819643927052255LLU, 2048, 64, 1, 0, 0.00001}, {9507819643927052259LLU, 2048, 128, 1, 1, 0.00001}, {9507819643927052251LLU, 2048, 256, 1, 1, 0.00001}, {9507819643927052258LLU, 2048, 512, 1, 5, 0.00001}, + {9507819643927052260LLU, 2050, 128, 1, 1, 0.00001}, }; const std::vector poisson_objective_test_parameters = { @@ -1613,6 +1651,7 @@ const std::vector poisson_objective_test_parameters = { {9507819643927052259LLU, 2048, 128, 1, 1, 0.00001}, {9507819643927052251LLU, 2048, 256, 1, 1, 0.00001}, {9507819643927052258LLU, 2048, 512, 1, 5, 0.00001}, + {9507819643927052260LLU, 2050, 128, 1, 1, 0.00001}, }; const std::vector gamma_objective_test_parameters = { @@ -1620,6 +1659,7 @@ const std::vector gamma_objective_test_parameters = { {9507819643927052259LLU, 2048, 128, 1, 1, 0.00001}, {9507819643927052251LLU, 2048, 256, 1, 1, 0.00001}, {9507819643927052258LLU, 2048, 512, 1, 5, 0.00001}, + {9507819643927052260LLU, 2050, 128, 1, 1, 0.00001}, }; const std::vector invgauss_objective_test_parameters = { @@ -1627,6 +1667,7 @@ const std::vector invgauss_objective_test_parameters = {9507819643927052259LLU, 2048, 128, 1, 1, 0.00001}, {9507819643927052251LLU, 2048, 256, 1, 1, 0.00001}, {9507819643927052258LLU, 2048, 512, 1, 5, 0.00001}, + {9507819643927052260LLU, 2050, 128, 1, 1, 0.00001}, }; const std::vector entropy_objective_test_parameters = { @@ -1634,6 +1675,7 @@ const std::vector entropy_objective_test_parameters = { {9507819643927052256LLU, 2048, 128, 10, 1, 0.00001}, {9507819643927052257LLU, 2048, 256, 100, 1, 0.00001}, {9507819643927052258LLU, 2048, 512, 100, 5, 0.00001}, + {9507819643927052260LLU, 2050, 128, 10, 1, 0.00001}, }; const std::vector gini_objective_test_parameters = { @@ -1641,6 +1683,7 @@ const std::vector gini_objective_test_parameters = { {9507819643927052256LLU, 2048, 128, 10, 1, 0.00001}, {9507819643927052257LLU, 2048, 256, 100, 1, 0.00001}, {9507819643927052258LLU, 2048, 512, 100, 5, 0.00001}, + {9507819643927052260LLU, 2050, 128, 10, 1, 0.00001}, }; // mse objective test