Skip to content
Merged
Show file tree
Hide file tree
Changes from all commits
Commits
File filter

Filter by extension

Filter by extension

Conversations
Failed to load comments.
Loading
Jump to
Jump to file
Failed to load files.
Loading
Diff view
Diff view
23 changes: 18 additions & 5 deletions onnxruntime/core/providers/cpu/ml/tree_ensemble_aggregator.h
Original file line number Diff line number Diff line change
Expand Up @@ -480,6 +480,7 @@ class TreeAggregatorClassifier : public TreeAggregatorSum<InputType, ThresholdTy
private:
const std::vector<int64_t>& class_labels_;
bool binary_case_;
bool weights_are_all_positive_;
int64_t positive_label_;
int64_t negative_label_;

Expand All @@ -490,11 +491,13 @@ class TreeAggregatorClassifier : public TreeAggregatorSum<InputType, ThresholdTy
const std::vector<ThresholdType>& base_values,
const std::vector<int64_t>& class_labels,
bool binary_case,
bool weights_are_all_positive,
int64_t positive_label = 1,
int64_t negative_label = 0) : TreeAggregatorSum<InputType, ThresholdType, OutputType>(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) {}

Expand Down Expand Up @@ -523,12 +526,22 @@ class TreeAggregatorClassifier : public TreeAggregatorSum<InputType, ThresholdTy
ThresholdType score1, unsigned char has_score1) const {
ThresholdType pos_weight = has_score1 ? score1 : (has_score0 ? score0 : 0); // only 1 class
if (binary_case_) {
if (pos_weight > 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)
Expand Down
17 changes: 14 additions & 3 deletions onnxruntime/core/providers/cpu/ml/tree_ensemble_common.h
Original file line number Diff line number Diff line change
Expand Up @@ -983,6 +983,7 @@ TreeEnsembleCommon<InputType, ThresholdType, OutputType>::ProcessTreeNodeLeave(
template <typename InputType, typename ThresholdType, typename OutputType>
class TreeEnsembleCommonClassifier : public TreeEnsembleCommon<InputType, ThresholdType, OutputType> {
private:
bool weights_are_all_positive_;
bool binary_case_;
std::vector<std::string> classlabels_strings_;
std::vector<int64_t> classlabels_int64s_;
Expand Down Expand Up @@ -1018,7 +1019,15 @@ Status TreeEnsembleCommonClassifier<InputType, ThresholdType, OutputType>::Init(

InlinedHashSet<int64_t> 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<ThresholdType>(attributes.target_class_weights[i])
: attributes.target_class_weights_as_tensor[i]) < 0) {
Comment thread
tianleiwu marked this conversation as resolved.
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());
Expand All @@ -1039,7 +1048,8 @@ Status TreeEnsembleCommonClassifier<InputType, ThresholdType, OutputType>::compu
TreeAggregatorClassifier<InputType, ThresholdType, OutputType>(
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;
Expand All @@ -1050,7 +1060,8 @@ Status TreeEnsembleCommonClassifier<InputType, ThresholdType, OutputType>::compu
TreeAggregatorClassifier<InputType, ThresholdType, OutputType>(
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<int64_t>();
std::string* labels = label->MutableData<std::string>();
for (size_t i = 0; i < (size_t)N; ++i)
Expand Down
110 changes: 110 additions & 0 deletions onnxruntime/test/providers/cpu/ml/tree_ensembler_classifier_test.cc
Original file line number Diff line number Diff line change
Expand Up @@ -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<int64_t> treeids = {0, 0, 0};
std::vector<int64_t> nodeids = {0, 1, 2};
std::vector<int64_t> featureids = {0, -2, -2};
std::vector<float> thresholds = {0.5f, -2.f, -2.f};
std::vector<std::string> modes = {"BRANCH_LEQ", "LEAF", "LEAF"};
std::vector<int64_t> truenodeids = {1, -1, -1};
std::vector<int64_t> 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<int64_t> class_treeids = {0, 0};
std::vector<int64_t> class_nodeids = {1, 2};
std::vector<int64_t> class_classids = {0, 0};
std::vector<float> class_weights = {0.3f, 0.8f};
std::vector<int64_t> 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<float>("X", {1, 1}, {0.0f});
test.AddOutput<int64_t>("Y", {1}, {0});
// write_additional_scores=1: complement = 1-score → [1-0.3, 0.3] = [0.7, 0.3]
test.AddOutput<float>("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<float>("X", {1, 1}, {1.0f});
test.AddOutput<int64_t>("Y", {1}, {1});
// write_additional_scores=0: complement = 1-score → [1-0.8, 0.8] = [0.2, 0.8]
test.AddOutput<float>("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<float> 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<float>("X", {1, 1}, {0.0f});
test.AddOutput<int64_t>("Y", {1}, {0});
// write_additional_scores=1: [1-0.5, 0.5] = [0.5, 0.5]
test.AddOutput<float>("Z", {1, 2}, {0.5f, 0.5f});
test.Run();
}
}

} // namespace test
} // namespace onnxruntime
Loading