From a270d8407efb0a3e67724b6581b05a4a319bc8a3 Mon Sep 17 00:00:00 2001 From: Pranav Sharma Date: Wed, 21 Jun 2023 22:46:26 -0700 Subject: [PATCH] Allow saving of large models after optimization (github issue 12882) (#16440) ### Description Allow saving of large models after optimization. ### Motivation and Context Addresses https://github.com/microsoft/onnxruntime/issues/12882 --- .../onnxruntime_session_options_config_keys.h | 10 ++++++- onnxruntime/core/session/inference_session.cc | 15 ++++++++++- .../test/python/onnxruntime_test_python.py | 27 +++++++++++++++++++ 3 files changed, 50 insertions(+), 2 deletions(-) diff --git a/include/onnxruntime/core/session/onnxruntime_session_options_config_keys.h b/include/onnxruntime/core/session/onnxruntime_session_options_config_keys.h index 1ef821e7c9..37545f41b4 100644 --- a/include/onnxruntime/core/session/onnxruntime_session_options_config_keys.h +++ b/include/onnxruntime/core/session/onnxruntime_session_options_config_keys.h @@ -216,4 +216,12 @@ static const char* const kDebugLayoutTransformation = "session.debug_layout_tran // Option values: // - "0": CPU EP fallback is not disabled. [DEFAULT] // - "1": CPU EP fallback is disabled. -static const char* const kOrtSessionOptionsDisableCPUEPFallback = "session.disable_cpu_ep_fallback"; \ No newline at end of file +static const char* const kOrtSessionOptionsDisableCPUEPFallback = "session.disable_cpu_ep_fallback"; + +// Use this config when serializing a large model after optimization to specify an external initializers file +static const char* const kOrtSessionOptionsOptimizedModelExternalInitializersFileName = + "session.optimized_model_external_initializers_file_name"; + +// Use this config to control the minimum size of the initializer when externalizing it during serialization +static const char* const kOrtSessionOptionsOptimizedModelExternalInitializersMinSizeInBytes = + "session.optimized_model_external_initializers_min_size_in_bytes"; diff --git a/onnxruntime/core/session/inference_session.cc b/onnxruntime/core/session/inference_session.cc index 99308e0df6..aa29a20c87 100644 --- a/onnxruntime/core/session/inference_session.cc +++ b/onnxruntime/core/session/inference_session.cc @@ -1725,7 +1725,20 @@ common::Status InferenceSession::Initialize() { if (saving_ort_format) { ORT_RETURN_IF_ERROR_SESSIONID_(SaveToOrtFormat(session_options_.optimized_model_filepath)); } else { - ORT_RETURN_IF_ERROR_SESSIONID_(Model::Save(*model_, session_options_.optimized_model_filepath)); + const std::string optimized_model_external_initializers_file_name = + session_options_.config_options.GetConfigOrDefault( + kOrtSessionOptionsOptimizedModelExternalInitializersFileName, ""); + if (optimized_model_external_initializers_file_name.empty()) { + ORT_RETURN_IF_ERROR_SESSIONID_(Model::Save(*model_, session_options_.optimized_model_filepath)); + } else { + const size_t optimized_model_external_initializers_min_size_in_bytes = + ParseStringWithClassicLocale(session_options_.config_options.GetConfigOrDefault( + kOrtSessionOptionsOptimizedModelExternalInitializersMinSizeInBytes, "1024")); + ORT_RETURN_IF_ERROR_SESSIONID_(Model::SaveWithExternalInitializers(*model_, + session_options_.optimized_model_filepath, + optimized_model_external_initializers_file_name, + optimized_model_external_initializers_min_size_in_bytes)); + } } } diff --git a/onnxruntime/test/python/onnxruntime_test_python.py b/onnxruntime/test/python/onnxruntime_test_python.py index 5e02906896..e18a6276cd 100644 --- a/onnxruntime/test/python/onnxruntime_test_python.py +++ b/onnxruntime/test/python/onnxruntime_test_python.py @@ -94,6 +94,33 @@ class TestInferenceSession(unittest.TestCase): else: raise onnxruntime_error + def testModelSerializationWithExternalInitializers(self): # noqa: N802 + try: + so = onnxrt.SessionOptions() + so.log_severity_level = 1 + so.logid = "TestModelSerializationWithExternalInitializers" + so.optimized_model_filepath = "./model_with_external_initializers.onnx" + external_initializers_file = "external_initializers.bin" + so.add_session_config_entry( + "session.optimized_model_external_initializers_file_name", external_initializers_file + ) + so.add_session_config_entry("session.optimized_model_external_initializers_min_size_in_bytes", "100") + onnxrt.InferenceSession( + get_name("mnist.onnx"), + sess_options=so, + providers=["CPUExecutionProvider"], + ) + self.assertTrue(os.path.isfile(so.optimized_model_filepath)) + self.assertTrue(os.path.isfile(external_initializers_file)) + except Fail as onnxruntime_error: + if ( + str(onnxruntime_error) == "[ONNXRuntimeError] : 1 : FAIL : Unable to serialize model as it contains" + " compiled nodes. Please disable any execution providers which generate compiled nodes." + ): + pass + else: + raise onnxruntime_error + def testGetProviders(self): # noqa: N802 self.assertTrue("CPUExecutionProvider" in onnxrt.get_available_providers()) # get_all_providers() returns the default EP order from highest to lowest.