diff --git a/modules/dnn/src/net_impl.cpp b/modules/dnn/src/net_impl.cpp index 210e04b487..c1f1ec07d0 100644 --- a/modules/dnn/src/net_impl.cpp +++ b/modules/dnn/src/net_impl.cpp @@ -2999,6 +2999,20 @@ void Net::Impl::getLayerTypes(std::vector& layersTypes) const // TODO drop? int Net::Impl::getLayersCount(const String& layerType) const { + if (mainGraph) { + int count = 0; + for (const Ptr& g: allgraphs) { + const std::vector >& prog = g->prog(); + for (const Ptr& layer: prog) { + if (!layer) + continue; + if (layer->type == layerType) + count++; + } + } + return count; + } + int count = 0; for (Impl::MapIdToLayerData::const_iterator it = layers.begin(); it != layers.end(); it++) diff --git a/modules/dnn/src/onnx/onnx_graph_simplifier.cpp b/modules/dnn/src/onnx/onnx_graph_simplifier.cpp index 03bb8bf3c4..b271c5d9cb 100644 --- a/modules/dnn/src/onnx/onnx_graph_simplifier.cpp +++ b/modules/dnn/src/onnx/onnx_graph_simplifier.cpp @@ -22,6 +22,71 @@ CV__DNN_INLINE_NS_BEGIN extern bool DNN_DIAGNOSTICS_RUN; +static bool isValidPerm(const std::vector& perm, int rank) +{ + if (rank <= 0 || perm.size() != static_cast(rank)) + return false; + + std::vector used(rank, false); + for (size_t i = 0; i < perm.size(); ++i) + { + if (perm[i] < 0 || perm[i] >= rank || used[perm[i]]) + return false; + used[perm[i]] = true; + } + return true; +} + +static bool getPermAttr(opencv_onnx::NodeProto* node, std::vector& perm, bool& hasPerm) +{ + CV_Assert(node); + + hasPerm = false; + for (int i = 0; i < node->attribute_size(); ++i) + { + opencv_onnx::AttributeProto attr = node->attribute(i); + // ONNX Transpose uses "perm"; OpenCV internal Permute layer uses "order" after import. + if (attr.name() != "perm") + continue; + + hasPerm = true; + perm.clear(); + for (int j = 0; j < attr.ints_size(); ++j) + { + int64_t axis = attr.ints(j); + if (axis < std::numeric_limits::min() || axis > std::numeric_limits::max()) + return false; + perm.push_back(static_cast(axis)); + } + return true; + } + + perm.clear(); + return true; +} + +static void getDefaultPerm(int rank, std::vector& perm) +{ + perm.resize(rank); + // ONNX Transpose without "perm" reverses all dimensions. + for (int i = 0; i < rank; ++i) + perm[i] = rank - 1 - i; +} + +static bool isIdentityPerm(const std::vector& p2, const std::vector& p1) +{ + if (p1.size() != p2.size()) + return false; + + for (size_t i = 0; i < p2.size(); ++i) + { + // For ONNX Transpose, p1 followed by p2 composes as p1[p2[i]]. + if (p2[i] < 0 || p2[i] >= static_cast(p1.size()) || p1[p2[i]] != static_cast(i)) + return false; + } + return true; +} + // This wrapper can behave differently for fake input nodes and real graph nodes. class ONNXNodeWrapper : public ImportNodeWrapper { @@ -163,6 +228,34 @@ public: return net.node(nodeId - numInputs - numInitializers).output(outId); } + bool hasSingleConsumer(const std::string& name) const + { + int count = 0; + for (int i = 0; i < net.node_size(); ++i) + { + const opencv_onnx::NodeProto& node = net.node(i); + for (int j = 0; j < node.input_size(); ++j) + { + if (node.input(j) == name) + { + if (++count > 1) + return false; + } + } + } + return count == 1; + } + + bool isGraphOutput(const std::string& name) const + { + for (int i = 0; i < net.output_size(); ++i) + { + if (net.output(i).name() == name) + return true; + } + return false; + } + virtual void removeNode(int idx) CV_OVERRIDE { if (idx >= numInputs + numInitializers) @@ -1754,6 +1847,72 @@ public: } }; +class ConsecutiveTransposePairsSubgraph : public Subgraph +{ +public: + ConsecutiveTransposePairsSubgraph() + { + input = addNodeToMatch(""); + transpose1 = addNodeToMatch("Transpose", input); + transpose2 = addNodeToMatch("Transpose", transpose1); + setFusedNode("Identity", input); + } + + virtual bool match(const Ptr& net, int nodeId, + std::vector& matchedNodesIds) CV_OVERRIDE + { + if (!Subgraph::match(net, nodeId, matchedNodesIds)) + return false; + + Ptr onnxNet = net.dynamicCast(); + if (!onnxNet) + return false; + + Ptr nodeWrapper1 = net->getNode(matchedNodesIds[transpose1]).dynamicCast(); + Ptr nodeWrapper2 = net->getNode(matchedNodesIds[transpose2]).dynamicCast(); + if (!nodeWrapper1 || !nodeWrapper1->node || !nodeWrapper2 || !nodeWrapper2->node) + return false; + + if (nodeWrapper1->node->output_size() != 1 || nodeWrapper2->node->output_size() != 1) + return false; + const std::string& intermediate = nodeWrapper1->node->output(0); + if (!onnxNet->hasSingleConsumer(intermediate) || onnxNet->isGraphOutput(intermediate)) + return false; + + std::vector perm1, perm2; + bool hasPerm1, hasPerm2; + if (!getPermAttr(nodeWrapper1->node, perm1, hasPerm1) || + !getPermAttr(nodeWrapper2->node, perm2, hasPerm2)) + return false; + + int inputRank = -1; + if (hasPerm1 && hasPerm2) + { + if (perm1.size() != perm2.size()) + return false; + inputRank = static_cast(perm1.size()); + } + else + { + inputRank = onnxNet->getTensorShapeSize(matchedNodesIds[transpose1], 0); + if (inputRank <= 0) + return false; + if (!hasPerm1) + getDefaultPerm(inputRank, perm1); + if (!hasPerm2) + getDefaultPerm(inputRank, perm2); + } + + if (!isValidPerm(perm1, inputRank) || !isValidPerm(perm2, inputRank)) + return false; + + return isIdentityPerm(perm2, perm1); + } + +protected: + int input, transpose1, transpose2; +}; + void simplifySubgraphs(opencv_onnx::GraphProto& net, const std::string& basePath) { std::vector > subgraphs; @@ -1790,6 +1949,8 @@ void simplifySubgraphs(opencv_onnx::GraphProto& net, const std::string& basePath subgraphs.push_back(makePtr()); subgraphs.push_back(makePtr()); } + // Cleanup pass: remove identity-equivalent consecutive Transpose nodes after larger fusions. + subgraphs.push_back(makePtr()); simplifySubgraphs(Ptr(new ONNXGraphWrapper(net, basePath)), subgraphs); } diff --git a/modules/dnn/test/test_onnx_importer.cpp b/modules/dnn/test/test_onnx_importer.cpp index 57ba14794a..a6ee94f5b0 100644 --- a/modules/dnn/test/test_onnx_importer.cpp +++ b/modules/dnn/test/test_onnx_importer.cpp @@ -3517,6 +3517,70 @@ TEST_P(Test_ONNX_layers, ClipDivSharedConstant) { testONNXModels("clip_div_shared_constant"); } +static Mat makeConsecutiveTransposeInput() +{ + int inputShape[] = {1, 3, 4, 4}; + Mat input(4, inputShape, CV_32F); + float* data = input.ptr(); + for (size_t i = 0; i < input.total(); ++i) + data[i] = static_cast(i); + return input; +} + +static void testConsecutiveTransposeModel(const String& basename, + int backendId, int targetId, + int expectedTransposeLayers, + const Mat& input, const Mat& ref) +{ + Net net = readNetFromONNX(_tf("models/" + basename + ".onnx")); + ASSERT_FALSE(net.empty()); + int numTransposeLayers = net.getLayersCount("Permute") + net.getLayersCount("Transpose"); + if (expectedTransposeLayers >= 0) + EXPECT_EQ(numTransposeLayers, expectedTransposeLayers); + else + // Non-identity transpose pairs must not be eliminated completely. + // They may still be fused into a single equivalent Transpose later. + EXPECT_GT(numTransposeLayers, 0); + + if (net.getMainGraph()) + net.setPreferableBackend(DNN_BACKEND_OPENCV); + else + { + net.setPreferableBackend(backendId); + net.setPreferableTarget(targetId); + } + + net.setInput(input); + Mat output = net.forward(); + + EXPECT_EQ(shape(output), shape(ref)); + EXPECT_LT(cv::norm(output, ref, NORM_INF), 1e-5); +} + +TEST_P(Test_ONNX_layers, ConsecutiveTransposeIdentity) +{ + Mat input = makeConsecutiveTransposeInput(); + + testConsecutiveTransposeModel("transpose_identity", backend, target, 0, input, input); +} + +TEST_P(Test_ONNX_layers, ConsecutiveTransposeDefaultPerm) +{ + Mat input = makeConsecutiveTransposeInput(); + + testConsecutiveTransposeModel("transpose_default_perm", backend, target, 0, input, input); +} + +TEST_P(Test_ONNX_layers, ConsecutiveTransposeNonIdentity) +{ + Mat input = makeConsecutiveTransposeInput(); + Mat ref1, ref; + cv::transposeND(input, std::vector{0, 2, 3, 1}, ref1); + cv::transposeND(ref1, std::vector{0, 1, 3, 2}, ref); + + testConsecutiveTransposeModel("transpose_non_identity", backend, target, -1, input, ref); +} + TEST_P(Test_ONNX_layers, TopK) { if (backend == DNN_BACKEND_INFERENCE_ENGINE_NGRAPH || backend == DNN_BACKEND_INFERENCE_ENGINE_NN_BUILDER_2019 ||