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

如何在Python中生成类R风格的Scikit-Learn决策树可视化?

实现类R风格的Scikit-Learn决策树可视化(泰坦尼克号数据集示例)

方案一:用dtreeviz快速实现(推荐)

dtreeviz是专门为Scikit-Learn决策树设计的可视化库,默认风格接近R的rpart包,支持高度自定义,刚好能满足你的需求:替换分支的True/False为特征实际取值,移除节点里的阈值描述。

  1. 先安装库:
pip install dtreeviz
  1. 泰坦尼克号数据集示例代码:
    假设你已经训练好决策树分类器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手动绘制,完全控制每个元素的显示:

  1. 核心思路:
  • 通过clf.tree_属性提取决策树的节点、分支信息
  • 自定义节点绘制逻辑,只保留类别、样本数等关键信息
  • 替换分支的True/False标签为特征的实际取值
  1. 示例代码片段:
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

相关产品推荐
方舟 Agent Plan

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

最近更新时间:2026.07.28 14:47:44