File size: 6,793 Bytes
5eee449 | 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 136 137 138 139 140 141 142 143 144 145 146 147 148 149 150 151 152 153 154 155 156 157 158 159 160 161 162 163 164 165 166 167 168 169 170 171 172 173 174 175 176 177 178 179 180 181 182 183 184 185 186 187 188 189 | #include "ax_runner.h"
#include <stdexcept>
#ifdef INFLECT_WITH_AX_ENGINE
#include <ax_engine_api.h>
#include <ax_sys_api.h>
#include <cstring>
#include <fstream>
#include <iterator>
namespace {
std::vector<char> read_binary(const std::string& path) {
std::ifstream file(path, std::ios::binary);
if (!file) {
throw std::runtime_error("failed to open " + path);
}
return std::vector<char>(std::istreambuf_iterator<char>(file),
std::istreambuf_iterator<char>());
}
void check_ax(int ret, const char* message) {
if (ret != 0) {
throw std::runtime_error(message);
}
}
} // namespace
struct AxRunner::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;
explicit Impl(const std::string& model_path) : model(read_binary(model_path)) {
check_ax(AX_SYS_Init(), "AX_SYS_Init failed");
AX_ENGINE_NPU_ATTR_T npu_attr;
std::memset(&npu_attr, 0, sizeof(npu_attr));
npu_attr.eHardMode = static_cast<AX_ENGINE_NPU_MODE_T>(0);
check_ax(AX_ENGINE_Init(&npu_attr), "AX_ENGINE_Init failed");
AX_ENGINE_HANDLE_EXTRA_T extra;
std::memset(&extra, 0, sizeof(extra));
char model_name[] = "inflect_tts";
extra.pName = reinterpret_cast<AX_S8*>(model_name);
check_ax(AX_ENGINE_CreateHandleV2(&handle, model.data(),
static_cast<AX_U32>(model.size()), &extra),
"AX_ENGINE_CreateHandleV2 failed");
check_ax(AX_ENGINE_CreateContextV2(handle, &context),
"AX_ENGINE_CreateContextV2 failed");
check_ax(AX_ENGINE_GetIOInfo(handle, &info), "AX_ENGINE_GetIOInfo failed");
if (!info || info->nInputSize < 1 || info->nOutputSize < 1) {
throw std::runtime_error("model has no input or output tensors");
}
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) {
allocate(inputs[i], info->pInputs[i].nSize, "inflect_input");
}
for (AX_U32 i = 0; i < info->nOutputSize; ++i) {
allocate(outputs[i], info->pOutputs[i].nSize, "inflect_output");
}
}
~Impl() {
for (auto& item : inputs) {
if (item.phyAddr) AX_SYS_MemFree(item.phyAddr, item.pVirAddr);
}
for (auto& item : outputs) {
if (item.phyAddr) AX_SYS_MemFree(item.phyAddr, item.pVirAddr);
}
if (handle) AX_ENGINE_DestroyHandle(handle);
AX_ENGINE_Deinit();
AX_SYS_Deinit();
}
static void allocate(AX_ENGINE_IO_BUFFER_T& buffer, AX_U32 size, const char* token) {
std::memset(&buffer, 0, sizeof(buffer));
buffer.nSize = size;
check_ax(AX_SYS_MemAllocCached(&buffer.phyAddr, &buffer.pVirAddr, buffer.nSize,
128, reinterpret_cast<const AX_S8*>(token)),
"AX_SYS_MemAllocCached failed");
}
};
AxRunner::AxRunner(const std::string& model_path) : impl_(new Impl(model_path)) {}
AxRunner::~AxRunner() { delete impl_; }
std::vector<size_t> AxRunner::input_sizes() const {
std::vector<size_t> sizes;
for (AX_U32 i = 0; i < impl_->info->nInputSize; ++i) {
sizes.push_back(impl_->info->pInputs[i].nSize);
}
return sizes;
}
std::vector<size_t> AxRunner::output_sizes() const {
std::vector<size_t> sizes;
for (AX_U32 i = 0; i < impl_->info->nOutputSize; ++i) {
sizes.push_back(impl_->info->pOutputs[i].nSize);
}
return sizes;
}
std::vector<std::string> AxRunner::input_names() const {
std::vector<std::string> names;
for (AX_U32 i = 0; i < impl_->info->nInputSize; ++i) {
names.emplace_back(impl_->info->pInputs[i].pName
? reinterpret_cast<const char*>(impl_->info->pInputs[i].pName)
: "");
}
return names;
}
std::vector<std::string> AxRunner::output_names() const {
std::vector<std::string> names;
for (AX_U32 i = 0; i < impl_->info->nOutputSize; ++i) {
names.emplace_back(impl_->info->pOutputs[i].pName
? reinterpret_cast<const char*>(impl_->info->pOutputs[i].pName)
: "");
}
return names;
}
std::vector<std::vector<uint8_t>> AxRunner::run(
const std::vector<std::pair<const void*, size_t>>& feeds) {
if (feeds.size() != impl_->inputs.size()) {
throw std::runtime_error("feed count != model input count");
}
for (size_t i = 0; i < feeds.size(); ++i) {
if (feeds[i].second != impl_->inputs[i].nSize) {
throw std::runtime_error("feed byte size != model input buffer size");
}
std::memcpy(impl_->inputs[i].pVirAddr, feeds[i].first, feeds[i].second);
AX_SYS_MflushCache(impl_->inputs[i].phyAddr, impl_->inputs[i].pVirAddr,
impl_->inputs[i].nSize);
}
check_ax(AX_ENGINE_RunSyncV2(impl_->handle, impl_->context, &impl_->io),
"AX_ENGINE_RunSyncV2 failed");
std::vector<std::vector<uint8_t>> result(impl_->outputs.size());
for (size_t i = 0; i < impl_->outputs.size(); ++i) {
AX_SYS_MinvalidateCache(impl_->outputs[i].phyAddr, impl_->outputs[i].pVirAddr,
impl_->outputs[i].nSize);
result[i].resize(impl_->outputs[i].nSize);
std::memcpy(result[i].data(), impl_->outputs[i].pVirAddr,
impl_->outputs[i].nSize);
}
return result;
}
#else // !INFLECT_WITH_AX_ENGINE — host stub (configure/build only)
struct AxRunner::Impl {};
AxRunner::AxRunner(const std::string&) : impl_(nullptr) {
throw std::runtime_error(
"inflect_tts was built without the AX runtime; reconfigure with "
"-DAX_RUNTIME_ROOT=<ax bsp root> (see sdk/cpp/README.md)");
}
AxRunner::~AxRunner() = default;
std::vector<size_t> AxRunner::input_sizes() const { return {}; }
std::vector<size_t> AxRunner::output_sizes() const { return {}; }
std::vector<std::string> AxRunner::input_names() const { return {}; }
std::vector<std::string> AxRunner::output_names() const { return {}; }
std::vector<std::vector<uint8_t>> AxRunner::run(
const std::vector<std::pair<const void*, size_t>>&) {
throw std::runtime_error("AX runtime not available in this build");
}
#endif // INFLECT_WITH_AX_ENGINE
|