package onnx /* #cgo CXXFLAGS: -std=c++17 -I/usr/local/onnxruntime/include #cgo linux LDFLAGS: -L/usr/local/onnxruntime/lib -lonnxruntime -ldl -lstdc++ #cgo darwin LDFLAGS: -lonnxruntime -framework CoreFoundation -lstdc++ #include typedef struct OnnxSession OnnxSession; // 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(); */ import "C" import ( "errors" "fmt" "sync" "unsafe" ) // Session ONNX 会话 type Session struct { session *C.OnnxSession mu sync.Mutex } var ( sessions = make(map[string]*Session) mu sync.RWMutex ) // LoadModel 加载 ONNX 模型 func LoadModel(name, path string) error { cPath := C.CString(path) defer C.free(unsafe.Pointer(cPath)) session := C.onnx_create_session(cPath) if session == nil { errMsg := C.GoString(C.onnx_get_last_error()) return fmt.Errorf("加载模型失败: %s", errMsg) } mu.Lock() sessions[name] = &Session{session: session} mu.Unlock() return nil } // GetSession 获取已加载的会话 func GetSession(name string) (*Session, bool) { mu.RLock() s, ok := sessions[name] mu.RUnlock() 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() for _, s := range sessions { s.Close() } sessions = make(map[string]*Session) mu.Unlock() } func init() { C.onnx_init() }