sklearn plot_tree绘制单节点决策树时不显示类别问题咨询
问题原因
这是sklearn的plot_tree函数默认行为导致的:当决策树仅存在单个根节点(未发生任何分裂)时,默认的节点文本输出不会包含类别信息。函数仅在存在分裂分支的叶节点上展示class_names对应的类别标签,单节点场景下只会显示样本数、不纯度(impurity)等基础信息,不会主动映射并展示类别。
解决办法
可以通过手动添加文本标注的方式,在单节点树的可视化结果中补充类别信息,具体修改代码如下:
plt.figure(figsize=(12, 12)) # 绘制决策树并保存返回的文本对象 tree_elements = plot_tree(estimator, feature_names=feature_names, label= 'all', class_names=[f'Class {k}' for k in range(2)], filled=True, rounded=True, impurity = True ) # 判断是否为单节点树 if estimator.get_depth() == 0: # 获取根节点的预测类别:从树的value中找出占比最高的类别 class_idx = estimator.tree_.value.argmax() target_class = [f'Class {k}' for k in range(2)][class_idx] # 获取根节点文本的位置,添加类别标注 node_text = tree_elements[0] plt.text(node_text.get_x(), node_text.get_y() - 0.06, f"Class: {target_class}", ha='center', fontsize=14, bbox=dict(facecolor='white', edgecolor='none')) plt.title(f"Decision Tree for Effect {effect_name}", fontsize = 40) file_name = f"DTs_action_{str(num_action)}/decision_tree_effect_{effect_name}.png" plt.savefig(file_name, format="png", dpi=300)
这段代码会先检查树的深度,确认是单节点后,从决策树的内部数据中提取出预测类别,再在节点文本的下方添加类别标注,确保单节点场景下也能看到对应的类别信息。
内容的提问来源于stack exchange,提问作者Addon
相关产品推荐
相关产品推荐

