mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-23 19:32:23 +00:00
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
This commit is contained in:
parent
89f8f20a61
commit
a270d8407e
3 changed files with 50 additions and 2 deletions
|
|
@ -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";
|
||||
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";
|
||||
|
|
|
|||
|
|
@ -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<size_t>(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));
|
||||
}
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
|
|
@ -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.
|
||||
|
|
|
|||
Loading…
Reference in a new issue