Files
admin 5c3b7ef81d
Build and Deploy / build (push) Failing after 2m55s
Build and Deploy / deploy (push) Has been skipped
fix: CastTypeInfoToTensorInfo 使用 const 指针
2026-07-16 22:37:14 +00:00

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;
}