inflect_micro_v2 / cpp /src /ax_runner.cpp
inoryQwQ's picture
三芯片合并:AX620E/AX637 升级 encoder+decoder 全 NPU,新增新一代 SDK;AX650 保持老 SDK
5eee449 verified
Raw
History Blame Contribute Delete
6.79 kB
#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