共计 1619 个字符,预计需要花费 5 分钟才能阅读完成。
为什么需要可视化决策树?
刚学机器学习时,我发现决策树算法虽然容易调用(clf.fit(X,y)一行代码就能跑起来),但模型到底怎么选特征、如何做判断完全是个黑箱。直到看到导师演示决策树可视化,才恍然大悟——原来每个节点的 gini 值、samples数量、value分布都能直观呈现,这对理解模型行为至关重要。

可视化方案对比
- 文本表示(
sklearn.tree.export_text) - 优点:无需额外依赖库,适合快速查看
-
缺点:复杂树结构时难以阅读,缺乏颜色标记
-
Graphviz(本文推荐方案)
- 优点:专业级图形渲染,支持节点颜色填充
-
缺点:需要单独安装 Graphviz 软件(后面会教避坑)
-
dtreeviz
- 优点:交互式可视化,动态显示数据分布
- 缺点:安装复杂(需安装 Java 环境)
核心实现步骤
环境准备
首先安装必要库(建议使用 conda 虚拟环境):
conda install python-graphviz scikit-learn
完整代码示例(鸢尾花数据集)
# 导入所需库
from sklearn.datasets import load_iris
from sklearn.tree import DecisionTreeClassifier, export_graphviz
import graphviz
# 加载数据(故意取子集制造过拟合)iris = load_iris()
X, y = iris.data[::2], iris.target[::2] # 每隔一行取一个样本
# 训练模型(关键参数 max_depth= 3 防止过深)clf = DecisionTreeClassifier(max_depth=3, criterion="gini")
clf.fit(X, y)
# 生成可视化图形
dot_data = export_graphviz(
clf,
out_file=None,
feature_names=iris.feature_names, # 特征名称对应
class_names=iris.target_names, # 分类标签映射
filled=True, # 颜色填充
rounded=True, # 圆角矩形
special_characters=True
)
graph = graphviz.Source(dot_data)
graph.render("iris_tree") # 保存为 PDF 文件
关键参数详解
feature_names:必须与训练数据X的列顺序完全一致class_names:分类任务的类别标签(回归问题不需要)filled=True:用颜色深浅表示节点纯度(越深纯度越高)rounded=True:审美党必备,节点变圆角
三大常见坑与解决方案
- Graphviz 报错:
- Windows 用户需单独安装Graphviz 软件
-
安装后添加 bin 目录到系统 PATH
-
中文乱码:
dot_data = dot_data.replace('fontname=Helvetica', 'fontname="Microsoft YaHei"') -
图形溢出:
- 控制
max_depth(建议 3 - 5 层) - 设置
proportion=True按比例显示样本量
延伸思考
尝试用 dtreeviz 库实现交互式可视化(鼠标悬停看数据分布):
from dtreeviz.trees import dtreeviz
viz = dtreeviz(clf, X, y,
target_name="species",
feature_names=iris.feature_names,
class_names=list(iris.target_names))
viz.view()
观察图形时注意比较:
– 使用 criterion="gini" 和"entropy"时节点的分裂顺序差异
– 过深的树如何在末端节点出现极端样本比例(这就是过拟合!)
结语
第一次成功看到自己训练的决策树展开时,那种 ” 原来如此 ” 的快乐至今难忘。建议读者动手调整 max_depth 参数,观察树结构变化——这是理解模型复杂度的最佳实验。
正文完
