使用dart booster的XGBClassifier转ONNX报错,求解决方法
解决XGBoost Dart Booster转ONNX时的KeyError问题
报错的核心原因是onnxmltools的XGBoost转换模块仅支持gbtree和gblinear两种booster类型,当使用dart时,模型配置结构里没有它预期的gbtree_model_param字段,因此抛出KeyError。
改用skl2onnx库即可解决这个问题,它对XGBoost的dart booster支持更完善,具体步骤如下:
安装必要依赖:
pip install skl2onnx onnxruntime xgboost scikit-learn修改代码实现转换:
from skl2onnx import convert_estimator from skl2onnx.common.data_types import FloatTensorType from sklearn.datasets import load_digits from sklearn.model_selection import train_test_split import xgboost as xgb import onnxruntime as rt import numpy as np # 加载数据并训练Dart模型 digits = load_digits() X, y = digits.data, digits.target X_train, X_test, y_train, y_test = train_test_split(X, y) booster = xgb.XGBClassifier(booster='dart', n_estimators=50, max_depth=10, learning_rate=0.1, random_state=42) booster.fit(X_train, y_train) # 定义输入张量类型,[None, 64]支持批量输入 initial_type = [('float_input', FloatTensorType([None, 64]))] # 执行模型转换 booster_onnx = convert_estimator(booster, initial_types=initial_type) # 保存ONNX模型到文件 with open("xgb_dart.onnx", "wb") as f: f.write(booster_onnx.SerializeToString()) # 验证转换后模型的正确性 sess = rt.InferenceSession("xgb_dart.onnx") input_name = sess.get_inputs()[0].name output_name = sess.get_outputs()[0].name pred_onnx = sess.run([output_name], {input_name: X_test.astype(np.float32)})[0] pred_xgb = booster.predict(X_test) # 对比原生模型与ONNX模型的预测结果 print(f"模型转换结果一致:{np.allclose(pred_onnx, pred_xgb)}")
注意事项
- 确保
skl2onnx与xgboost、scikit-learn版本兼容,建议使用各库的最新稳定版 - 若需指定ONNX算子集版本,可在
convert_estimator中添加参数target_opset=17(版本号可按需调整)
内容的提问来源于stack exchange,提问作者JoGe
相关产品推荐
相关产品推荐

