ONNX Runtime 1.16.3 C++如何获取可读模型输入名称?
ONNX Runtime 1.16+ 正确获取模型输入名称的方法
在ONNX Runtime 1.16.3版本中调用GetInputNameAllocated获取模型输入名称时会出现乱码,而在1.12版本使用GetInputName接口则能正常得到真实名称。这是因为新版本API的字符串编码和内存管理逻辑发生了变化,以下是两种可行的解决方法:
方法一:直接使用宽字符输出
ONNX Runtime 1.13+的GetInputNameAllocated返回的是宽字符格式的字符串,直接用宽字符输出流std::wcout即可正常显示:
#include <iostream> #include <onnxruntime_cxx_api.h> int main() { const char *model_path = "/home/roroco/Downloads/mix/test_ai/resnet18-v1-7.onnx"; Ort::Env env(ORT_LOGGING_LEVEL_WARNING, "ONNXExample"); Ort::SessionOptions session_options; Ort::Session session(env, model_path, session_options); size_t input_count = session.GetInputCount(); Ort::AllocatorWithDefaultOptions allocator; for (size_t i = 0; i < input_count; ++i) { auto input_name = session.GetInputNameAllocated(i, allocator); std::wcout << L"Input name " << i << L": " << input_name.get() << std::endl; } return 0; }
方法二:转换为UTF-8字符串输出
如果需要输出普通char类型的UTF-8字符串,可以使用ONNX Runtime提供的Ort::ConvertWStringToString函数完成转换:
#include <iostream> #include <onnxruntime_cxx_api.h> int main() { const char *model_path = "/home/roroco/Downloads/mix/test_ai/resnet18-v1-7.onnx"; Ort::Env env(ORT_LOGGING_LEVEL_WARNING, "ONNXExample"); Ort::SessionOptions session_options; Ort::Session session(env, model_path, session_options); size_t input_count = session.GetInputCount(); Ort::AllocatorWithDefaultOptions allocator; for (size_t i = 0; i < input_count; ++i) { auto input_name_w = session.GetInputNameAllocated(i, allocator); std::string input_name = Ort::ConvertWStringToString(input_name_w.get()); std::cout << "Input name " << i << ": " << input_name << std::endl; } return 0; }
原因说明
ONNX Runtime从1.13版本开始,为了更好地支持多语言字符,将C++ API中的字符串接口统一改为宽字符(wchar_t)类型,旧的char*类型GetInputName接口被标记为废弃。直接将宽字符指针强制转为char*输出会因编码不匹配导致乱码,必须使用宽字符输出或完成编码转换才能正常显示。
内容的提问来源于stack exchange,提问作者chikadance
相关产品推荐
相关产品推荐

