学习总结:机器学习——决策树可视化
·
from xml.sax.handler import feature_namespaces
from sklearn.datasets import load_breast_cancer
from sklearn.tree import DecisionTreeClassifier
from sklearn.model_selection import train_test_split
import numpy as np
import graphviz #决策树可视化
from sklearn import tree
import matplotlib.pyplot as plt
#加载乳腺癌数据集
cancer = load_breast_cancer()
print("canser.keys():\n{}".format(cancer.keys()))
print("shape of cancer data:{}".format(cancer.data.shape))
#cancer.target_names:返回类别名称列表:['malignant'(恶性), 'benign'(良性)]
#np.bincount(cancer.target):统计数组中每个值的出现次数:
#字典推导式 {n: v for n,v in ...}生成格式化的字典:{'malignant': 357, 'benign': 212}
print("sample counts per class:\n{}".format({n: v for n,v in zip(cancer.target_names,np.bincount(cancer.target))}))
# stratify=cancer.target:保持类别分布一致
#test_size=0.2:明确测试集比例
X_train,X_test,y_train,y_test = train_test_split(
cancer.data,cancer.target,stratify=cancer.target,test_size=0.2,random_state=42
)
print("——————未使用树的深度限制————————")
clf = DecisionTreeClassifier(random_state=0).fit(X_train,y_train)
print("Accuracy on training set:{:.3f}".format(clf.score(X_train,y_train)))
print("Accuracy on test set:{:.3f}".format(clf.score(X_test,y_test)))
#更换使用到达一定深度停止树的展开,通过限制树的深度来减低过拟合的风险
print("——————使用树的深度限制————————")
clf = DecisionTreeClassifier(max_depth=4,random_state=0).fit(X_train,y_train)
print("Accuracy on training set:{:.3f}".format(clf.score(X_train,y_train)))
print("Accuracy on test set:{:.3f}".format(clf.score(X_test,y_test)))
#使用tree模块的export_graphviz来对决策树进行可视化(需要graphviz软件包的支持)
#函数export_graphviz()会生成一个.dot的文本文件格式
dot_data = tree.export_graphviz(clf,
out_file = None,#out_file=None 输出到内存而非文件 返回DOT字符串供后续处理
class_names = ["malignant","benign"],#class_names 定义类别标签 ["malignant","benign"]对应乳腺癌良恶性分类
feature_names = cancer.feature_names,#feature_names 指定特征名称 使用cancer.feature_names(乳腺癌数据集的30个特征名)
impurity = False,#impurity=False 隐藏节点不纯度 避免显示基尼系数/信息增益等数值
filled = True)#filled=True 节点颜色填充 根据多数类别自动着色(如恶性为红色,良性为蓝色)
#rounded=True(默认) 节点圆角化 使决策树呈现更友好的可视化效果
graph = graphviz.Source(dot_data) #graphviz.Source对象:将DOT字符串包装为可渲染的Graphviz源对象
graph.render("cancer")#graph.render()方法:调用Graphviz引擎生成图像文件(默认保存为PDF格式,文件名cancer.pdf)
#因为有些树的分支很多,很难观察,所以利用一些有用的属性来总结树的工作原理。
#其中最常见的是特征重要性,它为每个特征对树的决策的重要性进行排序。
#对于每个特征来说,他都是一个0-1的数字;0表示“根本没用到“,1表示"完全预测目标值"。特种重要性的求和始终为一
print("特征重要性:\n{}".format(clf.feature_importances_))
#将特征重要性用条形图来可视化
def plot_feature_importances_cancer(model):
#cancer.data 是乳腺癌数据集的特征矩阵(二维NumPy数组)
#shape 属性返回 (样本数, 特征数)。shape[1] 提取特征维度数量
n_feature = cancer.data.shape[1]
#plt.barh() 绘制水平条形图
#y轴位置由 range(n_feature) 生成(0到29的整数序列)
#条形长度由 model.feature_importances_ 决定
#align="center" 使条形在y轴刻度上居中对齐,避免偏移
plt.barh(range(n_feature),model.feature_importances_,align="center")
#np.arange(n_feature) 生成与y轴位置匹配的刻度序列(0到29)
#cancer.feature_names 提供特征名称(如 'mean radius'、'worst texture'),将数字索引映射为可读性强的文本标签
plt.yticks(np.arange(n_feature),cancer.feature_names)
# 添加轴标签
plt.xlabel("Feature importance")
plt.ylabel("Feature")
# 自动调整布局防止标签重叠
plt.tight_layout()
# 保存图像到当前目录(PNG格式)
plt.savefig('feature_importance_plot.png')
print("特征重要性图已保存为 feature_importance_plot.png")
# 显示图像(在Jupyter等环境中直接展示)
plt.show()
plot_feature_importances_cancer(clf)
更多推荐



所有评论(0)