基于ONNX与MIGraphX的AMD GPU推理:如何避免重复编译?
解决方案:AMD GPU上ONNX Runtime加载预编译.mxr模型避免重复编译
核心问题分析
你遇到的两个错误根源如下:
- 序列化冲突:同时启用ONNX Runtime通用模型缓存(
SetOptimizedModelFilePath)和MIGraphX执行提供者的编译节点,两者机制不兼容,无法序列化包含编译节点的模型。 - 无效路径报错:
migraphx_save_model_path设为空字符串,即使migraphx_save_compiled_model=0,部分MIGraphX版本仍会尝试生成临时编译文件,导致路径无效。
步骤1:移除冲突的通用缓存设置
完全删除sessionOptions.SetOptimizedModelFilePath("/home/erikbylow/C++Libs/optimized_model.mxr")这一行。MIGraphX的预编译模型机制与ONNX Runtime通用优化缓存不能同时使用。
步骤2:修正MIGraphX配置选项
调整OrtMIGraphXProviderOptions参数,确保预编译模型正确加载:
OrtMIGraphXProviderOptions migraphxOptions{}; migraphxOptions.device_id = 0; // 与预编译.mxr的精度匹配(编译默认FP32,故设为0) migraphxOptions.migraphx_fp16_enable = 0; // 与预编译.mxr的INT8设置匹配 migraphxOptions.migraphx_int8_enable = 0; // 启用加载预编译模型 migraphxOptions.migraphx_load_compiled_model = 1; // 禁用保存编译模型,避免临时文件写入 migraphxOptions.migraphx_save_compiled_model = 0; migraphxOptions.migraphx_exhaustive_tune = false; // 确保路径为绝对路径,且当前用户有读取权限 migraphxOptions.migraphx_load_model_path = "/home/erikbylow/C++Libs/optimized_model.mxr"; // 改用nullptr而非空字符串,避免路径解析错误 migraphxOptions.migraphx_save_model_path = nullptr;
步骤3:验证预编译模型兼容性
- 版本一致:编译.mxr的
migraphx-driver版本必须与运行时ONNX Runtime集成的MIGraphX版本完全匹配,跨版本的.mxr文件无法兼容。 - 硬件匹配:编译模型的GPU架构需与运行时GPU一致(如同为AMD RDNA 2/RDNA 3),跨架构预编译模型无法加载。
- 精度匹配:编译时未加
--fp16则默认FP32,代码中migraphx_fp16_enable必须设为0;若编译用了--fp16,则需设为1。
步骤4:测试预编译模型有效性
先通过migraphx-driver直接运行.mxr验证模型本身无问题:
migraphx-driver run optimized_model.mxr --input <输入张量路径>
若此命令正常执行,说明模型有效;若失败,需重新编译原始ONNX模型。
修正后的完整Session创建代码
Ort::Session createSession(Ort::Env &env, const char *modelFilepath) { Ort::SessionOptions sessionOptions; sessionOptions.SetIntraOpNumThreads(8); sessionOptions.SetGraphOptimizationLevel(GraphOptimizationLevel::ORT_ENABLE_ALL); sessionOptions.SetLogSeverityLevel(0); // verbose #ifdef __ANDROID__ uint32_t nnapi_flags = 0; Ort::ThrowOnError(OrtSessionOptionsAppendExecutionProvider_Nnapi(sessionOptions, nnapi_flags)); #endif OrtMIGraphXProviderOptions migraphxOptions{}; migraphxOptions.device_id = 0; migraphxOptions.migraphx_fp16_enable = 0; migraphxOptions.migraphx_int8_enable = 0; migraphxOptions.migraphx_load_compiled_model = 1; migraphxOptions.migraphx_save_compiled_model = 0; migraphxOptions.migraphx_exhaustive_tune = false; migraphxOptions.migraphx_load_model_path = "/home/erikbylow/C++Libs/optimized_model.mxr"; migraphxOptions.migraphx_save_model_path = nullptr; sessionOptions.AppendExecutionProvider_MIGraphX(migraphxOptions); Ort::Session session(env, modelFilepath, sessionOptions); return session; }
内容的提问来源于stack exchange,提问作者MrB
相关产品推荐
相关产品推荐

