203 lines
4.9 KiB
Go
203 lines
4.9 KiB
Go
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 <stdlib.h>
|
|
|
|
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()
|
|
} |