如何在Python中生成类R风格的Scikit-Learn决策树可视化?
实现类R风格的Scikit-Learn决策树可视化(泰坦尼克号数据集示例)
方案一:用dtreeviz快速实现(推荐)
dtreeviz是专门为Scikit-Learn决策树设计的可视化库,默认风格接近R的rpart包,支持高度自定义,刚好能满足你的需求:替换分支的True/False为特征实际取值,移除节点里的阈值描述。
- 先安装库:
pip install dtreeviz
- 泰坦尼克号数据集示例代码:
假设你已经训练好决策树分类器clf,特征名称列表为feature_names(比如['Pclass', 'Sex', 'Age', 'SibSp', 'Parch', 'Fare', 'Embarked']),目标类别为target_names=['Died', 'Survived']:
from dtreeviz.trees import dtreeviz # 自定义分支标签,替换默认的True/False为特征实际含义 def custom_branch_label(edge): feat_name = edge.feature threshold = edge.threshold if feat_name == 'Sex': # 假设Sex特征0=男性,1=女性 return "Male" if edge.left else "Female" elif feat_name == 'Pclass': return f"Class ≤{int(threshold)}" if edge.left else f"Class >{int(threshold)}" # 其他特征可按需扩展 return f"{feat_name} ≤{threshold:.1f}" if edge.left else f"{feat_name} >{threshold:.1f}" viz = dtreeviz( clf, X_train, y_train, feature_names=feature_names, target_name="Survived", class_names=target_names, # 隐藏节点内的阈值条件描述 show_node_labels=False, # 应用自定义分支标签 edge_notation_func=custom_branch_label ) viz.view() # 打开可视化窗口 viz.save("titanic_tree.svg") # 保存为矢量图
这个方案能直接生成类R风格的清晰决策树,分支显示特征的实际含义,节点仅展示类别分布和样本数,完全匹配你的需求。
方案二:手动用Matplotlib自定义绘制(高度定制场景)
如果不想依赖第三方库,可以结合Scikit-Learn的树结构API,用Matplotlib手动绘制,完全控制每个元素的显示:
- 核心思路:
- 通过
clf.tree_属性提取决策树的节点、分支信息 - 自定义节点绘制逻辑,只保留类别、样本数等关键信息
- 替换分支的True/False标签为特征的实际取值
- 示例代码片段:
import matplotlib.pyplot as plt from matplotlib.patches import Rectangle, Arrow def plot_custom_tree(clf, feature_names, target_names, ax=None): if ax is None: ax = plt.gca() ax.clear() ax.set_axis_off() # 递归绘制节点与分支 def draw_node(node_id, x, y, width): tree = clf.tree_ # 提取节点核心信息:样本数、预测类别 n_samples = tree.n_node_samples[node_id] pred_class = target_names[tree.value[node_id].argmax()] # 绘制节点矩形框 ax.add_patch(Rectangle((x - width/2, y - 0.1), width, 0.2, fill=False, edgecolor='black')) # 节点文本:仅显示类别和样本数 ax.text(x, y, f"{pred_class}\nn={n_samples}", ha='center', va='center') # 叶子节点停止递归 if tree.children_left[node_id] == tree.children_right[node_id]: return # 获取当前节点的特征与阈值 feat_idx = tree.feature[node_id] feat_name = feature_names[feat_idx] threshold = tree.threshold[node_id] # 自定义分支标签 if feat_name == 'Sex': left_label = "Male" right_label = "Female" elif feat_name == 'Pclass': left_label = f"Class ≤{int(threshold)}" right_label = f"Class >{int(threshold)}" else: left_label = f"{feat_name} ≤{threshold:.1f}" right_label = f"{feat_name} >{threshold:.1f}" # 绘制左分支 left_x = x - width/2 left_y = y - 0.5 ax.arrow(x, y - 0.1, left_x - x, left_y - y + 0.1, head_width=0.02, color='black') ax.text((x + left_x)/2, (y + left_y)/2, left_label, ha='center', va='center') draw_node(tree.children_left[node_id], left_x, left_y, width/2) # 绘制右分支 right_x = x + width/2 right_y = y - 0.5 ax.arrow(x, y - 0.1, right_x - x, right_y - y + 0.1, head_width=0.02, color='black') ax.text((x + right_x)/2, (y + right_y)/2, right_label, ha='center', va='center') draw_node(tree.children_right[node_id], right_x, right_y, width/2) # 从根节点开始绘制 draw_node(0, 0.5, 0.9, 0.8) plt.show() # 调用绘制函数 plot_custom_tree(clf, feature_names, target_names)
这个方法完全手动控制每个元素,你可以根据需求调整节点样式、分支标签、文本内容,灵活性更高。
内容的提问来源于stack exchange,提问作者Fish1996
相关产品推荐
相关产品推荐

