Download cpp/src/ax_runner.cpp from AXERA-TECH/inflect_micro_v2: direct link, hf CLI and curl.
- Browser
- Download file 6.79 kB
-
https://huggingface.co/AXERA-TECH/inflect_micro_v2/resolve/main/cpp/src/ax_runner.cpp
- Command line
-
hf download hf://AXERA-TECH/inflect_micro_v2/cpp/src/ax_runner.cpp
-
curl -L -o ax_runner.cpp https://huggingface.co/AXERA-TECH/inflect_micro_v2/resolve/main/cpp/src/ax_runner.cpp
6.79 kB
| 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; | |
| } | |
| 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"); | |
| } | |