MATLAB实现决策树分类器

实现决策树分类器的关键MATLAB函数和概念

模块 核心函数/概念 简要说明
模型训练 fitctree 根据训练数据创建决策树分类器
模型预测 predict 使用训练好的决策树对新的数据进行分类预测
性能评估 resubLoss , crossval , kfoldLoss 计算再代入误差、进行交叉验证、计算交叉验证误差
模型可视化 view 以图形化方式或文本形式查看决策树的结构和决策规则
过拟合处理 cvloss 评估剪枝后的损失,帮助找到泛化能力最佳的树

从基础代码开始

一个完整的决策树仿真通常包括数据准备、模型训练、性能评估和结果可视化步骤。下面以经典的鸢尾花数据集为例,展示基础的实现流程:

% 加载鸢尾花数据集
load fisheriris;

% 为了演示,这里仅使用萼片长度和萼片宽度两个特征
X = meas(:, 1:2); % 特征矩阵
Y = species;      % 类别标签

% 创建决策树分类器
treeModel = fitctree(X, Y, 'PredictorNames', {'Sepal_Length', 'Sepal_Width'});

% 使用训练好的模型对训练数据进行预测(用于评估再代入精度)
Y_pred_train = predict(treeModel, X);

% 计算训练集上的准确率
trainAccuracy = sum(strcmp(Y_pred_train, Y)) / length(Y);
fprintf('训练集准确率: %.2f%%\n', trainAccuracy * 100);

% 可视化决策树的图形化结构
view(treeModel, 'Mode', 'graph'); % 会打开一个图形窗口显示树结构

% 可视化决策边界
figure;
gscatter(X(:,1), X(:,2), Y, 'rgb', 'osd');
xlabel('Sepal Length (cm)');
ylabel('Sepal Width (cm)');
title('决策树分类结果');
hold on;

% 创建网格点覆盖整个特征空间,用于绘制决策区域
xRange = linspace(min(X(:,1)), max(X(:,1)), 200);
yRange = linspace(min(X(:,2)), max(X(:,2)), 200);
[xx, yy] = meshgrid(xRange, yRange);
gridPoints = [xx(:), yy(:)];

% 预测网格点的类别
predGrid = predict(treeModel, gridPoints);

% 绘制决策区域
gscatter(gridPoints(:,1), gridPoints(:,2), predGrid, 'rgb', '.', 1);
legend('Location', 'best');
hold off;

评估模型性能

训练好模型后,我们需要科学地评估其性能,并解决可能存在的过拟合问题。

再代入误差与交叉验证误差

% 计算再代入误差 (在训练集上的误差)
resubError = resubLoss(treeModel);
fprintf('再代入误差 (训练集误分率): %.2f%%\n', resubError * 100);

% 执行10折交叉验证,更可靠地估计模型在新数据上的表现
cvTree = crossval(treeModel, 'KFold', 10);

% 计算交叉验证误差
cvError = kfoldLoss(cvTree);
fprintf('10折交叉验证误差: %.2f%%\n', cvError * 100);

理解过拟合与进行剪枝

决策树容易在训练集上表现很好(再代入误差低)但对新数据泛化能力差(交叉验证误差高),这就是过拟合。剪枝是解决过拟合的常用方法。

% 评估不同剪枝级别下的误差,帮助选择最优子树
resubCost = resubLoss(treeModel, 'Subtrees', 'all'); % 所有子树级别的再代入误差
[cvCost, stdError, nLeaves, bestLevel] = cvloss(treeModel, 'Subtrees', 'all'); % 所有子树级别的交叉验证误差

% 绘制误差随树复杂度(叶子节点数)的变化
figure;
plot(nLeaves, cvCost, 'b.-', 'LineWidth', 2, 'MarkerSize', 15); 
hold on;
plot(nLeaves, resubCost, 'r.--', 'LineWidth', 2, 'MarkerSize', 15);
xlabel('Number of Leaf Nodes (Tree Complexity)');
ylabel('Misclassification Error');
title('过拟合与剪枝分析');
legend('Cross-Validation Error', 'Resubstitution Error', 'Location', 'best');
grid on;

% 自动选择"最佳"剪枝级别(默认选择误差在一个标准差范围内最简单的树)
optimalTree = prune(treeModel, 'Level', bestLevel);
fprintf('原始树叶节点数: %d\n', size(treeModel.PruneList, 1));
fprintf('最优剪枝后叶节点数: %d\n', size(optimalTree.PruneList, 1));

% 你也可以手动指定剪枝级别,例如修剪到第5级
% manuallyPrunedTree = prune(treeModel, 'Level', 5);

决策树的核心:分裂准则

决策树如何选择最佳分裂点和特征至关重要,这通常由分裂准则决定。

% 在创建决策树时,你可以指定不同的分裂准则,例如基于基尼不纯度的分裂
% treeModel_Gini = fitctree(X, Y, 'SplitCriterion', 'gini');

% 或者,你也可以尝试基于信息增益(熵)的分裂
treeModel_Entropy = fitctree(X, Y, 'SplitCriterion', 'deviance'); % 在MATLAB中,'deviance'通常用于分类树的信息增益计算

% 比较不同分裂准则对模型性能的影响
% cvError_Gini = kfoldLoss(crossval(treeModel_Gini));
cvError_Entropy = kfoldLoss(crossval(treeModel_Entropy));
% fprintf('基尼系数准则 - 交叉验证误差: %.2f%%\n', cvError_Gini * 100);
fprintf('信息增益准则 - 交叉验证误差: %.2f%%\n', cvError_Entropy * 100);

自定义实现与高级应用

除了使用内置函数,你还可以尝试自定义实现或探索高级应用。

自定义决策树实现

搜索结果中提到了一些教学用的自定义决策树MATLAB代码文件(如 build_tree.m, ent.m, cond_ent.m 等)。这些代码通常更注重教学清晰度,能帮助你理解决策树构建、熵和条件熵计算等底层原理。如果你对算法细节感兴趣,可以寻找这类教学代码进行学习。

处理类别不平衡

如果数据集中不同类别的样本数量差距很大,可以引入先验概率。

% 计算每个类的先验概率(可以根据数据分布调整)
classPrior = 'empirical'; % 使用数据中的经验分布
% 或者手动指定先验概率,例如:classPrior = [0.3, 0.4, 0.3];

% 在训练模型时指定先验概率
treeModel_Prior = fitctree(X, Y, 'Prior', classPrior);

特征重要性分析

决策树也可以用于评估特征的重要性。

% 特征重要性可以通过比较分裂前后不纯度的减少量来间接评估
% 更直接的方法是使用随机森林等进行评估,但决策树本身也能提供一定洞察

% 通过查看模型结构,或者使用相关函数(如predictorImportance)来分析
% 注意:predictorImportance 适用于集成树如随机森林,对于单棵树意义有限

参考代码 matlab实现的决策树仿真的代码 www.3dddown.com/cna/82338.html

实践提示与常见问题

  • 数据预处理:确保类别标签是分类变量(categorical array)或字符串细胞数组,数值型特征最好进行归一化。
  • 模型复杂性控制:除了剪枝,还可以在 fitctree 中直接设置 MaxDepth(最大深度)或 MinLeafSize(最小叶子节点样本数)来控制树生长,防止过拟合。
  • 结果解读view 函数生成的图形化树结构可能很复杂,对于大型树,可以重点理解决策路径。
posted @ 2025-12-31 11:50  u95900090  阅读(64)  评论(0)    收藏  举报