mirror of
https://github.com/saymrwulf/onnxruntime.git
synced 2026-07-30 20:18:08 +00:00
update the code to use C++ API
This commit is contained in:
parent
9ec8b1fd01
commit
906cfd5cf5
1 changed files with 13 additions and 40 deletions
|
|
@ -3,24 +3,19 @@
|
|||
|
||||
// onnxruntime dependencies
|
||||
#include "test_configuration.h"
|
||||
#include <core/session/onnxruntime_c_api.h>
|
||||
#include <core/session/onnxruntime_cxx_api.h>
|
||||
#include <random>
|
||||
#include "command_args_parser.h"
|
||||
#include <google/protobuf/stubs/common.h>
|
||||
|
||||
#include "core/session/onnxruntime_session_options_config_keys.h"
|
||||
#include "core/session/inference_session.h"
|
||||
#include "core/session/ort_env.h"
|
||||
#include "core/providers/provider_factory_creators.h"
|
||||
#include "core/common/logging/sinks/clog_sink.h"
|
||||
|
||||
#include "core/graph/model.h"
|
||||
#include "core/session/environment.h"
|
||||
#include "core/common/logging/logging.h"
|
||||
|
||||
using namespace onnxruntime;
|
||||
const OrtApi* g_ort = NULL;
|
||||
using ProviderOptions = std::unordered_map<std::string, std::string>;
|
||||
|
||||
std::unique_ptr<Ort::Env> ort_env;
|
||||
|
||||
static void CheckStatus(const Status& status) {
|
||||
|
|
@ -110,41 +105,23 @@ int real_main(int argc, wchar_t* argv[]) {
|
|||
#else
|
||||
int real_main(int argc, char* argv[]) {
|
||||
#endif
|
||||
g_ort = OrtGetApiBase()->GetApi(ORT_API_VERSION);
|
||||
qnnctxgen::TestConfig test_config;
|
||||
if (!qnnctxgen::CommandLineParser::ParseArguments(test_config, argc, argv)) {
|
||||
qnnctxgen::CommandLineParser::ShowUsage();
|
||||
return -1;
|
||||
}
|
||||
|
||||
{
|
||||
bool failed = false;
|
||||
ORT_TRY {
|
||||
OrtLoggingLevel logging_level = test_config.run_config.f_verbose
|
||||
? ORT_LOGGING_LEVEL_VERBOSE
|
||||
: ORT_LOGGING_LEVEL_WARNING;
|
||||
|
||||
ort_env = std::make_unique<Ort::Env>(logging_level, "Default");
|
||||
}
|
||||
ORT_CATCH(const Ort::Exception& e) {
|
||||
ORT_HANDLE_EXCEPTION([&]() {
|
||||
fprintf(stderr, "Error creating environment. Error-> %s \n", e.what());
|
||||
failed = true;
|
||||
});
|
||||
}
|
||||
|
||||
if (failed)
|
||||
return -1;
|
||||
}
|
||||
OrtLoggingLevel logging_level = test_config.run_config.f_verbose
|
||||
? ORT_LOGGING_LEVEL_VERBOSE
|
||||
: ORT_LOGGING_LEVEL_WARNING;
|
||||
Ort::Env env(logging_level, "ep_weight_sharing");
|
||||
|
||||
ORT_TRY {
|
||||
SessionOptions so;
|
||||
so.session_logid = "ep_weight_sharing_ctx_gen_session_logger";
|
||||
Ort::SessionOptions so;
|
||||
so.SetLogId("ep_weight_sharing_ctx_gen_session_logger");
|
||||
// Set default session option to dump QNN context model with non-embed mode
|
||||
CheckStatus(so.config_options.AddConfigEntry(kOrtSessionOptionEpContextEnable, "1"));
|
||||
CheckStatus(so.config_options.AddConfigEntry(kOrtSessionOptionEpContextEmbedMode, "0"));
|
||||
RunOptions run_options;
|
||||
run_options.run_tag = so.session_logid;
|
||||
so.AddConfigEntry(kOrtSessionOptionEpContextEnable, "1");
|
||||
so.AddConfigEntry(kOrtSessionOptionEpContextEmbedMode, "0");
|
||||
|
||||
ProviderOptions provider_options;
|
||||
#if defined(_WIN32)
|
||||
|
|
@ -172,7 +149,7 @@ int real_main(int argc, char* argv[]) {
|
|||
std::cout << "Not support to specify the generated Onnx context cache file name." << std::endl;
|
||||
continue;
|
||||
}
|
||||
CheckStatus(so.config_options.AddConfigEntry(it.first.c_str(), it.second.c_str()));
|
||||
so.AddConfigEntry(it.first.c_str(), it.second.c_str());
|
||||
}
|
||||
|
||||
for (auto model_path : test_config.model_file_paths) {
|
||||
|
|
@ -182,15 +159,11 @@ int real_main(int argc, char* argv[]) {
|
|||
// Generate context cache model files with QNN context binary files
|
||||
// The context binary file generated later includes all graphs from previous models
|
||||
{
|
||||
auto ep = QNNProviderFactoryCreator::Create(provider_options, &so)->CreateProvider();
|
||||
std::shared_ptr<IExecutionProvider> qnn_ep(std::move(ep));
|
||||
so.AppendExecutionProvider("QNN", provider_options);
|
||||
|
||||
for (auto model_path : test_config.model_file_paths) {
|
||||
std::cout << "Generate context cache model for: " << ToUTF8String(model_path) << std::endl;
|
||||
InferenceSession session_object{so, ((OrtEnv*)*ort_env.get())->GetEnvironment()};
|
||||
CheckStatus(session_object.RegisterExecutionProvider(qnn_ep));
|
||||
CheckStatus(session_object.Load(ToPathString(model_path)));
|
||||
CheckStatus(session_object.Initialize());
|
||||
Ort::Session session(env, model_path.c_str(), so);
|
||||
}
|
||||
}
|
||||
|
||||
|
|
|
|||
Loading…
Reference in a new issue