From b9f90e297e0f2c7199cd8215ece7e37f852207d7 Mon Sep 17 00:00:00 2001 From: Hariharan Seshadri Date: Fri, 9 Oct 2020 20:26:28 -0700 Subject: [PATCH] Support sharing of initializers between session via the Python API (#5407) --- .../InferenceTest.cs | 106 ++++++++++-------- .../python/onnxruntime_pybind_mlvalue.cc | 7 +- .../python/onnxruntime_pybind_mlvalue.h | 3 + .../python/onnxruntime_pybind_state.cc | 15 ++- .../python/onnxruntime_pybind_state_common.h | 6 +- .../test/python/onnxruntime_test_python.py | 22 +++- onnxruntime/test/shared_lib/test_inference.cc | 5 +- 7 files changed, 102 insertions(+), 62 deletions(-) diff --git a/csharp/test/Microsoft.ML.OnnxRuntime.Tests/InferenceTest.cs b/csharp/test/Microsoft.ML.OnnxRuntime.Tests/InferenceTest.cs index 40503ee8f4..3744325ae7 100644 --- a/csharp/test/Microsoft.ML.OnnxRuntime.Tests/InferenceTest.cs +++ b/csharp/test/Microsoft.ML.OnnxRuntime.Tests/InferenceTest.cs @@ -1848,60 +1848,68 @@ namespace Microsoft.ML.OnnxRuntime.Tests var ortCpuMemInfo = OrtMemoryInfo.DefaultInstance; var dims = new long[] { 3, 2 }; var dataBuffer = new float[] { 1.0F, 2.0F, 3.0F, 4.0F, 5.0F, 6.0F }; - var allocator = OrtAllocator.DefaultInstance; - var ortAllocationInput = allocator.Allocate((uint)dataBuffer.Length * sizeof(float)); - unsafe - { - float* p = (float*)ortAllocationInput.DangerousGetHandle(); - for (int i = 0; i < dataBuffer.Length; ++i) + var dataHandle = GCHandle.Alloc(dataBuffer, GCHandleType.Pinned); + + try + { + unsafe { - *p++ = dataBuffer[i]; + float* p = (float*)dataHandle.AddrOfPinnedObject(); + for (int i = 0; i < dataBuffer.Length; ++i) + { + *p++ = dataBuffer[i]; + } + } + var dataBufferNumBytes = (uint)dataBuffer.Length * sizeof(float); + var sharedInitializer = OrtValue.CreateTensorValueWithData(ortCpuMemInfo, Tensors.TensorElementType.Float, + dims, dataHandle.AddrOfPinnedObject(), dataBufferNumBytes); + + SessionOptions options = new SessionOptions(); + options.AddInitializer("W", sharedInitializer); + + float[] expectedOutput = { 1.0F, 4.0F, 9.0F, 16.0F, 25.0F, 36.0F }; + int[] expectedDimensions = { 3, 2 }; + + using (var session = new InferenceSession(modelPath, options)) + using (var session2 = new InferenceSession(modelPath, options)) + { + var inputMeta = session.InputMetadata; + var container = new List(); + + foreach (var name in inputMeta.Keys) + { + Assert.Equal(typeof(float), inputMeta[name].ElementType); + Assert.True(inputMeta[name].IsTensor); + var tensor = new DenseTensor(dataBuffer, inputMeta[name].Dimensions); + container.Add(NamedOnnxValue.CreateFromTensor(name, tensor)); + } + + ReadOnlySpan expectedOutputDimensions = new int[] { 1, 1000, 1, 1 }; + string[] expectedOutputNames = new string[] { "Y" }; + + // Run inference with named inputs and outputs created with in Run() + using (var results = session.Run(container)) // results is an IReadOnlyList container + { + foreach (var r in results) + { + validateRunResultData(r.AsTensor(), expectedOutput, expectedDimensions); + } + } + + // Run inference with named inputs and outputs created with in Run() + using (var results2 = session2.Run(container)) // results is an IReadOnlyList container + { + foreach (var r in results2) + { + validateRunResultData(r.AsTensor(), expectedOutput, expectedDimensions); + } + } } } - var dataBufferNumBytes = (uint)dataBuffer.Length * sizeof(float); - var sharedInitializer = OrtValue.CreateTensorValueWithData(ortCpuMemInfo, Tensors.TensorElementType.Float, - dims, ortAllocationInput.DangerousGetHandle(), dataBufferNumBytes); - SessionOptions options = new SessionOptions(); - options.AddInitializer("W", sharedInitializer); - - float[] expectedOutput = { 1.0F, 4.0F, 9.0F, 16.0F, 25.0F, 36.0F }; - int[] expectedDimensions = { 3, 2 }; - - using (var session = new InferenceSession(modelPath, options)) - using (var session2 = new InferenceSession(modelPath, options)) + finally { - var inputMeta = session.InputMetadata; - var container = new List(); - - foreach (var name in inputMeta.Keys) - { - Assert.Equal(typeof(float), inputMeta[name].ElementType); - Assert.True(inputMeta[name].IsTensor); - var tensor = new DenseTensor(dataBuffer, inputMeta[name].Dimensions); - container.Add(NamedOnnxValue.CreateFromTensor(name, tensor)); - } - - ReadOnlySpan expectedOutputDimensions = new int[] { 1, 1000, 1, 1 }; - string[] expectedOutputNames = new string[] { "Y" }; - - // Run inference with named inputs and outputs created with in Run() - using (var results = session.Run(container)) // results is an IReadOnlyList container - { - foreach (var r in results) - { - validateRunResultData(r.AsTensor(), expectedOutput, expectedDimensions); - } - } - - // Run inference with named inputs and outputs created with in Run() - using (var results2 = session2.Run(container)) // results is an IReadOnlyList container - { - foreach (var r in results2) - { - validateRunResultData(r.AsTensor(), expectedOutput, expectedDimensions); - } - } + dataHandle.Free(); } } diff --git a/onnxruntime/python/onnxruntime_pybind_mlvalue.cc b/onnxruntime/python/onnxruntime_pybind_mlvalue.cc index c768569470..b673549cda 100644 --- a/onnxruntime/python/onnxruntime_pybind_mlvalue.cc +++ b/onnxruntime/python/onnxruntime_pybind_mlvalue.cc @@ -23,6 +23,9 @@ namespace python { namespace py = pybind11; using namespace onnxruntime::logging; +const char* PYTHON_ORTVALUE_OBJECT_NAME = "OrtValue"; +const char* PYTHON_ORTVALUE_NATIVE_OBJECT_ATTR = "_ortvalue"; + static bool PyObjectCheck_NumpyArray(PyObject* o) { return PyObject_HasAttrString(o, "__array_finalize__"); } @@ -656,12 +659,12 @@ void CreateGenericMLValue(const onnxruntime::InputDefList* input_def_list, const throw std::runtime_error("Map type is not supported in this build."); #endif - } else if (!accept_only_numpy_array && strcmp(Py_TYPE(value.ptr())->tp_name, "OrtValue") == 0) { + } else if (!accept_only_numpy_array && strcmp(Py_TYPE(value.ptr())->tp_name, PYTHON_ORTVALUE_OBJECT_NAME) == 0) { // This is an OrtValue coming in directly from Python, so assign the underlying native OrtValue handle // to the OrtValue object that we are going to use for Run(). // This should just increase the ref counts of the underlying shared_ptrs in the native OrtValue // and the ref count will be decreased when the OrtValue used for Run() is destroyed upon exit. - *p_mlvalue = value.attr("_ortvalue").cast(); + *p_mlvalue = *value.attr(PYTHON_ORTVALUE_NATIVE_OBJECT_ATTR).cast(); } else if (!accept_only_numpy_array) { auto iterator = PyObject_GetIter(value.ptr()); if (iterator == NULL) { diff --git a/onnxruntime/python/onnxruntime_pybind_mlvalue.h b/onnxruntime/python/onnxruntime_pybind_mlvalue.h index e940dc1fac..72d5ac568e 100644 --- a/onnxruntime/python/onnxruntime_pybind_mlvalue.h +++ b/onnxruntime/python/onnxruntime_pybind_mlvalue.h @@ -21,6 +21,9 @@ namespace python { namespace py = pybind11; +extern const char* PYTHON_ORTVALUE_OBJECT_NAME; +extern const char* PYTHON_ORTVALUE_NATIVE_OBJECT_ATTR; + bool IsNumericNumpyType(int npy_type); bool IsNumericNumpyArray(py::object& py_object); diff --git a/onnxruntime/python/onnxruntime_pybind_state.cc b/onnxruntime/python/onnxruntime_pybind_state.cc index 5c95b25c1e..6c8de4b33e 100644 --- a/onnxruntime/python/onnxruntime_pybind_state.cc +++ b/onnxruntime/python/onnxruntime_pybind_state.cc @@ -1465,7 +1465,7 @@ Applies to session load, initialization, etc. Default is 0.)pbdoc") "Rpbdoc(Get a single session configuration value using the given configuration key.)pbdoc") .def( "register_custom_ops_library", - [](PySessionOptions* options, const std::string& library_path) + [](PySessionOptions* options, const char* library_path) -> void { #if !defined(ORT_MINIMAL_BUILD) // We need to pass in an `OrtSessionOptions` instance because the exported method in the shared library expects that @@ -1473,7 +1473,7 @@ Applies to session load, initialization, etc. Default is 0.)pbdoc") // into the container we are maintaining for that very purpose and the `ortSessionoptions` instance can go out of scope. OrtSessionOptions s; - options->custom_op_libraries_.emplace_back(std::make_shared(library_path.c_str(), s)); + options->custom_op_libraries_.emplace_back(std::make_shared(library_path, s)); // reserve enough memory to hold current contents and the new incoming contents options->custom_op_domains_.reserve(options->custom_op_domains_.size() + s.custom_op_domains_.size()); @@ -1486,7 +1486,16 @@ Applies to session load, initialization, etc. Default is 0.)pbdoc") ORT_THROW("Custom Ops are not supported in this build."); #endif }, - "Rpbdoc(Specify the path to the shared library containing the custom op kernels required to run a model.)pbdoc"); + "Rpbdoc(Specify the path to the shared library containing the custom op kernels required to run a model.)pbdoc") + .def( + "add_initializer", [](PySessionOptions* options, const char* name, py::object& ml_value_pyobject) -> void { + ORT_ENFORCE(strcmp(Py_TYPE(ml_value_pyobject.ptr())->tp_name, PYTHON_ORTVALUE_OBJECT_NAME) == 0, "The provided Python object must be an OrtValue"); + // The user needs to ensure that the python OrtValue being provided as an overriding initializer + // is not destructed as long as any session that uses the provided OrtValue initializer is still in scope + // This is no different than the native APIs + OrtValue* ml_value = ml_value_pyobject.attr(PYTHON_ORTVALUE_NATIVE_OBJECT_ATTR).cast(); + options->AddInitializer(name, ml_value); + }); py::class_(m, "RunOptions", R"pbdoc(Configuration information for a single Run.)pbdoc") .def(py::init()) diff --git a/onnxruntime/python/onnxruntime_pybind_state_common.h b/onnxruntime/python/onnxruntime_pybind_state_common.h index f8dbba8726..6e555438cb 100644 --- a/onnxruntime/python/onnxruntime_pybind_state_common.h +++ b/onnxruntime/python/onnxruntime_pybind_state_common.h @@ -64,7 +64,7 @@ struct PyInferenceSession { void AddCustomOpLibraries(const std::vector>& custom_op_libraries) { if (!custom_op_libraries.empty()) { - custom_op_libraries_.reserve(custom_op_libraries_.size() + custom_op_libraries.size()); + custom_op_libraries_.reserve(custom_op_libraries.size()); for (size_t i = 0; i < custom_op_libraries.size(); ++i) { custom_op_libraries_.push_back(custom_op_libraries[i]); } @@ -82,9 +82,9 @@ struct PyInferenceSession { } private: - #if !defined(ORT_MINIMAL_BUILD) +#if !defined(ORT_MINIMAL_BUILD) // Hold CustomOpLibrary resources so as to tie it to the life cycle of the InferenceSession needing it. - // NOTE: Declare this above `sess_` so that this is destructed AFTER the InferenceSession instance - + // NOTE: Define this above `sess_` so that this is destructed AFTER the InferenceSession instance - // this is so that the custom ops held by the InferenceSession gets destroyed prior to the library getting unloaded // (if ref count of the shared_ptr reaches 0) std::vector> custom_op_libraries_; diff --git a/onnxruntime/test/python/onnxruntime_test_python.py b/onnxruntime/test/python/onnxruntime_test_python.py index 4646103c1e..f1526a7e70 100644 --- a/onnxruntime/test/python/onnxruntime_test_python.py +++ b/onnxruntime/test/python/onnxruntime_test_python.py @@ -664,6 +664,22 @@ class TestInferenceSession(unittest.TestCase): self.assertTrue( 'SessionOptions does not have configuration with key: ' + invalide_key in str(context.exception)) + def testSessionOptionsAddInitializer(self): + # Create an initializer and add it to a SessionOptions instance + so = onnxrt.SessionOptions() + # This initializer is different from the actual initializer in the model for "W" + ortvalue_initializer = onnxrt.OrtValue.ortvalue_from_numpy(np.array([[2.0, 1.0], [4.0, 3.0], [6.0, 5.0]], dtype=np.float32)) + # The user should manage the life cycle of this OrtValue and should keep it in scope + # as long as any session that is going to be reliant on it is in scope + so.add_initializer("W", ortvalue_initializer) + + # Create an InferenceSession that only uses the CPU EP and validate that it uses the + # initializer provided via the SessionOptions instance (overriding the model initializer) + # We only use the CPU EP because the initializer we created is on CPU and we want the model to use that + sess = onnxrt.InferenceSession(get_name("mul_1.onnx"), so, ['CPUExecutionProvider']) + res = sess.run(["Y"], {"X": np.array([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]], dtype=np.float32)}) + self.assertTrue(np.array_equal(res[0], np.array([[2.0, 2.0], [12.0, 12.0], [30.0, 30.0]], dtype=np.float32))) + def testRegisterCustomOpsLibrary(self): if sys.platform.startswith("win"): shared_library = 'custom_op_library.dll' @@ -714,14 +730,14 @@ class TestInferenceSession(unittest.TestCase): def testOrtValue(self): + numpy_arr_input = np.array([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]], dtype=np.float32) + numpy_arr_output = np.array([[1.0, 4.0], [9.0, 16.0], [25.0, 36.0]], dtype=np.float32) + def test_session_with_ortvalue_input(ortvalue): sess = onnxrt.InferenceSession(get_name("mul_1.onnx")) res = sess.run(["Y"], {"X": ortvalue}) self.assertTrue(np.array_equal(res[0], numpy_arr_output)) - numpy_arr_input = np.array([[1.0, 2.0], [3.0, 4.0], [5.0, 6.0]], dtype=np.float32) - numpy_arr_output = np.array([[1.0, 4.0], [9.0, 16.0], [25.0, 36.0]], dtype=np.float32) - ortvalue1 = onnxrt.OrtValue.ortvalue_from_numpy(numpy_arr_input) self.assertEqual(ortvalue1.device_name(), "cpu") self.assertEqual(ortvalue1.shape(), [3, 2]) diff --git a/onnxruntime/test/shared_lib/test_inference.cc b/onnxruntime/test/shared_lib/test_inference.cc index 891d13737f..d1fae73651 100644 --- a/onnxruntime/test/shared_lib/test_inference.cc +++ b/onnxruntime/test/shared_lib/test_inference.cc @@ -949,11 +949,12 @@ TEST(CApiTest, TestSharingOfInitializer) { // prepare expected inputs and outputs std::vector expected_dims_y = {3, 2}; - std::vector expected_values_y = {1.0f, 4.0f, 9.0f, 16.0f, 25.0f, 36.0f}; + std::vector expected_values_y = {2.0f, 2.0f, 12.0f, 12.0f, 30.0f, 30.0f}; Ort::SessionOptions session_options; Ort::MemoryInfo mem_info = Ort::MemoryInfo::CreateCpu(OrtArenaAllocator, OrtMemTypeDefault); - float data[] = {1., 2., 3., 4., 5., 6.}; + // These values are different from the actual initializer values in the model + float data[] = {2., 1., 4., 3., 6., 5.}; const int data_len = sizeof(data) / sizeof(data[0]); const int64_t shape[] = {3, 2}; const size_t shape_len = sizeof(shape) / sizeof(shape[0]);