#include "ax_runner.h" #include #ifdef INFLECT_WITH_AX_ENGINE #include #include #include #include #include namespace { std::vector 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(std::istreambuf_iterator(file), std::istreambuf_iterator()); } 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 inputs; std::vector outputs; std::vector 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(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(model_name); check_ax(AX_ENGINE_CreateHandleV2(&handle, model.data(), static_cast(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(token)), "AX_SYS_MemAllocCached failed"); } }; AxRunner::AxRunner(const std::string& model_path) : impl_(new Impl(model_path)) {} AxRunner::~AxRunner() { delete impl_; } std::vector AxRunner::input_sizes() const { std::vector sizes; for (AX_U32 i = 0; i < impl_->info->nInputSize; ++i) { sizes.push_back(impl_->info->pInputs[i].nSize); } return sizes; } std::vector AxRunner::output_sizes() const { std::vector sizes; for (AX_U32 i = 0; i < impl_->info->nOutputSize; ++i) { sizes.push_back(impl_->info->pOutputs[i].nSize); } return sizes; } std::vector AxRunner::input_names() const { std::vector names; for (AX_U32 i = 0; i < impl_->info->nInputSize; ++i) { names.emplace_back(impl_->info->pInputs[i].pName ? reinterpret_cast(impl_->info->pInputs[i].pName) : ""); } return names; } std::vector AxRunner::output_names() const { std::vector names; for (AX_U32 i = 0; i < impl_->info->nOutputSize; ++i) { names.emplace_back(impl_->info->pOutputs[i].pName ? reinterpret_cast(impl_->info->pOutputs[i].pName) : ""); } return names; } std::vector> AxRunner::run( const std::vector>& 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> 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= (see sdk/cpp/README.md)"); } AxRunner::~AxRunner() = default; std::vector AxRunner::input_sizes() const { return {}; } std::vector AxRunner::output_sizes() const { return {}; } std::vector AxRunner::input_names() const { return {}; } std::vector AxRunner::output_names() const { return {}; } std::vector> AxRunner::run( const std::vector>&) { throw std::runtime_error("AX runtime not available in this build"); } #endif // INFLECT_WITH_AX_ENGINE