升级onnxruntime至v1.14.0后CatBoost ONNX模型无法读取求助
ONNX Runtime升级v1.14.0后CatBoost模型失效问题排查与解决
问题背景
作为ONNX Runtime新手,使用旧C++代码基于PyTorch和CatBoost二分类模型做数据评估,代码在ONNX Runtime v1.6.0下运行正常,但升级到v1.14.0后CatBoost模型无法工作。
模型导出代码
CatBoost模型按原有方式导出为ONNX格式:
model.save_model( f'{filename}', format="onnx", export_parameters={ 'onnx_domain': 'ai.catboost', } )
C++端输出类型处理代码
初始化ONNX Runtime时,通过以下逻辑判断输出类型:
Ort::TypeInfo typeInfo = _session->GetOutputTypeInfo(i); if (typeInfo.GetONNXType() == ONNX_TYPE_SEQUENCE) { // <- 针对CatBoost模型的分支 const OrtSequenceTypeInfo *sequence{}; Ort::ThrowOnError(Ort::GetApi().CastTypeInfoToSequenceTypeInfo(typeInfo, &sequence)); OrtTypeInfo *sequenceTypeInfo{}; Ort::ThrowOnError(Ort::GetApi().GetSequenceElementType(sequence, &sequenceTypeInfo)); Ort::TypeInfo sequenceElementInfo = Ort::TypeInfo{sequenceTypeInfo}; if (sequenceElementInfo.GetONNXType() == ONNX_TYPE_TENSOR) { _outputNodeTypes.push_back(ONNX_TYPE_SEQUENCE); } else { _outputNodeTypes.push_back(ONNX_TYPE_MAP); } } else { _outputNodeTypes.push_back(ONNX_TYPE_TENSOR); // <- 针对PyTorch模型的分支,两个版本均正常 }
崩溃触发代码
访问输出值时在指定行抛出Ort::Exception():
if (_outputNodeTypes.front() == ONNX_TYPE_MAP) { Ort::Value &value = outputTensors.front(); Ort::Value map = value.GetValue(0, allocator); // <-- 触发异常的代码行 Ort::Value mapValue = map.GetValue(1, allocator); auto probabilities = gsl::span<float>(mapValue.GetTensorMutableData<float>(), 2); return probabilities[1]; }
关键现象
- 初始化时报错:
what(): Input is not of type sequence or map - 崩溃前打印输出值类型:v1.6.0显示为
ONNX_TYPE_SEQUENCE,v1.14.0显示为ONNX_TYPE_TENSOR,触发类型不匹配异常
原因分析
ONNX Runtime在v1.6.0到v1.14.0的迭代中,对自定义域(此处为ai.catboost)的ONNX模型输出类型解析逻辑发生了核心变更:
- 早期版本中,CatBoost导出的ONNX模型输出被识别为**序列(Sequence)**类型,旧代码的对应分支会被触发;
- 新版本中,ONNX Runtime对CatBoost导出模型的输出类型解析更贴合标准ONNX规范,直接识别为**张量(Tensor)**类型,导致旧代码中依赖
ONNX_TYPE_SEQUENCE的分支不再执行,后续却仍按Map类型处理张量,最终引发类型不匹配异常。
解决方案
1. 适配输出类型判定逻辑
修改输出类型检测代码,优先判断张量类型,同时兼容旧版本的序列/Map类型:
Ort::TypeInfo typeInfo = _session->GetOutputTypeInfo(i); auto onnxType = typeInfo.GetONNXType(); if (onnxType == ONNX_TYPE_TENSOR) { _outputNodeTypes.push_back(ONNX_TYPE_TENSOR); } else if (onnxType == ONNX_TYPE_SEQUENCE) { const OrtSequenceTypeInfo *sequence{}; Ort::ThrowOnError(Ort::GetApi().CastTypeInfoToSequenceTypeInfo(typeInfo, &sequence)); OrtTypeInfo *sequenceTypeInfo{}; Ort::ThrowOnError(Ort::GetApi().GetSequenceElementType(sequence, &sequenceTypeInfo)); Ort::TypeInfo sequenceElementInfo = Ort::TypeInfo{sequenceTypeInfo}; if (sequenceElementInfo.GetONNXType() == ONNX_TYPE_TENSOR) { _outputNodeTypes.push_back(ONNX_TYPE_SEQUENCE); } else { _outputNodeTypes.push_back(ONNX_TYPE_MAP); } } else { _outputNodeTypes.push_back(onnxType); }
2. 修改输出值访问逻辑
新增张量类型的处理分支,直接从张量中读取概率值,同时保留旧版本的兼容逻辑:
if (_outputNodeTypes.front() == ONNX_TYPE_TENSOR) { Ort::Value &value = outputTensors.front(); auto probabilities = gsl::span<float>(value.GetTensorMutableData<float>(), 2); return probabilities[1]; } else if (_outputNodeTypes.front() == ONNX_TYPE_MAP) { // 保留原有Map类型处理逻辑,兼容旧版本 Ort::Value &value = outputTensors.front(); Ort::Value map = value.GetValue(0, allocator); Ort::Value mapValue = map.GetValue(1, allocator); auto probabilities = gsl::span<float>(mapValue.GetTensorMutableData<float>(), 2); return probabilities[1]; } else if (_outputNodeTypes.front() == ONNX_TYPE_SEQUENCE) { // 保留原有Sequence类型处理逻辑,兼容旧版本 // 可根据实际序列结构补充读取逻辑 }
3. 验证模型格式
使用ONNX Runtime工具确认模型输出类型:
onnxruntime-inspect --model path/to/your/catboost_model.onnx
若模型导出格式存在兼容性问题,可尝试升级CatBoost版本后重新导出,确保与新版本ONNX Runtime适配。
内容的提问来源于stack exchange,提问作者gasar8
相关产品推荐
相关产品推荐

