146 lines
3.2 KiB
Go
146 lines
3.2 KiB
Go
package rknn
|
|
|
|
/*
|
|
#cgo linux LDFLAGS: -lrknnrt
|
|
#cgo linux,arm64 LDFLAGS: -lrknnrt
|
|
|
|
#include <stdint.h>
|
|
#include <stdlib.h>
|
|
|
|
typedef void* rknn_context;
|
|
|
|
typedef enum {
|
|
RKNN_QUERY_IN_OUT_NUM = 0,
|
|
RKNN_QUERY_INPUT_ATTR = 1,
|
|
RKNN_QUERY_OUTPUT_ATTR = 2,
|
|
RKNN_QUERY_SDK_VERSION = 5,
|
|
} rknn_query_cmd;
|
|
|
|
typedef enum {
|
|
RKNN_TENSOR_UINT8 = 0,
|
|
RKNN_TENSOR_FLOAT32 = 2,
|
|
} rknn_tensor_type;
|
|
|
|
typedef enum {
|
|
RKNN_TENSOR_NCHW = 0,
|
|
RKNN_TENSOR_NHWC = 1,
|
|
} rknn_tensor_format;
|
|
|
|
typedef struct {
|
|
uint32_t index;
|
|
int32_t type;
|
|
int32_t fmt;
|
|
void* buf;
|
|
uint32_t size;
|
|
uint8_t want_float;
|
|
uint8_t is_prealloc;
|
|
} rknn_output;
|
|
|
|
typedef struct {
|
|
uint32_t n_input;
|
|
uint32_t n_output;
|
|
} rknn_input_output_num;
|
|
|
|
typedef struct {
|
|
void* buf;
|
|
uint32_t size;
|
|
uint8_t pass_through;
|
|
uint32_t type;
|
|
uint32_t fmt;
|
|
} rknn_input;
|
|
|
|
extern int rknn_init(rknn_context* ctx, void* model, uint32_t size, uint32_t flag, rknn_input_output_num* io_num);
|
|
extern int rknn_query(rknn_context ctx, rknn_query_cmd cmd, void* info, uint32_t size);
|
|
extern int rknn_inputs_set(rknn_context ctx, uint32_t n_inputs, rknn_input inputs[]);
|
|
extern int rknn_run(rknn_context ctx, void* ext);
|
|
extern int rknn_outputs_get(rknn_context ctx, uint32_t n_outputs, rknn_output outputs[], void* ext);
|
|
extern int rknn_outputs_release(rknn_context ctx, uint32_t n_outputs, rknn_output outputs[]);
|
|
extern int rknn_destroy(rknn_context ctx);
|
|
*/
|
|
import "C"
|
|
import (
|
|
"fmt"
|
|
"os"
|
|
"unsafe"
|
|
)
|
|
|
|
type Context struct {
|
|
ctx C.rknn_context
|
|
nInput uint32
|
|
nOutput uint32
|
|
}
|
|
|
|
func LoadModel(modelPath string) (*Context, error) {
|
|
data, err := os.ReadFile(modelPath)
|
|
if err != nil {
|
|
return nil, fmt.Errorf("rknn: read model file: %w", err)
|
|
}
|
|
|
|
var ctx C.rknn_context
|
|
var ioNum C.rknn_input_output_num
|
|
|
|
ret := C.rknn_init(
|
|
&ctx,
|
|
unsafe.Pointer(&data[0]),
|
|
C.uint32_t(len(data)),
|
|
0,
|
|
&ioNum,
|
|
)
|
|
if ret != 0 {
|
|
return nil, fmt.Errorf("rknn_init: error %d", int(ret))
|
|
}
|
|
|
|
return &Context{
|
|
ctx: ctx,
|
|
nInput: uint32(ioNum.n_input),
|
|
nOutput: uint32(ioNum.n_output),
|
|
}, nil
|
|
}
|
|
|
|
func (c *Context) InferenceRGB(inputData []uint8, height, width int) ([][]float32, error) {
|
|
rknnIn := C.rknn_input{
|
|
buf: unsafe.Pointer(&inputData[0]),
|
|
size: C.uint32_t(len(inputData)),
|
|
pass_through: 0,
|
|
_type: C.RKNN_TENSOR_UINT8,
|
|
fmt: C.RKNN_TENSOR_NHWC,
|
|
}
|
|
|
|
ret := C.rknn_inputs_set(c.ctx, 1, &rknnIn)
|
|
if ret < 0 {
|
|
return nil, fmt.Errorf("rknn_inputs_set: error %d", int(ret))
|
|
}
|
|
|
|
ret = C.rknn_run(c.ctx, nil)
|
|
if ret < 0 {
|
|
return nil, fmt.Errorf("rknn_run: error %d", int(ret))
|
|
}
|
|
|
|
outputs := make([]C.rknn_output, c.nOutput)
|
|
for i := range outputs {
|
|
outputs[i].want_float = 1
|
|
}
|
|
|
|
ret = C.rknn_outputs_get(c.ctx, C.uint32_t(c.nOutput), &outputs[0], nil)
|
|
if ret < 0 {
|
|
return nil, fmt.Errorf("rknn_outputs_get: error %d", int(ret))
|
|
}
|
|
|
|
results := make([][]float32, c.nOutput)
|
|
for i := range outputs {
|
|
nVals := int(outputs[i].size) / 4
|
|
vals := make([]float32, nVals)
|
|
src := unsafe.Slice((*float32)(outputs[i].buf), nVals)
|
|
copy(vals, src)
|
|
results[i] = vals
|
|
}
|
|
|
|
C.rknn_outputs_release(c.ctx, C.uint32_t(c.nOutput), &outputs[0])
|
|
|
|
return results, nil
|
|
}
|
|
|
|
func (c *Context) Release() {
|
|
C.rknn_destroy(c.ctx)
|
|
}
|