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)
Logo

码道开发者社区,聚焦华为云码道 CodeArts 代码智能体,沉淀 Agent、Skill、鸿蒙开发实战内容,供开发者查阅资料、交流技术、分享工程实践

更多推荐