使用Onnx转换Sklearn决策树后如何获取样本预测对应的叶节点索引
解决方案
1. 调整ONNX转换参数
skl2onnx原生支持树模型输出叶节点索引,只需在convert_sklearn的配置项中新增output_leaf_nodes=True即可,修改后的转换代码如下:
# Convert into ONNX format from skl2onnx import convert_sklearn from skl2onnx.common.data_types import FloatTensorType initial_type = [('float_input', FloatTensorType([None, X_train.shape[1]]))] # 新增output_leaf_nodes参数开启叶节点索引输出 onx = convert_sklearn(clf, initial_types=initial_type, options={type(clf): {'nocl': True, 'output_leaf_nodes': True}}) with open(file_loc, "wb") as f: f.write(onx.SerializeToString())
2. 调整推理代码获取叶节点索引
开启上述参数后,导出的ONNX模型会新增第三个输出项,对应每个样本的预测叶节点ID,直接读取该输出即可:
# Compute the prediction with ONNX Runtime import onnxruntime as rt import numpy as np sess = rt.InferenceSession(file_loc) input_name = sess.get_inputs()[0].name # 第三个输出(索引为2)对应叶节点索引 leaf_node_name = sess.get_outputs()[2].name # 如需同时获取概率和叶节点,可在run的第一个参数中传入多个输出名 leaf_nodes = sess.run([leaf_node_name], {input_name: X_test.values.astype(np.float32)})[0]
补充说明
- 输出的叶节点索引和Sklearn原生
clf.apply(X_test)返回的结果完全一致 - 如果是多输出决策树,叶节点索引会按输出顺序依次返回,对应每个输出的预测节点位置
内容的提问来源于stack exchange,提问作者Ilan12
相关产品推荐
相关产品推荐

