File size: 5,130 Bytes
ba07985
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8104c58
ba07985
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
 
8104c58
ba07985
 
 
 
 
 
 
 
 
 
 
 
 
1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
20
21
22
23
24
25
26
27
28
29
30
31
32
33
34
35
36
37
38
39
40
41
42
43
44
45
46
47
48
49
50
51
52
53
54
55
56
57
58
59
60
61
62
63
64
65
66
67
68
69
70
71
72
73
74
75
76
77
78
79
80
81
82
83
84
85
86
87
88
89
90
91
92
93
94
95
96
97
98
99
100
101
102
103
104
105
106
107
108
109
110
111
112
113
114
115
116
117
118
119
120
121
122
123
124
125
126
127
128
129
130
131
132
133
134
135
#include "ax_engine.hpp"

#include <ax_engine_api.h>
#include <ax_sys_api.h>

#include <cstring>
#include <fstream>
#include <stdexcept>

namespace kantts {

namespace {

std::vector<char> ReadBinary(const std::string& path) {
    std::ifstream f(path, std::ios::binary);
    if (!f) throw std::runtime_error("cannot open " + path);
    return std::vector<char>(std::istreambuf_iterator<char>(f), std::istreambuf_iterator<char>());
}

void Check(int ret, const char* msg) {
    if (ret != 0) throw std::runtime_error(msg);
}

int g_refcount = 0;

}  // namespace

void AxRuntimeInit() {
    if (g_refcount++ > 0) return;
    Check(AX_SYS_Init(), "AX_SYS_Init failed");
    AX_ENGINE_NPU_ATTR_T attr;
    std::memset(&attr, 0, sizeof(attr));
    attr.eHardMode = AX_ENGINE_VIRTUAL_NPU_DISABLE;
    Check(AX_ENGINE_Init(&attr), "AX_ENGINE_Init failed");
}

void AxRuntimeDeinit() {
    if (--g_refcount > 0) return;
    AX_ENGINE_Deinit();
    AX_SYS_Deinit();
}

struct ModelSession::Impl {
    AX_ENGINE_HANDLE handle = nullptr;
    AX_ENGINE_CONTEXT_T context = nullptr;
    AX_ENGINE_IO_INFO_T* info = nullptr;
    AX_ENGINE_IO_T io {};
    std::vector<AX_ENGINE_IO_BUFFER_T> inputs;
    std::vector<AX_ENGINE_IO_BUFFER_T> outputs;
    std::vector<char> model;
    std::map<std::string, AX_U32> input_index;
    std::map<std::string, AX_U32> output_index;

    explicit Impl(const std::string& path) : model(ReadBinary(path)) {
        AX_ENGINE_HANDLE_EXTRA_T extra;
        std::memset(&extra, 0, sizeof(extra));
        Check(AX_ENGINE_CreateHandleV2(&handle, model.data(),
                                       static_cast<AX_U32>(model.size()), &extra),
              "AX_ENGINE_CreateHandleV2 failed");
        Check(AX_ENGINE_CreateContextV2(handle, &context), "AX_ENGINE_CreateContextV2 failed");
        Check(AX_ENGINE_GetIOInfo(handle, &info), "AX_ENGINE_GetIOInfo failed");
        inputs.resize(info->nInputSize);
        outputs.resize(info->nOutputSize);
        io.pInputs = inputs.data();
        io.nInputSize = info->nInputSize;
        io.pOutputs = outputs.data();
        io.nOutputSize = info->nOutputSize;
        for (AX_U32 i = 0; i < info->nInputSize; ++i) {
            std::memset(&inputs[i], 0, sizeof(inputs[i]));
            inputs[i].nSize = info->pInputs[i].nSize;
            Check(AX_SYS_MemAllocCached(&inputs[i].phyAddr, &inputs[i].pVirAddr,
                                        inputs[i].nSize, 128, (AX_S8*)"kantts_in"),
                  "input alloc failed");
            input_index[info->pInputs[i].pName] = i;
        }
        for (AX_U32 i = 0; i < info->nOutputSize; ++i) {
            std::memset(&outputs[i], 0, sizeof(outputs[i]));
            outputs[i].nSize = info->pOutputs[i].nSize;
            Check(AX_SYS_MemAllocCached(&outputs[i].phyAddr, &outputs[i].pVirAddr,
                                        outputs[i].nSize, 128, (AX_S8*)"kantts_out"),
                  "output alloc failed");
            output_index[info->pOutputs[i].pName] = i;
        }
    }

    ~Impl() {
        for (auto& x : inputs) if (x.phyAddr) AX_SYS_MemFree(x.phyAddr, x.pVirAddr);
        for (auto& x : outputs) if (x.phyAddr) AX_SYS_MemFree(x.phyAddr, x.pVirAddr);
        if (handle) AX_ENGINE_DestroyHandle(handle);
    }
};

ModelSession::ModelSession(const std::string& model_path) : impl_(new Impl(model_path)) {}
ModelSession::~ModelSession() { delete impl_; }

void ModelSession::SetInput(const std::string& name, const void* data, size_t bytes) {
    auto it = impl_->input_index.find(name);
    if (it == impl_->input_index.end()) throw std::runtime_error("no input named " + name);
    auto& buf = impl_->inputs[it->second];
    if (bytes > buf.nSize) throw std::runtime_error("input too large " + name);
    std::memcpy(buf.pVirAddr, data, bytes);
    AX_SYS_MflushCache(buf.phyAddr, buf.pVirAddr, buf.nSize);
}

void ModelSession::Run() {
    Check(AX_ENGINE_RunSyncV2(impl_->handle, impl_->context, &impl_->io),
          "AX_ENGINE_RunSyncV2 failed");
}

size_t ModelSession::OutputBytes(const std::string& name) const {
    auto it = impl_->output_index.find(name);
    if (it == impl_->output_index.end()) throw std::runtime_error("no output named " + name);
    return impl_->outputs[it->second].nSize;
}

void ModelSession::GetOutput(const std::string& name, void* out, size_t bytes) const {
    auto it = impl_->output_index.find(name);
    if (it == impl_->output_index.end()) throw std::runtime_error("no output named " + name);
    auto& buf = impl_->outputs[it->second];
    if (bytes > buf.nSize) bytes = buf.nSize;
    AX_SYS_MinvalidateCache(buf.phyAddr, buf.pVirAddr, buf.nSize);
    std::memcpy(out, buf.pVirAddr, bytes);
}

std::vector<int64_t> ModelSession::OutputShape(const std::string& name) const {
    auto it = impl_->output_index.find(name);
    if (it == impl_->output_index.end()) throw std::runtime_error("no output named " + name);
    auto& t = impl_->info->pOutputs[it->second];
    std::vector<int64_t> shape(t.nShapeSize);
    for (AX_U32 i = 0; i < t.nShapeSize; ++i) shape[i] = t.pShape[i];
    return shape;
}

}  // namespace kantts