Removes omp for ThreadPool in TreeEsemble* (#3596)

* Removes omp to use ThreadPool

* removes unnecessary old OMP code

* rename compute_agg, use ThreadPool::NumThreads

Co-authored-by: xavier dupré <xavier.dupre@gmail.com>
This commit is contained in:
Xavier Dupré 2020-04-23 08:48:31 +02:00 committed by GitHub
parent 02bae6bd06
commit 5777fc18c3
No known key found for this signature in database
GPG key ID: 4AEE18F83AFDEB23
3 changed files with 79 additions and 75 deletions

View file

@ -174,7 +174,7 @@ common::Status TreeEnsembleClassifier<T>::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();
}

View file

@ -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<int64_t>& target_class_treeids,
const std::vector<OTYPE>& 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<OTYPE>* ProcessTreeNodeLeave(
TreeNodeElement<OTYPE>* root, const ITYPE* x_data) const;
template <typename AGG>
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 <typename ITYPE, typename OTYPE>
@ -213,32 +215,32 @@ TreeEnsembleCommon<ITYPE, OTYPE>::TreeEnsembleCommon(int parallel_tree, int para
}
template <typename ITYPE, typename OTYPE>
void TreeEnsembleCommon<ITYPE, OTYPE>::compute(const Tensor* X, Tensor* Z, Tensor* label) const {
void TreeEnsembleCommon<ITYPE, OTYPE>::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<ITYPE, OTYPE>(
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<ITYPE, OTYPE>(
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<ITYPE, OTYPE>(
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<ITYPE, OTYPE>(
roots_.size(), n_targets_or_classes_,
post_transform_, base_values_));
@ -250,7 +252,7 @@ void TreeEnsembleCommon<ITYPE, OTYPE>::compute(const Tensor* X, Tensor* Z, Tenso
template <typename ITYPE, typename OTYPE>
template <typename AGG>
void TreeEnsembleCommon<ITYPE, OTYPE>::compute_agg(const Tensor* X, Tensor* Z, Tensor* label, const AGG& agg) const {
void TreeEnsembleCommon<ITYPE, OTYPE>::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<ITYPE, OTYPE>::compute_agg(const Tensor* X, Tensor* Z, T
agg.ProcessTreeNodePrediction1(score, *ProcessTreeNodeLeave(roots_[j], x_data));
} else {
std::vector<ScoreValue<OTYPE>> 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<int32_t>(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<ITYPE, OTYPE>::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<OTYPE> score = {0, 0};
for (size_t j = 0; j < static_cast<size_t>(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<int32_t>(N),
[&](ptrdiff_t i) {
ScoreValue<OTYPE> score = {0, 0};
for (size_t j = 0; j < static_cast<size_t>(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<ITYPE, OTYPE>::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<ScoreValue<OTYPE>> 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<int32_t>(nth),
[&](ptrdiff_t th) {
std::vector<ScoreValue<OTYPE>> 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<ITYPE, OTYPE>::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<ScoreValue<OTYPE>> 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<OTYPE>({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<int32_t>(nth),
[&](ptrdiff_t th) {
size_t j;
std::vector<ScoreValue<OTYPE>> 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<OTYPE>({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<ITYPE, OTYPE> {
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 <typename ITYPE, typename OTYPE>
@ -571,10 +575,10 @@ TreeEnsembleCommonClassifier<ITYPE, OTYPE>::TreeEnsembleCommonClassifier(
}
template <typename ITYPE, typename OTYPE>
void TreeEnsembleCommonClassifier<ITYPE, OTYPE>::compute(const Tensor* X, Tensor* Z, Tensor* label) const {
void TreeEnsembleCommonClassifier<ITYPE, OTYPE>::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<ITYPE, OTYPE>(
this->roots_.size(), this->n_targets_or_classes_,
this->post_transform_, this->base_values_,
@ -584,8 +588,8 @@ void TreeEnsembleCommonClassifier<ITYPE, OTYPE>::compute(const Tensor* X, Tensor
int64_t N = X->Shape().NumDimensions() == 1 ? 1 : X->Shape()[0];
std::shared_ptr<IAllocator> allocator = std::make_shared<CPUAllocator>();
Tensor label_int64(DataTypeImpl::GetType<int64_t>(), TensorShape({N}), allocator);
this->compute_agg(
X, Z, &label_int64,
this->ComputeAgg(
ttp, X, Z, &label_int64,
TreeAggregatorClassifier<ITYPE, OTYPE>(
this->roots_.size(), this->n_targets_or_classes_,
this->post_transform_, this->base_values_,

View file

@ -56,10 +56,10 @@ common::Status TreeEnsembleRegressor<T>::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