fix: ONNX Runtime API 兼容 - 添加 OrtEnv 环境
This commit is contained in:
+19
-3
@@ -6,15 +6,27 @@
|
||||
#include <memory>
|
||||
|
||||
static const OrtApi* g_ort = nullptr;
|
||||
static OrtEnv* g_env = nullptr;
|
||||
static std::string last_error;
|
||||
|
||||
// 初始化 ONNX Runtime
|
||||
extern "C" int onnx_init() {
|
||||
if (g_ort) return 0; // 已初始化
|
||||
|
||||
g_ort = OrtGetApiBase()->GetApi(ORT_API_VERSION);
|
||||
if (!g_ort) {
|
||||
last_error = "Failed to get ONNX Runtime API";
|
||||
return -1;
|
||||
}
|
||||
|
||||
// 创建环境
|
||||
OrtStatus* status = g_ort->CreateEnv(ORT_LOGGING_LEVEL_WARNING, "anticaptcha", &g_env);
|
||||
if (status) {
|
||||
last_error = g_ort->GetErrorMessage(status);
|
||||
g_ort->ReleaseStatus(status);
|
||||
return -1;
|
||||
}
|
||||
|
||||
return 0;
|
||||
}
|
||||
|
||||
@@ -34,6 +46,10 @@ extern "C" OnnxSession* onnx_create_session(const char* model_path) {
|
||||
if (!g_ort) {
|
||||
onnx_init();
|
||||
}
|
||||
if (!g_env) {
|
||||
last_error = "ONNX Runtime environment not initialized";
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
auto* sess = new OnnxSession();
|
||||
|
||||
@@ -48,7 +64,7 @@ extern "C" OnnxSession* onnx_create_session(const char* model_path) {
|
||||
|
||||
// 设置 CPU 线程数
|
||||
g_ort->SetIntraOpNumThreads(sess->session_options, 4);
|
||||
g_ort->SetSessionGraphOptimizationLevel(sess->session_options, GraphOptimizationLevel::ORT_ENABLE_EXTENDED);
|
||||
g_ort->SetSessionGraphOptimizationLevel(sess->session_options, ORT_ENABLE_EXTENDED);
|
||||
|
||||
// 创建 memory info
|
||||
status = g_ort->CreateCpuMemoryInfo(OrtArenaAllocator, OrtMemTypeDefault, &sess->memory_info);
|
||||
@@ -60,8 +76,8 @@ extern "C" OnnxSession* onnx_create_session(const char* model_path) {
|
||||
return nullptr;
|
||||
}
|
||||
|
||||
// 创建 session
|
||||
status = g_ort->CreateSession(model_path, sess->session_options, &sess->session);
|
||||
// 创建 session (需要 env)
|
||||
status = g_ort->CreateSession(g_env, model_path, sess->session_options, &sess->session);
|
||||
if (status) {
|
||||
last_error = g_ort->GetErrorMessage(status);
|
||||
g_ort->ReleaseStatus(status);
|
||||
|
||||
Reference in New Issue
Block a user