feat: 实现 ONNX Runtime CGO 绑定和 OCR/Math 推理逻辑
Build and Deploy / build (push) Failing after 2m23s
Build and Deploy / deploy (push) Has been skipped

This commit is contained in:
2026-07-16 21:14:34 +00:00
parent 5dacb61122
commit 7afefe5e9f
6 changed files with 1094 additions and 85 deletions
+424
View File
@@ -0,0 +1,424 @@
#include <onnxruntime_c_api.h>
#include <vector>
#include <string>
#include <cstring>
#include <memory>
static const OrtApi* g_ort = nullptr;
static std::string last_error;
// 初始化 ONNX Runtime
extern "C" int onnx_init() {
g_ort = OrtGetApiBase()->GetApi(ORT_API_VERSION);
if (!g_ort) {
last_error = "Failed to get ONNX Runtime API";
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();
}
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, GraphOptimizationLevel::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
status = g_ort->CreateSession(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;
}
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(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;
}
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(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;
}
+144 -59
View File
@@ -1,35 +1,24 @@
package onnx
/*
#cgo CXXFLAGS: -std=c++17
#cgo linux LDFLAGS: -lonnxruntime -ldl
#cgo darwin LDFLAGS: -lonnxruntime -framework CoreFoundation
#cgo CXXFLAGS: -std=c++17 -I/usr/include/onnxruntime
#cgo linux LDFLAGS: -lonnxruntime -ldl -lstdc++
#cgo darwin LDFLAGS: -lonnxruntime -framework CoreFoundation -lstdc++
#include <stdlib.h>
#include <string.h>
// ONNX Runtime C API 声明
#ifdef __cplusplus
extern "C" {
#endif
typedef struct OnnxSession OnnxSession;
typedef void* OrtSession;
typedef void* OrtMemoryInfo;
typedef void* OrtValue;
// 初始化 ONNX 会话
OrtSession onnx_create_session(const char* model_path);
void onnx_destroy_session(OrtSession session);
// 运行推理
int onnx_run(OrtSession session, const float* input_data, int input_size, float* output_data, int output_size);
// 错误信息
// ONNX Runtime 函数
int onnx_init();
OnnxSession* onnx_create_session(const char* model_path);
void onnx_destroy_session(OnnxSession* session);
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);
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);
int onnx_get_input_shape(OnnxSession* sess, int input_index, int64_t* dims, int max_dims);
int onnx_get_output_shape(OnnxSession* sess, int output_index, int64_t* dims, int max_dims);
void onnx_free(void* ptr);
const char* onnx_get_last_error();
#ifdef __cplusplus
}
#endif
*/
import "C"
import (
@@ -39,8 +28,9 @@ import (
"unsafe"
)
// Session ONNX 会话
type Session struct {
session C.OrtSession
session *C.OnnxSession
mu sync.Mutex
}
@@ -56,7 +46,8 @@ func LoadModel(name, path string) error {
session := C.onnx_create_session(cPath)
if session == nil {
return fmt.Errorf("加载模型失败: %s", C.GoString(C.onnx_get_last_error()))
errMsg := C.GoString(C.onnx_get_last_error())
return fmt.Errorf("加载模型失败: %s", errMsg)
}
mu.Lock()
@@ -66,39 +57,6 @@ func LoadModel(name, path string) error {
return nil
}
// Run 执行模型推理
func (s *Session) Run(input []float32, inputSize int) ([]float32, error) {
s.mu.Lock()
defer s.mu.Unlock()
output := make([]float32, inputSize)
result := C.onnx_run(
s.session,
(*C.float)(unsafe.Pointer(&input[0])),
C.int(len(input)),
(*C.float)(unsafe.Pointer(&output[0])),
C.int(len(output)),
)
if result != 0 {
return nil, errors.New(C.GoString(C.onnx_get_last_error()))
}
return output, nil
}
// Close 关闭会话
func (s *Session) Close() {
s.mu.Lock()
defer s.mu.Unlock()
if s.session != nil {
C.onnx_destroy_session(s.session)
s.session = nil
}
}
// GetSession 获取已加载的会话
func GetSession(name string) (*Session, bool) {
mu.RLock()
@@ -107,6 +65,129 @@ func GetSession(name string) (*Session, bool) {
return s, ok
}
// Run 执行单输入推理
func (s *Session) Run(input []float32, dims []int64) ([]float32, error) {
s.mu.Lock()
defer s.mu.Unlock()
if len(input) == 0 {
return nil, errors.New("输入数据为空")
}
var outputSize C.int
output := C.onnx_run_float(
s.session,
(*C.float)(unsafe.Pointer(&input[0])),
C.int(len(input)),
(*C.int64_t)(unsafe.Pointer(&dims[0])),
C.int(len(dims)),
&outputSize,
)
if output == nil {
errMsg := C.GoString(C.onnx_get_last_error())
return nil, fmt.Errorf("推理失败: %s", errMsg)
}
defer C.onnx_free(unsafe.Pointer(output))
// 复制输出数据
result := make([]float32, int(outputSize))
outputSlice := (*[1 << 30]float32)(unsafe.Pointer(output))[:int(outputSize):int(outputSize)]
for i := 0; i < int(outputSize); i++ {
result[i] = outputSlice[i]
}
return result, nil
}
// RunDualInput 执行双输入推理(用于 Siamese 网络)
func (s *Session) RunDualInput(input1 []float32, dims1 []int64, input2 []float32, dims2 []int64) ([]float32, error) {
s.mu.Lock()
defer s.mu.Unlock()
if len(input1) == 0 || len(input2) == 0 {
return nil, errors.New("输入数据为空")
}
var outputSize C.int
output := C.onnx_run_dual_input(
s.session,
(*C.float)(unsafe.Pointer(&input1[0])),
C.int(len(input1)),
(*C.int64_t)(unsafe.Pointer(&dims1[0])),
C.int(len(dims1)),
(*C.float)(unsafe.Pointer(&input2[0])),
C.int(len(input2)),
(*C.int64_t)(unsafe.Pointer(&dims2[0])),
C.int(len(dims2)),
&outputSize,
)
if output == nil {
errMsg := C.GoString(C.onnx_get_last_error())
return nil, fmt.Errorf("推理失败: %s", errMsg)
}
defer C.onnx_free(unsafe.Pointer(output))
// 复制输出数据
result := make([]float32, int(outputSize))
outputSlice := (*[1 << 30]float32)(unsafe.Pointer(output))[:int(outputSize):int(outputSize)]
for i := 0; i < int(outputSize); i++ {
result[i] = outputSlice[i]
}
return result, nil
}
// GetInputShape 获取输入形状
func (s *Session) GetInputShape(index int) ([]int64, error) {
var dims [8]C.int64_t
dimCount := C.onnx_get_input_shape(s.session, C.int(index), &dims[0], 8)
if dimCount < 0 {
return nil, fmt.Errorf("获取输入形状失败: %s", C.GoString(C.onnx_get_last_error()))
}
result := make([]int64, dimCount)
for i := 0; i < int(dimCount); i++ {
result[i] = int64(dims[i])
}
return result, nil
}
// GetOutputShape 获取输出形状
func (s *Session) GetOutputShape(index int) ([]int64, error) {
var dims [8]C.int64_t
dimCount := C.onnx_get_output_shape(s.session, C.int(index), &dims[0], 8)
if dimCount < 0 {
return nil, fmt.Errorf("获取输出形状失败: %s", C.GoString(C.onnx_get_last_error()))
}
result := make([]int64, dimCount)
for i := 0; i < int(dimCount); i++ {
result[i] = int64(dims[i])
}
return result, nil
}
// Close 关闭会话
func (s *Session) Close() {
s.mu.Lock()
defer s.mu.Unlock()
if s.session != nil {
C.onnx_destroy_session(s.session)
s.session = nil
}
}
// CloseAll 关闭所有会话
func CloseAll() {
mu.Lock()
@@ -115,4 +196,8 @@ func CloseAll() {
}
sessions = make(map[string]*Session)
mu.Unlock()
}
func init() {
C.onnx_init()
}