决策树是一种基于特征对实例进行分类的树形结构,它通过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折交叉验证评估模型性能
提供模型准确率和特征重要性分析
生成决策树可视化文件便于模型解释
该实现展示了决策树算法在分类问题中的应用,通过可视化功能增强了模型的可解释性,符合决策树算法可读性强的特点。

Logo

Agent 垂直技术社区,欢迎活跃、内容共建。

更多推荐