register execution provider when onnxruntime server creating sessions

This commit is contained in:
zhijxu 2019-10-23 01:37:05 +00:00 committed by Changming Sun
parent 1fa956fb3f
commit be7c24247f
2 changed files with 35 additions and 1 deletions

View file

@ -5,6 +5,24 @@
#include "environment.h"
#include "core/session/onnxruntime_cxx_api.h"
#ifdef USE_MKLDNN
#include "core/providers/mkldnn/mkldnn_provider_factory.h"
#endif
#ifdef USE_NGRAPH
#include "core/providers/ngraph/ngraph_provider_factory.h"
#endif
#ifdef USE_NUPHAR
#include "core/providers/nuphar/nuphar_provider_factory.h"
#endif
namespace onnxruntime {
namespace server {
@ -42,8 +60,23 @@ ServerEnvironment::ServerEnvironment(OrtLoggingLevel severity, spdlog::sinks_ini
spdlog::initialize_logger(default_logger_);
}
void ServerEnvironment::RegisterEexcutionProviders(){
#ifdef USE_MKLDNN
Ort::ThrowOnError(OrtSessionOptionsAppendExecutionProvider_Mkldnn(options_, 1));
#endif
#ifdef USE_NGRAPH
Ort::ThrowOnError(OrtSessionOptionsAppendExecutionProvider_NGraph(options_, "CPU"));
#endif
#ifdef USE_NUPHAR
Ort::ThrowOnError(OrtSessionOptionsAppendExecutionProvider_Nuphar(options_, 1, ""));
#endif
}
void ServerEnvironment::InitializeModel(const std::string& model_path, const std::string& model_name, const std::string& model_version) {
auto result = sessions_.emplace(std::piecewise_construct, std::forward_as_tuple(model_name, model_version), std::forward_as_tuple(runtime_environment_, model_path.c_str(), Ort::SessionOptions()));
RegisterEexcutionProviders();
auto result = sessions_.emplace(std::piecewise_construct, std::forward_as_tuple(model_name, model_version), std::forward_as_tuple(runtime_environment_, model_path.c_str(), options_));
if (!result.second) {
throw Ort::Exception("Model of that name already loaded.", ORT_INVALID_ARGUMENT);

View file

@ -28,6 +28,7 @@ class ServerEnvironment {
std::shared_ptr<spdlog::logger> GetLogger(const std::string& request_id) const;
std::shared_ptr<spdlog::logger> GetAppLogger() const;
void UnloadModel(const std::string& model_name, const std::string& model_version);
void RegisterEexcutionProviders();
private:
const OrtLoggingLevel severity_;