Python 决策树分类算法分析与实现
·
决策树是一种基于特征对实例进行分类的树形结构,它通过if-then规则集合实现分类过程。该算法的主要优点是模型可读性强且分类速度快。
决策树核心组成
决策树由节点和有向边组成,其中包含两类特殊节点:根节点(第一个节点)和叶节点(最下层无子节点的节点)。节点代表决策边界(如年龄是否大于20),边代表条件关系。
算法原理与划分标准
决策树分类适用于目标值为分类型变量,特征值可以是分类型或连续型。其核心思想是通过特征变量划分来判断类别,采用贪婪算法实现局部最优分类。主要划分准则包括:
ID3算法:使用信息增益最大准则
C4.5算法:采用信息增益比最大准则
CART算法:基于基尼系数最小准则,能够实现更细致的划分
基尼系数和信息熵用于衡量节点纯度,计算公式
表示类别概率。
Python实现与可视化
scikit-learn库提供DecisionTreeClassifier类实现决策树分类,支持基尼系数和信息增益两种划分标准。通过export_graphviz函数可以生成决策树的可视化图形文件。
以下是完整的决策树分类实现代码:
from sklearn import datasets
from sklearn.model_selection import train_test_split, GridSearchCV
from sklearn.tree import DecisionTreeClassifier, export_graphviz
import matplotlib.pyplot as plt
import numpy as np
def load_iris_data():
"""加载鸢尾花数据集"""
iris = datasets.load_iris()
return iris.data, iris.target, iris.feature_names
def train_decision_tree(X_train, X_test, y_train, y_test):
"""训练决策树模型并进行参数调优"""
# 初始化决策树分类器
estimator = DecisionTreeClassifier(random_state=42)
# 设置参数网格
param_grid = {
"criterion": ['entropy', 'gini'],
"max_depth": [3, 5, 7, None]
}
# 使用网格搜索和10折交叉验证
grid_search = GridSearchCV(estimator, param_grid=param_grid, cv=10)
grid_search.fit(X_train, y_train)
return grid_search
def evaluate_model(model, X_test, y_test):
"""评估模型性能"""
y_pred = model.predict(X_test)
accuracy = model.score(X_test, y_test)
print("模型准确率: {:.2f}%".format(accuracy * 100))
print("最佳参数:", model.best_params_)
print("交叉验证最佳得分: {:.2f}%".format(model.best_score_ * 100))
return accuracy, y_pred
def visualize_decision_tree(model, feature_names, output_file='decision_tree.dot'):
"""生成决策树可视化文件"""
export_graphviz(model.best_estimator_,
out_file=output_file,
feature_names=feature_names,
class_names=['setosa', 'versicolor', 'virginica'],
filled=True,
rounded=True)
print(f"决策树可视化文件已生成: {output_file}")
print("可通过GraphViz工具转换为PDF或PNG格式查看")
def main():
"""主函数"""
print("=== 决策树分类算法实现 ===\n")
# 1. 加载数据
X, y, feature_names = load_iris_data()
print(f"数据集特征: {X.shape}")
print(f"特征名称: {feature_names}")
# 2. 划分训练集和测试集
X_train, X_test, y_train, y_test = train_test_split(
X, y, test_size=0.3, random_state=42
)
print(f"训练集大小: {X_train.shape}, 测试集大小: {X_test.shape}")
# 3. 训练模型
print("\n正在训练决策树模型...")
model = train_decision_tree(X_train, X_test, y_train, y_test)
# 4. 评估模型
print("\n模型评估结果:")
accuracy, y_pred = evaluate_model(model, X_test, y_test)
# 5. 生成可视化
print("\n生成决策树可视化...")
visualize_decision_tree(model, feature_names)
# 6. 特征重要性分析
print("\n特征重要性分析:")
importance = model.best_estimator_.feature_importances_
for i, (feature, imp) in enumerate(zip(feature_names, importance)):
print(f"{feature}: {imp:.3f}")
if __name__ == "__main__":
main()
代码实现功能包括:
使用鸢尾花数据集进行决策树分类训练
通过网格搜索优化决策树参数(划分准则和最大深度)
使用10折交叉验证评估模型性能
提供模型准确率和特征重要性分析
生成决策树可视化文件便于模型解释
该实现展示了决策树算法在分类问题中的应用,通过可视化功能增强了模型的可解释性,符合决策树算法可读性强的特点。
更多推荐


所有评论(0)