diff --git a/modules/dnn/src/net.cpp b/modules/dnn/src/net.cpp index 6f167df3bb..4e6e26808d 100644 --- a/modules/dnn/src/net.cpp +++ b/modules/dnn/src/net.cpp @@ -136,7 +136,7 @@ void Net::finalizeNet() CV_TRACE_FUNCTION(); CV_Assert(impl); #ifdef HAVE_ONNXRUNTIME - if (impl->mainGraph && impl->modelFormat == DNN_MODEL_ONNX && !impl->modelFileName.empty()) + if (impl->useOrtEngine && impl->mainGraph && impl->modelFormat == DNN_MODEL_ONNX && !impl->modelFileName.empty()) { impl->finalizeOrt(); return; diff --git a/modules/dnn/src/net_impl.hpp b/modules/dnn/src/net_impl.hpp index 435ad5b7f9..0b1e2673cb 100644 --- a/modules/dnn/src/net_impl.hpp +++ b/modules/dnn/src/net_impl.hpp @@ -258,7 +258,8 @@ struct Net::Impl : public detail::NetImplBase std::shared_ptr ort_env; std::shared_ptr ort_session; std::shared_ptr ort_names_cache; - bool ortNeedsReinit = true; // session needs (re)creation on next finalizeNet + bool useOrtEngine = false; // true only when user explicitly selected ENGINE_ORT + bool ortNeedsReinit = false; // session needs (re)creation on next finalizeNet #endif void allocateLayer(int lid, const LayersShapesMap& layersShapes); diff --git a/modules/dnn/src/net_impl2.cpp b/modules/dnn/src/net_impl2.cpp index 4207a92eb7..3e5883c0a4 100644 --- a/modules/dnn/src/net_impl2.cpp +++ b/modules/dnn/src/net_impl2.cpp @@ -619,7 +619,7 @@ void Net::Impl::allocateLayerOutputs( void Net::Impl::forwardMainGraph(InputArrayOfArrays inputs, OutputArrayOfArrays outputs) { #ifdef HAVE_ONNXRUNTIME - if (mainGraph && modelFormat == DNN_MODEL_ONNX && !modelFileName.empty()) + if (useOrtEngine && mainGraph && modelFormat == DNN_MODEL_ONNX && !modelFileName.empty()) finalizeOrt(); if (this->ort_session) { @@ -678,7 +678,7 @@ void Net::Impl::forwardMainGraph(InputArrayOfArrays inputs, OutputArrayOfArrays void Net::Impl::forwardWithSingleOutput(const std::string& outname, OutputArrayOfArrays outputBlobs) { #ifdef HAVE_ONNXRUNTIME - if (mainGraph && modelFormat == DNN_MODEL_ONNX && !modelFileName.empty()) + if (useOrtEngine && mainGraph && modelFormat == DNN_MODEL_ONNX && !modelFileName.empty()) finalizeOrt(); if (this->ort_session) { @@ -763,7 +763,7 @@ void Net::Impl::forwardWithSingleOutput(const std::string& outname, OutputArrayO void Net::Impl::forwardWithMultipleOutputs(OutputArrayOfArrays outblobs, const std::vector& outnames) { #ifdef HAVE_ONNXRUNTIME - if (mainGraph && modelFormat == DNN_MODEL_ONNX && !modelFileName.empty()) + if (useOrtEngine && mainGraph && modelFormat == DNN_MODEL_ONNX && !modelFileName.empty()) finalizeOrt(); if (this->ort_session) { @@ -947,7 +947,7 @@ void Net::Impl::traceArg(std::ostream& strm_, const char* prefix, size_t i, Arg void Net::Impl::setMainGraphInput(InputArray m, const std::string& inpname) { #ifdef HAVE_ONNXRUNTIME - if (ortNeedsReinit && mainGraph && modelFormat == DNN_MODEL_ONNX && !modelFileName.empty()) + if (useOrtEngine && ortNeedsReinit && mainGraph && modelFormat == DNN_MODEL_ONNX && !modelFileName.empty()) { Mat inputMat = m.getMat(); if (inputMat.empty()) diff --git a/modules/dnn/src/net_impl_backend.cpp b/modules/dnn/src/net_impl_backend.cpp index 1326681409..659cf8927e 100644 --- a/modules/dnn/src/net_impl_backend.cpp +++ b/modules/dnn/src/net_impl_backend.cpp @@ -273,7 +273,7 @@ void Net::Impl::setPreferableBackend(Net& net, int backendId) return; #ifdef HAVE_ONNXRUNTIME - if (mainGraph && modelFormat == DNN_MODEL_ONNX && !modelFileName.empty()) + if (useOrtEngine && mainGraph && modelFormat == DNN_MODEL_ONNX && !modelFileName.empty()) { preferableBackend = backendId; ortNeedsReinit = true; // will be applied on finalizeNet() @@ -317,7 +317,7 @@ void Net::Impl::setPreferableBackend(Net& net, int backendId) void Net::Impl::setPreferableTarget(int targetId) { #ifdef HAVE_ONNXRUNTIME - if (mainGraph && modelFormat == DNN_MODEL_ONNX && !modelFileName.empty()) + if (useOrtEngine && mainGraph && modelFormat == DNN_MODEL_ONNX && !modelFileName.empty()) { int resolved = IS_DNN_CUDA_TARGET(targetId) ? targetId : DNN_TARGET_CPU; if (preferableTarget != resolved) diff --git a/modules/dnn/src/onnx/onnx_importer2.cpp b/modules/dnn/src/onnx/onnx_importer2.cpp index 54a8604de6..acbd1d4c85 100644 --- a/modules/dnn/src/onnx/onnx_importer2.cpp +++ b/modules/dnn/src/onnx/onnx_importer2.cpp @@ -2847,6 +2847,7 @@ Net readNetFromONNX2_ORT(const String& onnxFile) auto impl = net.getImpl(); impl->modelFileName = onnxFile; impl->modelFormat = DNN_MODEL_ONNX; + impl->useOrtEngine = true; impl->ortNeedsReinit = true; // Create an empty main graph placeholder so that callers can detect