441 lines
12 KiB
C++
441 lines
12 KiB
C++
#include <fstream>
|
|
#include <onnxruntime_c_api.h>
|
|
#include <vector>
|
|
#include <string>
|
|
#include <cstring>
|
|
#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;
|
|
}
|
|
|
|
// ONNX Session 结构
|
|
typedef struct {
|
|
OrtSession* session;
|
|
OrtSessionOptions* session_options;
|
|
OrtMemoryInfo* memory_info;
|
|
std::vector<std::string> input_names;
|
|
std::vector<std::string> output_names;
|
|
std::vector<const char*> input_name_ptrs;
|
|
std::vector<const char*> output_name_ptrs;
|
|
} OnnxSession;
|
|
|
|
// 创建 Session
|
|
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();
|
|
|
|
// 创建 session options
|
|
OrtStatus* status = g_ort->CreateSessionOptions(&sess->session_options);
|
|
if (status) {
|
|
last_error = g_ort->GetErrorMessage(status);
|
|
g_ort->ReleaseStatus(status);
|
|
delete sess;
|
|
return nullptr;
|
|
}
|
|
|
|
// 设置 CPU 线程数
|
|
g_ort->SetIntraOpNumThreads(sess->session_options, 4);
|
|
g_ort->SetSessionGraphOptimizationLevel(sess->session_options, ORT_ENABLE_EXTENDED);
|
|
|
|
// 创建 memory info
|
|
status = g_ort->CreateCpuMemoryInfo(OrtArenaAllocator, OrtMemTypeDefault, &sess->memory_info);
|
|
if (status) {
|
|
last_error = g_ort->GetErrorMessage(status);
|
|
g_ort->ReleaseStatus(status);
|
|
g_ort->ReleaseSessionOptions(sess->session_options);
|
|
delete sess;
|
|
return nullptr;
|
|
}
|
|
|
|
// 创建 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);
|
|
g_ort->ReleaseMemoryInfo(sess->memory_info);
|
|
g_ort->ReleaseSessionOptions(sess->session_options);
|
|
delete sess;
|
|
return nullptr;
|
|
}
|
|
|
|
// 获取输入输出名称
|
|
OrtAllocator* allocator = nullptr;
|
|
status = g_ort->GetAllocatorWithDefaultOptions(&allocator);
|
|
if (status) {
|
|
last_error = g_ort->GetErrorMessage(status);
|
|
g_ort->ReleaseStatus(status);
|
|
} else {
|
|
// 获取输入数量和名称
|
|
size_t num_inputs = 0;
|
|
g_ort->SessionGetInputCount(sess->session, &num_inputs);
|
|
sess->input_names.resize(num_inputs);
|
|
sess->input_name_ptrs.resize(num_inputs);
|
|
|
|
for (size_t i = 0; i < num_inputs; i++) {
|
|
char* name = nullptr;
|
|
g_ort->SessionGetInputName(sess->session, i, allocator, &name);
|
|
sess->input_names[i] = name;
|
|
sess->input_name_ptrs[i] = sess->input_names[i].c_str();
|
|
g_ort->AllocatorFree(allocator, name);
|
|
}
|
|
|
|
// 获取输出数量和名称
|
|
size_t num_outputs = 0;
|
|
g_ort->SessionGetOutputCount(sess->session, &num_outputs);
|
|
sess->output_names.resize(num_outputs);
|
|
sess->output_name_ptrs.resize(num_outputs);
|
|
|
|
for (size_t i = 0; i < num_outputs; i++) {
|
|
char* name = nullptr;
|
|
g_ort->SessionGetOutputName(sess->session, i, allocator, &name);
|
|
sess->output_names[i] = name;
|
|
sess->output_name_ptrs[i] = sess->output_names[i].c_str();
|
|
g_ort->AllocatorFree(allocator, name);
|
|
}
|
|
}
|
|
|
|
return sess;
|
|
}
|
|
|
|
// 销毁 Session
|
|
extern "C" void onnx_destroy_session(OnnxSession* sess) {
|
|
if (sess) {
|
|
if (sess->session) g_ort->ReleaseSession(sess->session);
|
|
if (sess->memory_info) g_ort->ReleaseMemoryInfo(sess->memory_info);
|
|
if (sess->session_options) g_ort->ReleaseSessionOptions(sess->session_options);
|
|
delete sess;
|
|
}
|
|
}
|
|
|
|
// 运行推理 - 支持动态输入
|
|
extern "C" float* onnx_run_float(
|
|
OnnxSession* sess,
|
|
const float* input_data,
|
|
int input_size,
|
|
const int64_t* input_dims,
|
|
int input_dim_count,
|
|
int* output_size
|
|
) {
|
|
if (!sess || !sess->session || !input_data) {
|
|
last_error = "Invalid session or input";
|
|
return nullptr;
|
|
}
|
|
|
|
OrtStatus* status = nullptr;
|
|
|
|
// 创建输入 tensor
|
|
OrtValue* input_tensor = nullptr;
|
|
status = g_ort->CreateTensorWithDataAsOrtValue(
|
|
sess->memory_info,
|
|
(void*)input_data,
|
|
input_size * sizeof(float),
|
|
input_dims,
|
|
input_dim_count,
|
|
ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT,
|
|
&input_tensor
|
|
);
|
|
|
|
if (status) {
|
|
last_error = g_ort->GetErrorMessage(status);
|
|
g_ort->ReleaseStatus(status);
|
|
return nullptr;
|
|
}
|
|
|
|
// 运行推理
|
|
OrtValue* output_tensor = nullptr;
|
|
status = g_ort->Run(
|
|
sess->session,
|
|
nullptr,
|
|
sess->input_name_ptrs.data(),
|
|
&input_tensor,
|
|
1,
|
|
sess->output_name_ptrs.data(),
|
|
1,
|
|
&output_tensor
|
|
);
|
|
|
|
g_ort->ReleaseValue(input_tensor);
|
|
|
|
if (status) {
|
|
last_error = g_ort->GetErrorMessage(status);
|
|
g_ort->ReleaseStatus(status);
|
|
return nullptr;
|
|
}
|
|
|
|
// 获取输出数据
|
|
float* output_data = nullptr;
|
|
status = g_ort->GetTensorMutableData(output_tensor, (void**)&output_data);
|
|
if (status) {
|
|
last_error = g_ort->GetErrorMessage(status);
|
|
g_ort->ReleaseStatus(status);
|
|
g_ort->ReleaseValue(output_tensor);
|
|
return nullptr;
|
|
}
|
|
|
|
// 获取输出大小
|
|
OrtTensorTypeAndShapeInfo* type_info = nullptr;
|
|
status = g_ort->GetTensorTypeAndShape(output_tensor, &type_info);
|
|
if (status) {
|
|
last_error = g_ort->GetErrorMessage(status);
|
|
g_ort->ReleaseStatus(status);
|
|
g_ort->ReleaseValue(output_tensor);
|
|
return nullptr;
|
|
}
|
|
|
|
size_t element_count = 0;
|
|
g_ort->GetTensorShapeElementCount(type_info, &element_count);
|
|
*output_size = (int)element_count;
|
|
|
|
// 复制输出数据
|
|
float* result = (float*)malloc(element_count * sizeof(float));
|
|
memcpy(result, output_data, element_count * sizeof(float));
|
|
|
|
g_ort->ReleaseTensorTypeAndShapeInfo(type_info);
|
|
g_ort->ReleaseValue(output_tensor);
|
|
|
|
return result;
|
|
}
|
|
|
|
// 获取输入形状
|
|
extern "C" int onnx_get_input_shape(
|
|
OnnxSession* sess,
|
|
int input_index,
|
|
int64_t* dims,
|
|
int max_dims
|
|
) {
|
|
if (!sess || !sess->session) {
|
|
last_error = "Invalid session";
|
|
return -1;
|
|
}
|
|
|
|
OrtTypeInfo* type_info = nullptr;
|
|
OrtStatus* status = g_ort->SessionGetInputTypeInfo(sess->session, input_index, &type_info);
|
|
if (status) {
|
|
last_error = g_ort->GetErrorMessage(status);
|
|
g_ort->ReleaseStatus(status);
|
|
return -1;
|
|
}
|
|
|
|
const OrtTensorTypeAndShapeInfo* tensor_info = nullptr;
|
|
status = g_ort->CastTypeInfoToTensorInfo(type_info, &tensor_info);
|
|
if (status) {
|
|
last_error = g_ort->GetErrorMessage(status);
|
|
g_ort->ReleaseStatus(status);
|
|
g_ort->ReleaseTypeInfo(type_info);
|
|
return -1;
|
|
}
|
|
|
|
size_t dim_count = 0;
|
|
g_ort->GetDimensionsCount(tensor_info, &dim_count);
|
|
|
|
if ((int)dim_count > max_dims) {
|
|
dim_count = max_dims;
|
|
}
|
|
|
|
g_ort->GetDimensions(tensor_info, dims, dim_count);
|
|
|
|
g_ort->ReleaseTensorTypeAndShapeInfo((OrtTensorTypeAndShapeInfo*)tensor_info);
|
|
g_ort->ReleaseTypeInfo(type_info);
|
|
|
|
return (int)dim_count;
|
|
}
|
|
|
|
// 获取输出形状
|
|
extern "C" int onnx_get_output_shape(
|
|
OnnxSession* sess,
|
|
int output_index,
|
|
int64_t* dims,
|
|
int max_dims
|
|
) {
|
|
if (!sess || !sess->session) {
|
|
last_error = "Invalid session";
|
|
return -1;
|
|
}
|
|
|
|
OrtTypeInfo* type_info = nullptr;
|
|
OrtStatus* status = g_ort->SessionGetOutputTypeInfo(sess->session, output_index, &type_info);
|
|
if (status) {
|
|
last_error = g_ort->GetErrorMessage(status);
|
|
g_ort->ReleaseStatus(status);
|
|
return -1;
|
|
}
|
|
|
|
const OrtTensorTypeAndShapeInfo* tensor_info = nullptr;
|
|
status = g_ort->CastTypeInfoToTensorInfo(type_info, &tensor_info);
|
|
if (status) {
|
|
last_error = g_ort->GetErrorMessage(status);
|
|
g_ort->ReleaseStatus(status);
|
|
g_ort->ReleaseTypeInfo(type_info);
|
|
return -1;
|
|
}
|
|
|
|
size_t dim_count = 0;
|
|
g_ort->GetDimensionsCount(tensor_info, &dim_count);
|
|
|
|
if ((int)dim_count > max_dims) {
|
|
dim_count = max_dims;
|
|
}
|
|
|
|
g_ort->GetDimensions(tensor_info, dims, dim_count);
|
|
|
|
g_ort->ReleaseTensorTypeAndShapeInfo((OrtTensorTypeAndShapeInfo*)tensor_info);
|
|
g_ort->ReleaseTypeInfo(type_info);
|
|
|
|
return (int)dim_count;
|
|
}
|
|
|
|
// 释放内存
|
|
extern "C" void onnx_free(void* ptr) {
|
|
if (ptr) {
|
|
free(ptr);
|
|
}
|
|
}
|
|
|
|
// 获取错误信息
|
|
extern "C" const char* onnx_get_last_error() {
|
|
return last_error.c_str();
|
|
}
|
|
|
|
// 双输入推理 (用于 Siamese 网络)
|
|
extern "C" float* onnx_run_dual_input(
|
|
OnnxSession* sess,
|
|
const float* input1_data,
|
|
int input1_size,
|
|
const int64_t* input1_dims,
|
|
int input1_dim_count,
|
|
const float* input2_data,
|
|
int input2_size,
|
|
const int64_t* input2_dims,
|
|
int input2_dim_count,
|
|
int* output_size
|
|
) {
|
|
if (!sess || !sess->session || !input1_data || !input2_data) {
|
|
last_error = "Invalid session or inputs";
|
|
return nullptr;
|
|
}
|
|
|
|
OrtStatus* status = nullptr;
|
|
|
|
// 创建输入 tensor 1
|
|
OrtValue* input_tensor1 = nullptr;
|
|
status = g_ort->CreateTensorWithDataAsOrtValue(
|
|
sess->memory_info,
|
|
(void*)input1_data,
|
|
input1_size * sizeof(float),
|
|
input1_dims,
|
|
input1_dim_count,
|
|
ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT,
|
|
&input_tensor1
|
|
);
|
|
if (status) {
|
|
last_error = g_ort->GetErrorMessage(status);
|
|
g_ort->ReleaseStatus(status);
|
|
return nullptr;
|
|
}
|
|
|
|
// 创建输入 tensor 2
|
|
OrtValue* input_tensor2 = nullptr;
|
|
status = g_ort->CreateTensorWithDataAsOrtValue(
|
|
sess->memory_info,
|
|
(void*)input2_data,
|
|
input2_size * sizeof(float),
|
|
input2_dims,
|
|
input2_dim_count,
|
|
ONNX_TENSOR_ELEMENT_DATA_TYPE_FLOAT,
|
|
&input_tensor2
|
|
);
|
|
if (status) {
|
|
last_error = g_ort->GetErrorMessage(status);
|
|
g_ort->ReleaseStatus(status);
|
|
g_ort->ReleaseValue(input_tensor1);
|
|
return nullptr;
|
|
}
|
|
|
|
// 准备输入
|
|
const char* input_names[] = {sess->input_name_ptrs[0], sess->input_name_ptrs[1]};
|
|
OrtValue* input_tensors[] = {input_tensor1, input_tensor2};
|
|
|
|
// 运行推理
|
|
OrtValue* output_tensor = nullptr;
|
|
status = g_ort->Run(
|
|
sess->session,
|
|
nullptr,
|
|
input_names,
|
|
input_tensors,
|
|
2,
|
|
sess->output_name_ptrs.data(),
|
|
1,
|
|
&output_tensor
|
|
);
|
|
|
|
g_ort->ReleaseValue(input_tensor1);
|
|
g_ort->ReleaseValue(input_tensor2);
|
|
|
|
if (status) {
|
|
last_error = g_ort->GetErrorMessage(status);
|
|
g_ort->ReleaseStatus(status);
|
|
return nullptr;
|
|
}
|
|
|
|
// 获取输出数据
|
|
float* output_data = nullptr;
|
|
status = g_ort->GetTensorMutableData(output_tensor, (void**)&output_data);
|
|
if (status) {
|
|
last_error = g_ort->GetErrorMessage(status);
|
|
g_ort->ReleaseStatus(status);
|
|
g_ort->ReleaseValue(output_tensor);
|
|
return nullptr;
|
|
}
|
|
|
|
// 获取输出大小
|
|
OrtTensorTypeAndShapeInfo* type_info = nullptr;
|
|
status = g_ort->GetTensorTypeAndShape(output_tensor, &type_info);
|
|
if (status) {
|
|
last_error = g_ort->GetErrorMessage(status);
|
|
g_ort->ReleaseStatus(status);
|
|
g_ort->ReleaseValue(output_tensor);
|
|
return nullptr;
|
|
}
|
|
|
|
size_t element_count = 0;
|
|
g_ort->GetTensorShapeElementCount(type_info, &element_count);
|
|
*output_size = (int)element_count;
|
|
|
|
// 复制输出数据
|
|
float* result = (float*)malloc(element_count * sizeof(float));
|
|
memcpy(result, output_data, element_count * sizeof(float));
|
|
|
|
g_ort->ReleaseTensorTypeAndShapeInfo(type_info);
|
|
g_ort->ReleaseValue(output_tensor);
|
|
|
|
return result;
|
|
} |