如何用代码实现sklearn DecisionTreeClassifier特定路径遍历与类别预测?
可行
方法一:直接用原模型的predict方法(最简单)
如果你的DecisionTreeClassifier实例dt还存在,直接用官方提供的predict方法就能快速得到结果,无需手动解析树结构:
# 把样本特征整理成二维数组(sklearn要求输入为二维格式) sample = [[1, 0, 0]] # 注意特征顺序要和训练模型时的顺序完全一致 predicted_class = dt.predict(sample)[0] print(predicted_class)
方法二:解析tree_.__getstate__结果手动遍历
如果只有树的状态字典(比如原模型已销毁,仅保留了__getstate__的输出),可以通过解析字典内的核心字段实现自动遍历决策树:tree_.__getstate__()返回的字典包含几个关键字段:
children_left: 每个节点的左子节点索引(-1表示当前是叶子节点)children_right: 每个节点的右子节点索引(-1表示当前是叶子节点)feature: 非叶子节点用于分裂的特征索引threshold: 非叶子节点的分裂阈值value: 每个节点的类别计数数组,叶子节点取计数最大值对应的索引即为类别
以下是实现代码:
def predict_from_tree_state(tree_state, sample): current_node = 0 # 从根节点开始遍历 while True: # 判断当前节点是否为叶子节点 if tree_state['children_left'][current_node] == -1 and tree_state['children_right'][current_node] == -1: # 取计数最多的类别 return tree_state['value'][current_node].argmax() # 获取当前节点的分裂规则 split_feature_idx = tree_state['feature'][current_node] split_threshold = tree_state['threshold'][current_node] # 根据样本特征值选择子节点 if sample[split_feature_idx] <= split_threshold: current_node = tree_state['children_left'][current_node] else: current_node = tree_state['children_right'][current_node] # 假设树状态已保存到tree_state变量中 tree_state = dt.tree_.__getstate__() sample = [1, 0, 0] # 特征顺序必须和训练时一致 predicted_class = predict_from_tree_state(tree_state, sample) print(predicted_class)
注意:务必保证样本特征的顺序和模型训练时的特征顺序完全匹配,否则会得到错误的预测结果。
内容的提问来源于stack exchange,提问作者Winston Lee
相关产品推荐
相关产品推荐

