You need to enable JavaScript to run this app.
优惠活动
大模型
产品
解决方案
定价
更多

升级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模型输出类型解析逻辑发生了核心变更:

  1. 早期版本中,CatBoost导出的ONNX模型输出被识别为**序列(Sequence)**类型,旧代码的对应分支会被触发;
  2. 新版本中,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

相关产品推荐
方舟 Agent Plan

超全模态模型 × Harness 升级,最新支持 Deepseek-V4.1-Flash、GLM-5.3 系列、Doubao-Seedream-5.0-pro、Kimi-K3 (部分), 限时 9.9 元起

最近更新时间:2026.07.27 17:17:56