#include #include #include #include #include #include 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 input_names; std::vector output_names; std::vector input_name_ptrs; std::vector 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; }