diff --git a/onnxruntime/core/providers/cpu/ml/tree_ensemble_classifier.cc b/onnxruntime/core/providers/cpu/ml/tree_ensemble_classifier.cc index 1c418e7302..8960c40a01 100644 --- a/onnxruntime/core/providers/cpu/ml/tree_ensemble_classifier.cc +++ b/onnxruntime/core/providers/cpu/ml/tree_ensemble_classifier.cc @@ -174,7 +174,7 @@ common::Status TreeEnsembleClassifier::Compute(OpKernelContext* context) cons Tensor* Y = context->Output(0, TensorShape({N})); Tensor* Z = context->Output(1, TensorShape({N, tree_ensemble_.get_class_count()})); - tree_ensemble_.compute(&X, Z, Y); + tree_ensemble_.compute(context->GetOperatorThreadPool(), &X, Z, Y); return Status::OK(); } diff --git a/onnxruntime/core/providers/cpu/ml/tree_ensemble_common.h b/onnxruntime/core/providers/cpu/ml/tree_ensemble_common.h index e3f844e236..a88953832f 100644 --- a/onnxruntime/core/providers/cpu/ml/tree_ensemble_common.h +++ b/onnxruntime/core/providers/cpu/ml/tree_ensemble_common.h @@ -2,7 +2,9 @@ // Licensed under the MIT License. #pragma once + #include "tree_ensemble_aggregator.h" +#include "core/platform/threadpool.h" namespace onnxruntime { namespace ml { @@ -49,14 +51,14 @@ class TreeEnsembleCommon { const std::vector& target_class_treeids, const std::vector& target_class_weights); - void compute(const Tensor* X, Tensor* Z, Tensor* label) const; + void compute(concurrency::ThreadPool* ttp, const Tensor* X, Tensor* Z, Tensor* label) const; protected: TreeNodeElement* ProcessTreeNodeLeave( TreeNodeElement* root, const ITYPE* x_data) const; template - void compute_agg(const Tensor* X, Tensor* Z, Tensor* label, const AGG& agg) const; + void ComputeAgg(concurrency::ThreadPool* ttp, const Tensor* X, Tensor* Z, Tensor* label, const AGG& agg) const; }; template @@ -213,32 +215,32 @@ TreeEnsembleCommon::TreeEnsembleCommon(int parallel_tree, int para } template -void TreeEnsembleCommon::compute(const Tensor* X, Tensor* Z, Tensor* label) const { +void TreeEnsembleCommon::compute(concurrency::ThreadPool* ttp, const Tensor* X, Tensor* Z, Tensor* label) const { switch (aggregate_function_) { case AGGREGATE_FUNCTION::AVERAGE: - compute_agg( - X, Z, label, + ComputeAgg( + ttp, X, Z, label, TreeAggregatorAverage( roots_.size(), n_targets_or_classes_, post_transform_, base_values_)); return; case AGGREGATE_FUNCTION::SUM: - compute_agg( - X, Z, label, + ComputeAgg( + ttp, X, Z, label, TreeAggregatorSum( roots_.size(), n_targets_or_classes_, post_transform_, base_values_)); return; case AGGREGATE_FUNCTION::MIN: - compute_agg( - X, Z, label, + ComputeAgg( + ttp, X, Z, label, TreeAggregatorMin( roots_.size(), n_targets_or_classes_, post_transform_, base_values_)); return; case AGGREGATE_FUNCTION::MAX: - compute_agg( - X, Z, label, + ComputeAgg( + ttp, X, Z, label, TreeAggregatorMax( roots_.size(), n_targets_or_classes_, post_transform_, base_values_)); @@ -250,7 +252,7 @@ void TreeEnsembleCommon::compute(const Tensor* X, Tensor* Z, Tenso template template -void TreeEnsembleCommon::compute_agg(const Tensor* X, Tensor* Z, Tensor* label, const AGG& agg) const { +void TreeEnsembleCommon::ComputeAgg(concurrency::ThreadPool* ttp, const Tensor* X, Tensor* Z, Tensor* label, const AGG& agg) const { int64_t stride = X->Shape().NumDimensions() == 1 ? X->Shape()[0] : X->Shape()[1]; int64_t N = X->Shape().NumDimensions() == 1 ? 1 : X->Shape()[0]; @@ -266,12 +268,13 @@ void TreeEnsembleCommon::compute_agg(const Tensor* X, Tensor* Z, T agg.ProcessTreeNodePrediction1(score, *ProcessTreeNodeLeave(roots_[j], x_data)); } else { std::vector> scores_t(n_trees_, {0, 0}); -#ifdef _OPENMP -#pragma omp parallel for -#endif - for (int64_t j = 0; j < n_trees_; ++j) { - agg.ProcessTreeNodePrediction1(scores_t[j], *ProcessTreeNodeLeave(roots_[j], x_data)); - } + concurrency::ThreadPool::TryBatchParallelFor( + ttp, + static_cast(this->n_trees_), + [&](ptrdiff_t j) { + agg.ProcessTreeNodePrediction1(scores_t[j], *ProcessTreeNodeLeave(this->roots_[j], x_data)); + }, + 0); for (auto it = scores_t.cbegin(); it != scores_t.cend(); ++it) agg.MergePrediction1(score, *it); } @@ -290,16 +293,17 @@ void TreeEnsembleCommon::compute_agg(const Tensor* X, Tensor* Z, T label_data == NULL ? NULL : (label_data + i)); } } else { -#ifdef _OPENMP -#pragma omp parallel for -#endif - for (int64_t i = 0; i < N; ++i) { - ScoreValue score = {0, 0}; - for (size_t j = 0; j < static_cast(n_trees_); ++j) - agg.ProcessTreeNodePrediction1(score, *ProcessTreeNodeLeave(roots_[j], x_data + i * stride)); - agg.FinalizeScores1(z_data + i * n_targets_or_classes_, score, - label_data == NULL ? NULL : (label_data + i)); - } + concurrency::ThreadPool::TryBatchParallelFor( + ttp, + static_cast(N), + [&](ptrdiff_t i) { + ScoreValue score = {0, 0}; + for (size_t j = 0; j < static_cast(n_trees_); ++j) + agg.ProcessTreeNodePrediction1(score, *ProcessTreeNodeLeave(roots_[j], x_data + i * stride)); + agg.FinalizeScores1(z_data + i * n_targets_or_classes_, score, + label_data == NULL ? NULL : (label_data + i)); + }, + 0); } } } else { @@ -311,24 +315,22 @@ void TreeEnsembleCommon::compute_agg(const Tensor* X, Tensor* Z, T agg.ProcessTreeNodePrediction(scores, *ProcessTreeNodeLeave(roots_[j], x_data)); agg.FinalizeScores(scores, z_data, -1, label_data); } else { -#ifdef _OPENMP -#pragma omp parallel -#endif - { - std::vector> private_scores(n_targets_or_classes_, {0, 0}); -#ifdef _OPENMP -#pragma omp for -#endif - for (int64_t j = 0; j < n_trees_; ++j) { - agg.ProcessTreeNodePrediction(private_scores, *ProcessTreeNodeLeave(roots_[j], x_data)); - } - -#ifdef _OPENMP -#pragma omp critical -#endif - agg.MergePrediction(scores, private_scores); - } - + auto nth = n_trees_ * 2 / concurrency::ThreadPool::NumThreads(ttp); + if (n_trees_ % nth != 0) + ++nth; + concurrency::ThreadPool::TryBatchParallelFor( + ttp, + static_cast(nth), + [&](ptrdiff_t th) { + std::vector> private_scores(n_targets_or_classes_, {0, 0}); + auto end = nth * (th + 1); + end = end < N ? end : N; + for (int64_t j = nth * th; j < end; ++j) { + agg.ProcessTreeNodePrediction(private_scores, *ProcessTreeNodeLeave(roots_[j], x_data)); + } + agg.MergePrediction(scores, private_scores); + }, + 0); agg.FinalizeScores(scores, z_data, -1, label_data); } } else { @@ -346,29 +348,31 @@ void TreeEnsembleCommon::compute_agg(const Tensor* X, Tensor* Z, T ORT_ENFORCE((int64_t)scores.size() == n_targets_or_classes_); } } else { -#ifdef _OPENMP -#pragma omp parallel -#endif - { - std::vector> scores(n_targets_or_classes_); - size_t j; - -#ifdef _OPENMP -#pragma omp for -#endif - for (int64_t i = 0; i < N; ++i) { - std::fill(scores.begin(), scores.end(), ScoreValue({0, 0})); - for (j = 0; j < roots_.size(); ++j) - agg.ProcessTreeNodePrediction(scores, *ProcessTreeNodeLeave(roots_[j], x_data + i * stride)); - agg.FinalizeScores(scores, - z_data + i * n_targets_or_classes_, -1, - label_data == NULL ? NULL : (label_data + i)); - } - } + auto nth = N * 2 / concurrency::ThreadPool::NumThreads(ttp); + if (N % nth != 0) + ++nth; + concurrency::ThreadPool::TryBatchParallelFor( + ttp, + static_cast(nth), + [&](ptrdiff_t th) { + size_t j; + std::vector> scores(n_targets_or_classes_); + auto end = nth * (th + 1); + end = end < N ? end : N; + for (int64_t i = nth * th; i < end; ++i) { + std::fill(scores.begin(), scores.end(), ScoreValue({0, 0})); + for (j = 0; j < roots_.size(); ++j) + agg.ProcessTreeNodePrediction(scores, *ProcessTreeNodeLeave(roots_[j], x_data + i * stride)); + agg.FinalizeScores(scores, + z_data + i * n_targets_or_classes_, -1, + label_data == NULL ? NULL : (label_data + i)); + } + }, + 0); } } } -} +} // namespace detail #define TREE_FIND_VALUE(CMP) \ if (has_missing_tracks_) { \ @@ -509,7 +513,7 @@ class TreeEnsembleCommonClassifier : TreeEnsembleCommon { int64_t get_class_count() const { return this->n_targets_or_classes_; } - void compute(const Tensor* X, Tensor* Z, Tensor* label) const; + void compute(concurrency::ThreadPool* ttp, const Tensor* X, Tensor* Z, Tensor* label) const; }; template @@ -571,10 +575,10 @@ TreeEnsembleCommonClassifier::TreeEnsembleCommonClassifier( } template -void TreeEnsembleCommonClassifier::compute(const Tensor* X, Tensor* Z, Tensor* label) const { +void TreeEnsembleCommonClassifier::compute(concurrency::ThreadPool* ttp, const Tensor* X, Tensor* Z, Tensor* label) const { if (classlabels_strings_.size() == 0) { - this->compute_agg( - X, Z, label, + this->ComputeAgg( + ttp, X, Z, label, TreeAggregatorClassifier( this->roots_.size(), this->n_targets_or_classes_, this->post_transform_, this->base_values_, @@ -584,8 +588,8 @@ void TreeEnsembleCommonClassifier::compute(const Tensor* X, Tensor int64_t N = X->Shape().NumDimensions() == 1 ? 1 : X->Shape()[0]; std::shared_ptr allocator = std::make_shared(); Tensor label_int64(DataTypeImpl::GetType(), TensorShape({N}), allocator); - this->compute_agg( - X, Z, &label_int64, + this->ComputeAgg( + ttp, X, Z, &label_int64, TreeAggregatorClassifier( this->roots_.size(), this->n_targets_or_classes_, this->post_transform_, this->base_values_, diff --git a/onnxruntime/core/providers/cpu/ml/treeregressor.cc b/onnxruntime/core/providers/cpu/ml/treeregressor.cc index 55aba6f8c9..96a95e5a1e 100644 --- a/onnxruntime/core/providers/cpu/ml/treeregressor.cc +++ b/onnxruntime/core/providers/cpu/ml/treeregressor.cc @@ -56,10 +56,10 @@ common::Status TreeEnsembleRegressor::Compute(OpKernelContext* context) const int64_t N = X->Shape().NumDimensions() == 1 ? 1 : X->Shape()[0]; Tensor* Y = context->Output(0, TensorShape({N, tree_ensemble_.n_targets_or_classes_})); - tree_ensemble_.compute(X, Y, NULL); + tree_ensemble_.compute(context->GetOperatorThreadPool(), X, Y, NULL); return Status::OK(); } -} // namespace onnxruntime +} // namespace ml } // namespace onnxruntime