diff --git a/inference-engine/src/low_precision_transformations/src/split.cpp b/inference-engine/src/low_precision_transformations/src/split.cpp index 91a32f77868..4dcff39fc5d 100644 --- a/inference-engine/src/low_precision_transformations/src/split.cpp +++ b/inference-engine/src/low_precision_transformations/src/split.cpp @@ -83,7 +83,8 @@ bool SplitTransformation::transform(TransformationContext& context, ngraph::patt parent = subtract; } - const auto multiply = std::make_shared(parent, splitedMul[i]); + const auto multiply = std::make_shared>(parent, splitedMul[i]); + NetworkHelper::setOutDataPrecisionForTypeRelaxed(multiply, dequantization.multiply->get_output_element_type(0)); copy_runtime_info({ newSplit, multiply }, multiply); lastNodes.push_back(multiply); diff --git a/inference-engine/tests/functional/inference_engine/lp_transformations/split_transformation.cpp b/inference-engine/tests/functional/inference_engine/lp_transformations/split_transformation.cpp index 9f04c18a580..ae2ada651ef 100644 --- a/inference-engine/tests/functional/inference_engine/lp_transformations/split_transformation.cpp +++ b/inference-engine/tests/functional/inference_engine/lp_transformations/split_transformation.cpp @@ -57,12 +57,19 @@ inline std::ostream& operator<<(std::ostream& os, return os; } -class SplitTransformation : public LayerTransformation, public testing::WithParamInterface { +typedef std::tuple < + ngraph::element::Type, + SplitTransformationTestValues +> SplitTransformationParams; + +class SplitTransformation : public LayerTransformation, public testing::WithParamInterface { public: void SetUp() override { - SplitTransformationTestValues testValues = GetParam(); + ngraph::element::Type precision = std::get<0>(GetParam()); + SplitTransformationTestValues testValues = std::get<1>(GetParam()); actualFunction = ngraph::builder::subgraph::SplitFunction::getOriginal( + precision, testValues.inputShape, testValues.actual.precisionBeforeDequantization, testValues.actual.dequantization, @@ -74,6 +81,7 @@ public: transformer.transform(actualFunction); referenceFunction = ngraph::builder::subgraph::SplitFunction::getReference( + precision, testValues.inputShape, testValues.expected.inputPrecision, testValues.expected.dequantizationBefore, @@ -83,11 +91,13 @@ public: testValues.numSplits); } - static std::string getTestCaseName(testing::TestParamInfo obj) { - const SplitTransformationTestValues testValues = obj.param; + static std::string getTestCaseName(testing::TestParamInfo obj) { + ngraph::element::Type precision = std::get<0>(obj.param); + SplitTransformationTestValues testValues = std::get<1>(obj.param); std::ostringstream result; - result << toString(testValues.params) << "_" << + result << precision << "_" << + toString(testValues.params) << "_" << testValues.inputShape << "_" << testValues.actual.precisionBeforeDequantization << "_" << testValues.actual.dequantization << "_" << @@ -106,6 +116,11 @@ TEST_P(SplitTransformation, CompareFunctions) { ASSERT_TRUE(res.first) << res.second; } +const std::vector precisions = { + ngraph::element::f32, + ngraph::element::f16 +}; + const std::vector testValues = { // U8 per tensor quantization { @@ -425,6 +440,8 @@ const std::vector testValues = { INSTANTIATE_TEST_CASE_P( smoke_LPT, SplitTransformation, - ::testing::ValuesIn(testValues), + ::testing::Combine( + ::testing::ValuesIn(precisions), + ::testing::ValuesIn(testValues)), SplitTransformation::getTestCaseName); } // namespace diff --git a/inference-engine/tests/ngraph_helpers/lpt_ngraph_functions/include/lpt_ngraph_functions/split_function.hpp b/inference-engine/tests/ngraph_helpers/lpt_ngraph_functions/include/lpt_ngraph_functions/split_function.hpp index 661c4c2e80c..46fa7cb9d61 100644 --- a/inference-engine/tests/ngraph_helpers/lpt_ngraph_functions/include/lpt_ngraph_functions/split_function.hpp +++ b/inference-engine/tests/ngraph_helpers/lpt_ngraph_functions/include/lpt_ngraph_functions/split_function.hpp @@ -20,6 +20,7 @@ namespace subgraph { class SplitFunction { public: static std::shared_ptr getOriginal( + const element::Type& precision, const ngraph::Shape& inputShape, const ngraph::element::Type precisionBeforeDequantization, const ngraph::builder::subgraph::DequantizationOperations& dequantization, @@ -34,6 +35,7 @@ public: const size_t numSplit); static std::shared_ptr getReference( + const element::Type& precision, const ngraph::Shape& inputShape, const ngraph::element::Type inputPrecision, const ngraph::builder::subgraph::DequantizationOperations& dequantizationBefore, diff --git a/inference-engine/tests/ngraph_helpers/lpt_ngraph_functions/src/split_function.cpp b/inference-engine/tests/ngraph_helpers/lpt_ngraph_functions/src/split_function.cpp index 6820572c302..fe2e797cd32 100644 --- a/inference-engine/tests/ngraph_helpers/lpt_ngraph_functions/src/split_function.cpp +++ b/inference-engine/tests/ngraph_helpers/lpt_ngraph_functions/src/split_function.cpp @@ -17,26 +17,29 @@ namespace ngraph { namespace builder { namespace subgraph { - std::shared_ptr SplitFunction::getOriginal( - const ngraph::Shape& inputShape, - const ngraph::element::Type precisionBeforeDequantization, - const ngraph::builder::subgraph::DequantizationOperations& dequantization, - const int64_t splitedAxis, - const size_t numSplits) { - const std::shared_ptr input = std::make_shared( - precisionBeforeDequantization, - ngraph::Shape(inputShape)); +std::shared_ptr SplitFunction::getOriginal( + const element::Type& precision, + const ngraph::Shape& inputShape, + const ngraph::element::Type precisionBeforeDequantization, + const ngraph::builder::subgraph::DequantizationOperations& dequantization, + const int64_t splitedAxis, + const size_t numSplits) { + const std::shared_ptr input = std::make_shared( + precisionBeforeDequantization, + ngraph::Shape(inputShape)); - const std::shared_ptr dequantizationOp = makeDequantization(input, dequantization); - const auto constant = std::make_shared(element::i64, Shape{ }, splitedAxis); - const std::shared_ptr split = std::make_shared(dequantizationOp, constant, numSplits); + auto dequantizationStructure = dequantization; + dequantizationStructure.multiply.outPrecision = precision; + const std::shared_ptr dequantizationOp = makeDequantization(input, dequantization); + const auto constant = std::make_shared(element::i64, Shape{ }, splitedAxis); + const std::shared_ptr split = std::make_shared(dequantizationOp, constant, numSplits); - ngraph::ResultVector results; - for (size_t i = 0; i < numSplits; ++i) { - results.push_back(std::make_shared(split->output(i))); - } - return std::make_shared(results, ngraph::ParameterVector{ input }, "SplitFunction"); + ngraph::ResultVector results; + for (size_t i = 0; i < numSplits; ++i) { + results.push_back(std::make_shared(split->output(i))); } + return std::make_shared(results, ngraph::ParameterVector{ input }, "SplitFunction"); +} std::shared_ptr SplitFunction::getOriginal( const ngraph::element::Type originalFunctionPrecision, @@ -67,6 +70,7 @@ std::shared_ptr SplitFunction::getOriginal( } std::shared_ptr SplitFunction::getReference( + const element::Type& precision, const ngraph::Shape& inputShape, const ngraph::element::Type inputPrecision, const ngraph::builder::subgraph::DequantizationOperations& dequantizationBefore, @@ -86,8 +90,15 @@ std::shared_ptr SplitFunction::getReference( ngraph::ResultVector results; for (size_t i = 0; i < numSplit; ++i) { - results.push_back(std::make_shared( - dequantizationAfter.empty() ? split->output(i) : makeDequantization(split->output(i), dequantizationAfter[i]))); + if (!dequantizationAfter.empty()) { + auto dequantizationStructure = dequantizationAfter[i]; + if (!dequantizationStructure.multiply.empty()) { + dequantizationStructure.multiply.outPrecision = precision; + } + results.push_back(std::make_shared(makeDequantization(split->output(i), dequantizationAfter[i]))); + } else { + results.push_back(std::make_shared(split->output(i))); + } } return std::make_shared(results, ngraph::ParameterVector{ input }, "SplitTransformation"); }