diff --git a/onnxruntime/core/providers/cpu/ml/tree_ensemble_aggregator.h b/onnxruntime/core/providers/cpu/ml/tree_ensemble_aggregator.h index 3b7eb82294761..ebdac882bff09 100644 --- a/onnxruntime/core/providers/cpu/ml/tree_ensemble_aggregator.h +++ b/onnxruntime/core/providers/cpu/ml/tree_ensemble_aggregator.h @@ -480,6 +480,7 @@ class TreeAggregatorClassifier : public TreeAggregatorSum& class_labels_; bool binary_case_; + bool weights_are_all_positive_; int64_t positive_label_; int64_t negative_label_; @@ -490,11 +491,13 @@ class TreeAggregatorClassifier : public TreeAggregatorSum& base_values, const std::vector& class_labels, bool binary_case, + bool weights_are_all_positive, int64_t positive_label = 1, int64_t negative_label = 0) : TreeAggregatorSum(n_trees, n_targets_or_classes, post_transform, base_values), class_labels_(class_labels), binary_case_(binary_case), + weights_are_all_positive_(weights_are_all_positive), positive_label_(positive_label), negative_label_(negative_label) {} @@ -523,12 +526,22 @@ class TreeAggregatorClassifier : public TreeAggregatorSum 0) { - write_additional_scores = 2; - return class_labels_[1]; // positive label + if (weights_are_all_positive_) { + if (pos_weight > 0.5) { + write_additional_scores = 0; + return class_labels_[1]; // positive label + } else { + write_additional_scores = 1; + return class_labels_[0]; // negative label + } } else { - write_additional_scores = 3; - return class_labels_[0]; // negative label + if (pos_weight > 0) { + write_additional_scores = 2; + return class_labels_[1]; // positive label + } else { + write_additional_scores = 3; + return class_labels_[0]; // negative label + } } } return (pos_weight > 0) diff --git a/onnxruntime/core/providers/cpu/ml/tree_ensemble_common.h b/onnxruntime/core/providers/cpu/ml/tree_ensemble_common.h index 96a237f9dbf75..7ea9cc1edc478 100644 --- a/onnxruntime/core/providers/cpu/ml/tree_ensemble_common.h +++ b/onnxruntime/core/providers/cpu/ml/tree_ensemble_common.h @@ -983,6 +983,7 @@ TreeEnsembleCommon::ProcessTreeNodeLeave( template class TreeEnsembleCommonClassifier : public TreeEnsembleCommon { private: + bool weights_are_all_positive_; bool binary_case_; std::vector classlabels_strings_; std::vector classlabels_int64s_; @@ -1018,7 +1019,15 @@ Status TreeEnsembleCommonClassifier::Init( InlinedHashSet weights_classes; weights_classes.reserve(attributes.target_class_ids.size()); - weights_classes.insert(attributes.target_class_ids.begin(), attributes.target_class_ids.end()); + weights_are_all_positive_ = true; + for (size_t i = 0, end = attributes.target_class_ids.size(); i < end; ++i) { + weights_classes.insert(attributes.target_class_ids[i]); + if (weights_are_all_positive_ && (attributes.target_class_weights_as_tensor.empty() + ? static_cast(attributes.target_class_weights[i]) + : attributes.target_class_weights_as_tensor[i]) < 0) { + weights_are_all_positive_ = false; + } + } binary_case_ = this->n_targets_or_classes_ == 2 && weights_classes.size() == 1; if (!classlabels_strings_.empty()) { class_labels_.reserve(classlabels_strings_.size()); @@ -1039,7 +1048,8 @@ Status TreeEnsembleCommonClassifier::compu TreeAggregatorClassifier( this->roots_.size(), this->n_targets_or_classes_, this->post_transform_, this->base_values_, - classlabels_int64s_, binary_case_)); + classlabels_int64s_, binary_case_, + weights_are_all_positive_)); } else { int64_t N = X->Shape().NumDimensions() == 1 ? 1 : X->Shape()[0]; AllocatorPtr alloc; @@ -1050,7 +1060,8 @@ Status TreeEnsembleCommonClassifier::compu TreeAggregatorClassifier( this->roots_.size(), this->n_targets_or_classes_, this->post_transform_, this->base_values_, - class_labels_, binary_case_)); + class_labels_, binary_case_, + weights_are_all_positive_)); const int64_t* plabel = label_int64.Data(); std::string* labels = label->MutableData(); for (size_t i = 0; i < (size_t)N; ++i) diff --git a/onnxruntime/test/providers/cpu/ml/tree_ensembler_classifier_test.cc b/onnxruntime/test/providers/cpu/ml/tree_ensembler_classifier_test.cc index df97056284ae7..4a56579bcfca1 100644 --- a/onnxruntime/test/providers/cpu/ml/tree_ensembler_classifier_test.cc +++ b/onnxruntime/test/providers/cpu/ml/tree_ensembler_classifier_test.cc @@ -709,5 +709,115 @@ TEST(MLOpTest, TreeEnsembleNegativeFeatureIds) { test.Run(OpTester::ExpectResult::kExpectFailure, "nodes_featureids[0]=-2 must be in [0, 2147483647] for non-leaf nodes"); } +// Regression test for PR #27552 which inadvertently broke binary TreeEnsembleClassifier models +// with all-positive weights and non-LOGISTIC post_transform (e.g., sklearn RandomForestClassifier). +// When weights are all non-negative (probability-like values in [0,1]), the threshold for the +// positive class should be > 0.5 and the complement should be computed as 1-score (not -score). +// Without the weights_are_all_positive_ flag, scores in (0, 0.5) would produce negative +// "probabilities" and incorrect labels. +TEST(MLOpTest, TreeEnsembleClassifierBinaryAllPositiveWeightsNone) { + // 3 trees, each with a root branching on feature 0 at threshold 0.5. + // Left leaf (feature <= 0.5): weight = 0.2 (low probability for class 1) + // Right leaf (feature > 0.5): weight = 0.8 (high probability for class 1) + // + // For input x=0.0: all 3 trees go left, avg score = 0.2 (< 0.5 → label 0) + // For input x=1.0: all 3 trees go right, avg score = 0.8 (> 0.5 → label 1) + // For input x=0.0 with 1 tree going right: score = (0.2+0.2+0.8)/3 ≈ 0.4 (< 0.5 → label 0) + // + // Binary case: classlabels_int64s={0,1}, class_ids all 0, all weights in [0,1]. + // post_transform = NONE (default). + + // Single tree for simplicity: root → left leaf (weight 0.3) or right leaf (weight 0.8) + std::vector treeids = {0, 0, 0}; + std::vector nodeids = {0, 1, 2}; + std::vector featureids = {0, -2, -2}; + std::vector thresholds = {0.5f, -2.f, -2.f}; + std::vector modes = {"BRANCH_LEQ", "LEAF", "LEAF"}; + std::vector truenodeids = {1, -1, -1}; + std::vector falsenodeids = {2, -1, -1}; + + // Both leaves target class 0 (single-column binary encoding). + // All weights are positive — this triggers weights_are_all_positive_=true. + std::vector class_treeids = {0, 0}; + std::vector class_nodeids = {1, 2}; + std::vector class_classids = {0, 0}; + std::vector class_weights = {0.3f, 0.8f}; + std::vector classes = {0, 1}; + + // Test 1: score = 0.3 (< 0.5) → label should be 0, probabilities should be [0.7, 0.3] + { + OpTester test("TreeEnsembleClassifier", 1, onnxruntime::kMLDomain); + test.AddAttribute("nodes_treeids", treeids); + test.AddAttribute("nodes_nodeids", nodeids); + test.AddAttribute("nodes_featureids", featureids); + test.AddAttribute("nodes_values", thresholds); + test.AddAttribute("nodes_modes", modes); + test.AddAttribute("nodes_truenodeids", truenodeids); + test.AddAttribute("nodes_falsenodeids", falsenodeids); + test.AddAttribute("class_treeids", class_treeids); + test.AddAttribute("class_nodeids", class_nodeids); + test.AddAttribute("class_ids", class_classids); + test.AddAttribute("class_weights", class_weights); + test.AddAttribute("classlabels_int64s", classes); + + // x=0.0 → goes left (0.0 <= 0.5) → leaf 1, weight = 0.3 + test.AddInput("X", {1, 1}, {0.0f}); + test.AddOutput("Y", {1}, {0}); + // write_additional_scores=1: complement = 1-score → [1-0.3, 0.3] = [0.7, 0.3] + test.AddOutput("Z", {1, 2}, {0.7f, 0.3f}); + test.Run(); + } + + // Test 2: score = 0.8 (> 0.5) → label should be 1, probabilities should be [0.2, 0.8] + { + OpTester test("TreeEnsembleClassifier", 1, onnxruntime::kMLDomain); + test.AddAttribute("nodes_treeids", treeids); + test.AddAttribute("nodes_nodeids", nodeids); + test.AddAttribute("nodes_featureids", featureids); + test.AddAttribute("nodes_values", thresholds); + test.AddAttribute("nodes_modes", modes); + test.AddAttribute("nodes_truenodeids", truenodeids); + test.AddAttribute("nodes_falsenodeids", falsenodeids); + test.AddAttribute("class_treeids", class_treeids); + test.AddAttribute("class_nodeids", class_nodeids); + test.AddAttribute("class_ids", class_classids); + test.AddAttribute("class_weights", class_weights); + test.AddAttribute("classlabels_int64s", classes); + + // x=1.0 → goes right (1.0 > 0.5) → leaf 2, weight = 0.8 + test.AddInput("X", {1, 1}, {1.0f}); + test.AddOutput("Y", {1}, {1}); + // write_additional_scores=0: complement = 1-score → [1-0.8, 0.8] = [0.2, 0.8] + test.AddOutput("Z", {1, 2}, {0.2f, 0.8f}); + test.Run(); + } + + // Test 3: score = 0.5 (boundary, <= 0.5) → label should be 0 + { + OpTester test("TreeEnsembleClassifier", 1, onnxruntime::kMLDomain); + // Use weight = 0.5 for the left leaf to test the boundary + std::vector boundary_weights = {0.5f, 0.8f}; + test.AddAttribute("nodes_treeids", treeids); + test.AddAttribute("nodes_nodeids", nodeids); + test.AddAttribute("nodes_featureids", featureids); + test.AddAttribute("nodes_values", thresholds); + test.AddAttribute("nodes_modes", modes); + test.AddAttribute("nodes_truenodeids", truenodeids); + test.AddAttribute("nodes_falsenodeids", falsenodeids); + test.AddAttribute("class_treeids", class_treeids); + test.AddAttribute("class_nodeids", class_nodeids); + test.AddAttribute("class_ids", class_classids); + test.AddAttribute("class_weights", boundary_weights); + test.AddAttribute("classlabels_int64s", classes); + + // x=0.0 → leaf 1, weight = 0.5 (not > 0.5, so label = 0) + test.AddInput("X", {1, 1}, {0.0f}); + test.AddOutput("Y", {1}, {0}); + // write_additional_scores=1: [1-0.5, 0.5] = [0.5, 0.5] + test.AddOutput("Z", {1, 2}, {0.5f, 0.5f}); + test.Run(); + } +} + } // namespace test } // namespace onnxruntime