From ab5f79f52230ecc14dc10a899af303efe93fe7a4 Mon Sep 17 00:00:00 2001 From: zetongzhao Date: Mon, 14 Mar 2022 16:41:56 -0400 Subject: [PATCH] move common helper function --- .../ut/cpp/dataset/c_api_dataset_ops_test.cc | 64 ++----------------- tests/ut/cpp/dataset/common/common.cc | 39 +++++++++++ tests/ut/cpp/dataset/common/common.h | 17 +++++ .../skip_pushdown_optimization_pass_test.cc | 7 -- 4 files changed, 61 insertions(+), 66 deletions(-) diff --git a/tests/ut/cpp/dataset/c_api_dataset_ops_test.cc b/tests/ut/cpp/dataset/c_api_dataset_ops_test.cc index 2371ef058e8..83cd0dedd20 100644 --- a/tests/ut/cpp/dataset/c_api_dataset_ops_test.cc +++ b/tests/ut/cpp/dataset/c_api_dataset_ops_test.cc @@ -16,7 +16,6 @@ #include "common/common.h" #include "include/api/types.h" #include "minddata/dataset/core/tensor_row.h" -#include "minddata/dataset/engine/ir/datasetops/dataset_node.h" #include "minddata/dataset/include/dataset/datasets.h" #include "minddata/dataset/include/dataset/vision.h" #include "minddata/dataset/kernels/ir/data/transforms_ir.h" @@ -142,68 +141,15 @@ class MindDataTestPipeline : public UT::DatasetOpTesting { protected: }; -TensorRow VecToRow(const MSTensorVec &v) { - TensorRow row; - for (const mindspore::MSTensor &t : v) { - std::shared_ptr rt; - (void)Tensor::CreateFromMemory(TensorShape(t.Shape()), MSTypeToDEType(static_cast(t.DataType())), - (const uchar *)(t.Data().get()), t.DataSize(), &rt); - row.emplace_back(rt); - } - return row; -} -MSTensorVec RowToVec(const TensorRow &v) { - MSTensorVec rv; // std::make_shared(de_tensor) - std::transform(v.begin(), v.end(), std::back_inserter(rv), [](std::shared_ptr t) -> mindspore::MSTensor { - return mindspore::MSTensor(std::make_shared(t)); - }); - return rv; -} - MSTensorVec BucketBatchTestFunction(MSTensorVec input) { - mindspore::dataset::TensorRow output; - std::shared_ptr out; - (void)Tensor::CreateEmpty(mindspore::dataset::TensorShape({1}), - mindspore::dataset::DataType(mindspore::dataset::DataType::Type::DE_INT32), &out); - (void)out->SetItemAt({0}, 2); - output.push_back(out); - return RowToVec(output); -} - -MSTensorVec Predicate1(MSTensorVec in) { - // Return true if input is equal to 3 - uint64_t input_value; - TensorRow input = VecToRow(in); - (void)input.at(0)->GetItemAt(&input_value, {0}); - bool result = (input_value == 3); - - // Convert from boolean to TensorRow TensorRow output; std::shared_ptr out; - (void)Tensor::CreateEmpty(mindspore::dataset::TensorShape({}), - mindspore::dataset::DataType(mindspore::dataset::DataType::Type::DE_BOOL), &out); - (void)out->SetItemAt({}, result); + (void)Tensor::CreateEmpty( + TensorShape({1}), DataType(DataType::Type::DE_INT32), + &out); + constexpr int value = 2; + (void)out->SetItemAt({0}, value); output.push_back(out); - - return RowToVec(output); -} - -MSTensorVec Predicate2(MSTensorVec in) { - // Return true if label is more than 1 - // The index of label in input is 1 - uint64_t input_value; - TensorRow input = VecToRow(in); - (void)input.at(1)->GetItemAt(&input_value, {0}); - bool result = (input_value > 1); - - // Convert from boolean to TensorRow - TensorRow output; - std::shared_ptr out; - (void)Tensor::CreateEmpty(mindspore::dataset::TensorShape({}), - mindspore::dataset::DataType(mindspore::dataset::DataType::Type::DE_BOOL), &out); - (void)out->SetItemAt({}, result); - output.push_back(out); - return RowToVec(output); } diff --git a/tests/ut/cpp/dataset/common/common.cc b/tests/ut/cpp/dataset/common/common.cc index 1d224869430..5bd07aeccef 100644 --- a/tests/ut/cpp/dataset/common/common.cc +++ b/tests/ut/cpp/dataset/common/common.cc @@ -150,3 +150,42 @@ std::shared_ptr DatasetOpTesting::Build( #endif #endif } // namespace UT + +namespace mindspore { +namespace dataset { +MSTensorVec Predicate1(MSTensorVec in) { + // Return true if input is equal to 3 + uint64_t input_value; + TensorRow input = VecToRow(in); + (void)input.at(0)->GetItemAt(&input_value, {0}); + bool result = (input_value == 3); + + // Convert from boolean to TensorRow + TensorRow output; + std::shared_ptr out; + (void)Tensor::CreateEmpty(TensorShape({}), DataType(DataType::Type::DE_BOOL), &out); + (void)out->SetItemAt({}, result); + output.push_back(out); + + return RowToVec(output); +} + +MSTensorVec Predicate2(MSTensorVec in) { + // Return true if label is more than 1 + // The index of label in input is 1 + uint64_t input_value; + TensorRow input = VecToRow(in); + (void)input.at(1)->GetItemAt(&input_value, {0}); + bool result = (input_value > 1); + + // Convert from boolean to TensorRow + TensorRow output; + std::shared_ptr out; + (void)Tensor::CreateEmpty(TensorShape({}), DataType(mindspore::dataset::DataType::Type::DE_BOOL), &out); + (void)out->SetItemAt({}, result); + output.push_back(out); + + return RowToVec(output); +} +} // namespace dataset +} // namespace mindspore \ No newline at end of file diff --git a/tests/ut/cpp/dataset/common/common.h b/tests/ut/cpp/dataset/common/common.h index ff925f7fdd6..6e9a0c0478e 100644 --- a/tests/ut/cpp/dataset/common/common.h +++ b/tests/ut/cpp/dataset/common/common.h @@ -27,6 +27,7 @@ #include "minddata/dataset/engine/datasetops/batch_op.h" #include "minddata/dataset/engine/datasetops/repeat_op.h" #include "minddata/dataset/engine/datasetops/source/tf_reader_op.h" +#include "minddata/dataset/engine/ir/datasetops/dataset_node.h" using mindspore::Status; using mindspore::StatusCode; @@ -118,4 +119,20 @@ class DatasetOpTesting : public Common { void SetUp() override; }; } // namespace UT + +namespace mindspore { +namespace dataset { +// defined in datasets.cc code, and function prototypes added here for UT purposes +// convert MSTensorVec to DE TensorRow, return empty if fails +TensorRow VecToRow(const MSTensorVec &v); + +// defined in datasets.cc code, and function prototypes added here for UT purposes +// convert DE TensorRow to MSTensorVec, won't fail +MSTensorVec RowToVec(const TensorRow &v); + +MSTensorVec Predicate1(MSTensorVec in); + +MSTensorVec Predicate2(MSTensorVec in); +} // namespace dataset +} // namespace mindspore #endif // TESTS_UT_CPP_DATASET_COMMON_COMMON_H_ diff --git a/tests/ut/cpp/dataset/skip_pushdown_optimization_pass_test.cc b/tests/ut/cpp/dataset/skip_pushdown_optimization_pass_test.cc index e589565f391..9ec7c68a76f 100644 --- a/tests/ut/cpp/dataset/skip_pushdown_optimization_pass_test.cc +++ b/tests/ut/cpp/dataset/skip_pushdown_optimization_pass_test.cc @@ -18,7 +18,6 @@ #include #include "common/common.h" -#include "minddata/dataset/engine/ir/datasetops/dataset_node.h" #include "minddata/dataset/engine/opt/pre/skip_pushdown_pass.h" #include "minddata/dataset/include/dataset/samplers.h" #include "minddata/dataset/include/dataset/vision.h" @@ -107,12 +106,6 @@ class MindDataSkipPushdownTestOptimizationPass : public UT::DatasetOpTesting { } }; -TensorRow VecToRow(const MSTensorVec &v); - -MSTensorVec RowToVec(const TensorRow &v); - -MSTensorVec Predicate1(MSTensorVec in); - /// Feature: MindData Skip Pushdown Optimization Pass Test /// Description: Test MindData Skip Pushdown Optimization Pass with Sampler in MappableSourceNode /// Expectation: Skip node is pushed down and removed after optimization pass