多输入类型下XGBoost转ONNX模型失败,求转换示例
解决onnxmltools转换多输入XGBoost模型的问题
我完全理解你遇到的困境——onnxmltools的XGBoost转换器目前只支持单一输入张量,这就是你触发"仅支持单一输入类型"错误的核心原因。不过我们可以通过几个变通方案来处理多输入场景,下面是具体的思路和代码示例:
问题根源
XGBoost在训练时本质上接收的是一个合并后的特征矩阵(不管原始特征是浮点型还是整数型,最终都会被整合为单一矩阵),而onnxmltools的XGBoost转换器目前只适配了这种单输入的设计,所以直接传入多个TensorType会触发输入数量校验的错误。
方案1:合并多输入为单一张量(最简单直接)
这种方法是把不同类型的输入特征提前合并成一个大张量,转换和推理都基于这个合并后的张量操作,完全贴合XGBoost和onnxmltools的现有逻辑。
import xgboost as xgb import numpy as np from onnxmltools.convert import convert_xgboost from onnxconverter_common.data_types import FloatTensorType # 1. 准备多类型训练数据 float_features = np.random.rand(100, 2).astype(np.float32) # 2个浮点特征 int_features = np.random.randint(0, 10, size=(100, 1)).astype(np.int64) # 1个整数特征 X_combined = np.hstack([float_features, int_features]) # 合并为3列的特征矩阵 y = np.random.rand(100).astype(np.float32) # 标签数据 # 2. 训练XGBoost回归模型 xgb_reg = xgb.XGBRegressor(objective='reg:squarederror') xgb_reg.fit(X_combined, y) # 3. 转换为ONNX模型:使用合并后的输入类型 initial_types = [('combined_input', FloatTensorType([None, 3]))] # None支持可变批量,3是总特征数 onnx_model = convert_xgboost(xgb_reg, initial_types=initial_types) # 4. 保存ONNX模型 with open("xgb_combined_input.onnx", "wb") as f: f.write(onnx_model.SerializeToString())
推理时,你只需要把新的多输入特征合并成同样维度的张量,再传入ONNX模型即可。
方案2:保留多输入接口(通过ONNX节点拼接)
如果你的场景必须对外暴露多输入的接口(比如要和其他多输入的ONNX模型串联),可以手动构建ONNX图,先通过Cast和Concat节点把多输入合并,再喂给XGBoost的ONNX子模型。
import onnx from onnx import helper, TensorProto # 1. 先加载方案1中转换好的单输入XGBoost模型 xgb_onnx_model = onnx.load("xgb_combined_input.onnx") # 2. 定义多输入的张量信息 input_float = helper.make_tensor_value_info('input', TensorProto.FLOAT, [None, 2]) input_int = helper.make_tensor_value_info('another_input', TensorProto.INT64, [None, 1]) # 3. 添加预处理节点:把整数特征转为浮点型,再和浮点特征合并 cast_node = helper.make_node( 'Cast', inputs=['another_input'], outputs=['int_casted'], to=TensorProto.FLOAT ) concat_node = helper.make_node( 'Concat', inputs=['input', 'int_casted'], outputs=['combined_input'], axis=1 # 按列合并特征 ) # 4. 拼接新的ONNX图:预处理节点 + XGBoost原节点 xgb_nodes = xgb_onnx_model.graph.node xgb_output = xgb_onnx_model.graph.output combined_graph = helper.make_graph( [cast_node, concat_node] + xgb_nodes, 'MultiInputXGBRegressor', [input_float, input_int], xgb_output ) # 5. 构建并验证完整模型 combined_model = helper.make_model(combined_graph, producer_name='onnxmltools') onnx.checker.check_model(combined_model) # 检查模型合法性 # 6. 保存多输入版本的ONNX模型 onnx.save(combined_model, "xgb_multi_input.onnx")
这个方案既保留了多输入的接口,又适配了XGBoost模型的单输入要求,完美解决你的场景需求。
目前onnxmltools官方对XGBoost的多输入支持还在完善中,这两个变通方案是当前最可行的解决路径。
内容的提问来源于stack exchange,提问作者Almo Daved
相关产品推荐
相关产品推荐

