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

使用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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.06.30 11:41:16