diff --git a/onnxruntime/core/graph/model.cc b/onnxruntime/core/graph/model.cc index 3cc756a12a..8b43d7356b 100644 --- a/onnxruntime/core/graph/model.cc +++ b/onnxruntime/core/graph/model.cc @@ -96,7 +96,7 @@ Model::Model(std::unique_ptr model_proto, const IOnnxRuntimeOpSchema } if (!model_proto->has_ir_version() || model_proto->ir_version() > ONNX_NAMESPACE::Version::IR_VERSION) { - throw std::invalid_argument("Unknown model file format version."); + throw std::invalid_argument("Unknown model file format version."); } model_proto_ = std::move(model_proto); @@ -394,14 +394,14 @@ Status Model::LoadFromBytes(int count, void* p_bytes, /*out*/ ONNX_NAMESPACE::Mo Status Model::LoadFromBytes(int count, void* p_bytes, /*out*/ std::shared_ptr& p_model, const IOnnxRuntimeOpSchemaRegistryList* local_registries, const logging::Logger& logger) { - ModelProto model_proto; + auto model_proto = onnxruntime::make_unique(); - auto status = LoadFromBytes(count, p_bytes, model_proto); + auto status = LoadFromBytes(count, p_bytes, *model_proto); if (!status.IsOK()) { return status; } - p_model = std::make_shared(model_proto, local_registries, logger); + p_model = std::make_shared(std::move(model_proto), local_registries, logger); ORT_RETURN_IF_ERROR(p_model->MainGraph().Resolve(true)); @@ -441,11 +441,11 @@ Status Model::Load(int fd, ONNX_NAMESPACE::ModelProto& model_proto) { Status Model::Load(int fd, std::shared_ptr& p_model, const IOnnxRuntimeOpSchemaRegistryList* local_registries, const logging::Logger& logger) { - ModelProto model_proto; + auto model_proto = onnxruntime::make_unique(); - ORT_RETURN_IF_ERROR(Load(fd, model_proto)); + ORT_RETURN_IF_ERROR(Load(fd, *model_proto)); - p_model = std::make_shared(model_proto, local_registries, logger); + p_model = std::make_shared(std::move(model_proto), local_registries, logger); ORT_RETURN_IF_ERROR(p_model->MainGraph().Resolve(true)); diff --git a/onnxruntime/core/session/inference_session.cc b/onnxruntime/core/session/inference_session.cc index 8f46e9ee16..651a2eee18 100644 --- a/onnxruntime/core/session/inference_session.cc +++ b/onnxruntime/core/session/inference_session.cc @@ -372,9 +372,10 @@ common::Status InferenceSession::Load(std::function& model) { - ModelProto model_proto; + auto model_proto = onnxruntime::make_unique(); google::protobuf::io::IstreamInputStream zero_copy_input(&model_istream); - const bool result = model_proto.ParseFromZeroCopyStream(&zero_copy_input) && model_istream.eof(); + const bool result = model_proto->ParseFromZeroCopyStream(&zero_copy_input) && model_istream.eof(); if (!result) { return Status(common::ONNXRUNTIME, common::INVALID_PROTOBUF, "Failed to load model because protobuf parsing failed."); } #ifdef ENABLE_LANGUAGE_INTEROP_OPS - LoadInterOp(model_proto, interop_domains_, [&](const char* msg) { LOGS(*session_logger_, WARNING) << msg; }); + LoadInterOp(*model_proto, interop_domains_, [&](const char* msg) { LOGS(*session_logger_, WARNING) << msg; }); for (const auto& domain : interop_domains_) { AddCustomOpDomains({domain.get()}); } #endif - return onnxruntime::Model::Load(model_proto, model, HasLocalSchema() ? &custom_schema_registries_ : nullptr, + return onnxruntime::Model::Load(std::move(model_proto), model, HasLocalSchema() ? &custom_schema_registries_ : nullptr, *session_logger_); }; @@ -516,21 +518,21 @@ common::Status InferenceSession::Load(const void* model_data, int model_data_len } auto loader = [this, model_data, model_data_len](std::shared_ptr& model) { - ModelProto model_proto; + auto model_proto = onnxruntime::make_unique(); - const bool result = model_proto.ParseFromArray(model_data, model_data_len); + const bool result = model_proto->ParseFromArray(model_data, model_data_len); if (!result) { return Status(common::ONNXRUNTIME, common::INVALID_PROTOBUF, "Failed to load model because protobuf parsing failed."); } #ifdef ENABLE_LANGUAGE_INTEROP_OPS - LoadInterOp(model_proto, interop_domains_, [&](const char* msg) { LOGS(*session_logger_, WARNING) << msg; }); + LoadInterOp(*model_proto, interop_domains_, [&](const char* msg) { LOGS(*session_logger_, WARNING) << msg; }); for (const auto& domain : interop_domains_) { AddCustomOpDomains({domain.get()}); } #endif - return onnxruntime::Model::Load(model_proto, model, HasLocalSchema() ? &custom_schema_registries_ : nullptr, + return onnxruntime::Model::Load(std::move(model_proto), model, HasLocalSchema() ? &custom_schema_registries_ : nullptr, *session_logger_); }; @@ -551,7 +553,8 @@ common::Status InferenceSession::Load() { AddCustomOpDomains({domain.get()}); } #endif - return Model::Load(*this->model_proto_, model, HasLocalSchema() ? &custom_schema_registries_ : nullptr, + // Pass on ownership of the parsed ModelProto to the Model instance (its job here is done by this stage) + return Model::Load(std::move(this->model_proto_), model, HasLocalSchema() ? &custom_schema_registries_ : nullptr, *session_logger_); };