拓冰建站拓冰建站
首页 / 资讯中心 / 正文

C++手写决策树:从信息熵、基尼不纯度到鸢尾花分类实验

简介面向机器学习初学者的鸢尾花分类决策树实验基于C语言实现经典分类任务适合正在学习数据挖掘、监督学习或想要动手实践决策树算法的读者。资源内提供完整的C源码涵盖数据集读取与预处理、树节点结构定义、信息增益与基尼不纯度等分裂标准计算、递归建树、剪枝停止条件判断以及新样本预测等核心环节代码结构清晰便于对照学习。压缩包内共有1个cpp文件仅3KB小巧精炼适合直接阅读和二次修改。目前已有3413人学习/浏览是入门决策树原理与C算法实现的实用参考。通过阅读源码可以直观理解“如果-那么”规则如何形成树形结构掌握从特征选择到模型预测的完整流程同时也能提升数据处理与递归编程能力。1. 鸢尾花分类实验决策树在 C 里的复现思路鸢尾花数据集是机器学习里最常被拿来验证分类算法的“标准样本”150 条记录、4 个数值特征、3 个类别规模小到可以按行读进内存也足够看出算法差异。常见的教程大多用 Python 调 sklearn一行 fit 就能出树但决策树到底怎么选特征、怎么定阈值、怎么停止生长在这些调包代码里是看不到的。这个实验用 C 手写决策树从读取 iris.data 开始到信息增益计算、递归建树、分类预测全部自己实现适合正在做课程设计、想把树结构看明白的人。跑通之后你可以把每一步分裂的中间结果打印出来验证教材里“信息增益最大化”到底对应哪一列花瓣数据。2. 决策树分类的第一步熵、信息增益与基尼不纯度2.1 为什么分裂标准决定树的质量决策树的工作方式可以概括成一句话在当前节点上把样本按某个特征的一个阈值切成两堆期望切完之后每一堆都比切之前更“纯”。这里的纯度用信息熵、基尼不纯度或分类误差率来量化。信息熵公式是对每个类别的概率取 log2 再加和值越大表示类别越混乱基尼不纯度计算的是随机抽取两个样本、类别不一致的概率分类误差率则直接看多数类占比。三者都能衡量纯度但性质不同信息熵对类别分布的变化更敏感基尼不纯度在类别较多时计算更省分类误差率在二分类场景下有时会出现多个分裂点得分相同的问题。分裂标准计算表达式特点C 实现成本信息熵 H-Σ p_i * log2(p_i)对类别变化敏感更符合信息论解释需要 log2 调用基尼不纯度1 - Σ p_i^2计算简单适合大数据量乘加运算分类误差率1 - max(p_i)直观但二分类时区分度低最廉价鸢尾花实验里三者的差异不会特别明显因为 Setosa、Versicolour、Virginica 三类样本各 50 条类别分布均衡。真正影响树质量的是连续特征上的阈值搜索而不是这三个公式本身。不过理解它们仍然是必要的因为后面计算信息增益时父节点熵减去条件熵这一步所有标准都要用到。2.2 计算熵、基尼和条件熵的 C 代码下面这个函数用来计算一个节点内样本的熵。输入是标签数组标签用 0、1、2 表示三个鸢尾花品种#include vector #include cmath #include unordered_map // labels 中存放当前节点所有样本的类别编码 double calcEntropy(const std::vectorint labels) { std::unordered_mapint, int count; for (int label : labels) { count[label]; } if (count.empty()) return 0.0; double entropy 0.0; int total labels.size(); for (auto kv : count) { double p static_castdouble(kv.second) / total; entropy - p * log2(p); } return entropy; }这段代码先遍历所有标签做频次统计再按信息熵公式累加。需要注意count.empty()的提前返回避免除零log2在p为 0 时会崩溃但因为只统计出现过类别所以不会出现 p0。如果要在更大的分类问题里复用建议把labels.size()的类型改成size_t并显式处理 total 为 0 的情况。基尼不纯度的实现几乎一样只是把 log2 换成平方double calcGini(const std::vectorint labels) { std::unordered_mapint, int count; for (int label : labels) count[label]; double gini 1.0; int total labels.size(); for (auto kv : count) { double p static_castdouble(kv.second) / total; gini - p * p; } return gini; }在训练过程中父节点的熵只需要计算一次然后对每个候选特征、每个候选阈值计算条件熵也就是子节点熵按样本数加权平均。条件熵公式是Σ (|S_v| / |S|) * H(S_v)用代码实现时要注意加权因子用的是样本数占比不是类别占比。下面这个entropyFromCounts是calcEntropy的计数版本直接传入每个类别的计数数组和总样本数避免反复组装 labelsdouble entropyFromCounts(const std::vectorint counts, int total) { double entropy 0.0; for (int c : counts) { if (c 0) continue; double p static_castdouble(c) / total; entropy - p * log2(p); } return entropy; }这个函数在后面 bestSplit 里会频繁调用。注意 counts 长度必须覆盖类别数鸢尾花固定为 3如果改成别的数据集建议用 unordered_map 代替 vector。2.3 连续特征的分割点搜索鸢尾花的四个特征都是连续值不能像“颜色红色”那样枚举取值。常见做法是对当前节点内的样本按特征值排序把相邻两个值的中点作为候选阈值逐一计算信息增益选最大的那个。为了避免无效计算可以跳过标签相同的相邻样本因为同标签样本之间取中点分割不会增加信息量。#include algorithm struct DataItem { std::vectordouble features; int label; }; double bestSplit(const std::vectorDataItem data, int featureIdx, double parentEntropy, double bestThreshold) { int n data.size(); std::vectorstd::pairdouble,int sorted(n); for (int i 0; i n; i) { sorted[i] {data[i].features[featureIdx], data[i].label}; } std::sort(sorted.begin(), sorted.end()); double bestGain 0.0; for (int i 0; i n - 1; i) { if (sorted[i].second sorted[i1].second) continue; double threshold (sorted[i].first sorted[i1].first) / 2.0; std::vectorint leftCount(3, 0), rightCount(3, 0); int leftTotal 0, rightTotal 0; for (int j 0; j i; j) { leftCount[sorted[j].second]; leftTotal; } for (int j i 1; j n; j) { rightCount[sorted[j].second]; rightTotal; } double leftEnt entropyFromCounts(leftCount, leftTotal); double rightEnt entropyFromCounts(rightCount, rightTotal); double condEntropy (double)leftTotal / n * leftEnt (double)rightTotal / n * rightEnt; double gain parentEntropy - condEntropy; if (gain bestGain) { bestGain gain; bestThreshold threshold; } } return bestGain; }bestSplit返回最大信息增益并通过引用参数bestThreshold带回最佳阈值。每次尝试阈值时都要重新统计左右子集的频次复杂度是 O(n^2 log n)对 150 条样本没有压力如果数据量上万可以先排序再从左到右增量更新 leftCount 和 rightCount把单次统计降到 O(1)。跳过同标签相邻样本是一个小优化但要注意它不会漏掉最优阈值因为真正信息增益大于 0 的分割必然发生在类别变化处。3. 从数据到树C 构建决策树的递归过程3.1 鸢尾花数据读取与预处理原始数据集通常叫 iris.data每行是“花萼长度,花萼宽度,花瓣长度,花瓣宽度,类别”类别是 Iris-setosa、Iris-versicolor、Iris-virginica。C 没有现成的 CSV 读取器最简单的方式是把逗号替换成空格再用 istringstream 按顺序读。#include fstream #include sstream #include algorithm std::vectorDataItem loadIris(const std::string path) { std::ifstream file(path); std::vectorDataItem data; std::string line; while (std::getline(file, line)) { if (line.empty()) continue; std::replace(line.begin(), line.end(), ,, ); std::istringstream iss(line); DataItem item; double value; for (int i 0; i 4; i) { iss value; item.features.push_back(value); } std::string species; iss species; if (species Iris-setosa) item.label 0; else if (species Iris-versicolor) item.label 1; else item.label 2; data.push_back(item); } return data; }替换逗号为空格后字符串流可以按空白符自动分隔。四个特征依次读入类别名最后读入。这里把类别字符串直接映射成整数后面计算熵和混淆矩阵时都会用到。数据预处理不需要做归一化因为决策树只关心特征值的相对顺序和大小关系阈值搜索基于排序不在不同特征之间做距离比较。3.2 树节点与决策树类定义决策树的节点需要记录分裂特征、分裂阈值、叶子类别、左右孩子指针。由于递归建树需要随时判断是否达到停止条件节点里再加一个 isLeaf 标记会更直观。struct TreeNode { int featureIdx; // 内部节点分裂特征的下标 double threshold; // 内部节点分裂阈值 int label; // 叶节点多数类别 TreeNode* left; // 小于等于阈值的样本走左子树 TreeNode* right; // 大于阈值的样本走右子树 bool isLeaf; }; class DecisionTree { public: DecisionTree() : root_(nullptr) {} ~DecisionTree() { destroy(root_); } void train(const std::vectorDataItem data, int maxDepth, int minSamples); int predict(const std::vectordouble features) const; private: TreeNode* build(const std::vectorDataItem data, int depth, int maxDepth, int minSamples); void destroy(TreeNode* node); TreeNode* root_; };这里把 train 和 predict 作为公开接口内部 build 递归建树。树的左右子树规则必须和预测逻辑一致左子树存特征值小于等于阈值的样本右子树存大于阈值的样本。这个方向定下来之后预测函数里也要用同一个比较符号否则结果会完全颠倒。3.3 递归建树与停止条件build 函数是整棵树的构造核心。它先统计当前数据集的类别分布判断是否满足停止条件不满足则遍历四个特征调用 bestSplit 找最优特征和阈值然后按阈值把数据切成左右两部分递归建子树。TreeNode* DecisionTree::build(const std::vectorDataItem data, int depth, int maxDepth, int minSamples) { TreeNode* node new TreeNode(); std::vectorint counts(3, 0); for (auto item : data) counts[item.label]; int majority 0; for (int i 1; i 3; i) { if (counts[i] counts[majority]) majority i; } if (depth maxDepth || (int)data.size() minSamples || counts[majority] (int)data.size()) { node-isLeaf true; node-label majority; return node; } double parentEntropy entropyFromCounts(counts, data.size()); double bestGain 0.0; int bestFeature -1; double bestThreshold 0.0; for (int f 0; f 4; f) { double threshold 0.0; double gain bestSplit(data, f, parentEntropy, threshold); if (gain bestGain) { bestGain gain; bestFeature f; bestThreshold threshold; } } if (bestFeature -1 || bestGain 1e-6) { node-isLeaf true; node-label majority; return node; } node-featureIdx bestFeature; node-threshold bestThreshold; node-isLeaf false; node-left build(filterData(data, bestFeature, bestThreshold, true), depth 1, maxDepth, minSamples); node-right build(filterData(data, bestFeature, bestThreshold, false), depth 1, maxDepth, minSamples); return node; }三个停止条件的优先级要注意depth maxDepth放在最前避免继续深挖类别纯度达到 100% 时即使深度不够也要停minSamples控制叶子节点的最小样本量样本太少时继续分割统计意义不大。1e-6这个信息增益下界是一个经验值。如果设成 0那么信息增益刚好等于 0 的分割也会被选中树会一直长到深度上限设太大会导致树太浅。对 150 条鸢尾花数据1e-6 是稳妥的起点。filterData是辅助函数作用是把样本按阈值切到左子树或右子树std::vectorDataItem filterData(const std::vectorDataItem data, int featureIdx, double threshold, bool left) { std::vectorDataItem result; for (const auto item : data) { bool goLeft item.features[featureIdx] threshold; if (goLeft left) result.push_back(item); } return result; }filterData的返回值是一个新 vector会带来额外拷贝。对 150 条样本无所谓如果要在更大数据集上跑可以改成先收集索引数组递归时直接传递索引区间避免频繁分配内存。3.4 预剪枝和后剪枝的取舍决策树建好以后容易在训练集上过拟合。鸢尾花数据集只有 150 条样本如果完全不限制深度树的叶子节点可能精确到每一条训练样本测试集效果反而变差。常用的对策有两类预剪枝在建树过程中提前停止后剪枝先把树建完整再合并子树。剪枝方式优点缺点适用于鸢尾花实验预剪枝实现简单训练快可能过早停止错过好的分割推荐通过 maxDepth 和 minSamples 控制后剪枝保留更多信息泛化通常更好需要额外验证集代码复杂数据量小不适合单独划分验证集本实验采用预剪枝因为 150 条样本再分出验证集会进一步压缩训练数据。如果你想体会后剪枝可以留出 30 条样本做验证集在树构建后自底向上尝试删除子树比较验证集准确率不过在课程设计场景下预剪枝已经足够把决策树的核心流程讲清楚。4. 预测新样本与分类评估从单棵树的路径到泛化能力4.1 决策树预测的实现预测过程比训练简单得多从根节点出发根据样本的特征值与当前节点阈值的大小关系决定走左还是走右直到遇到叶子节点叶子节点的 label 就是预测类别。int DecisionTree::predict(const std::vectordouble features) const { TreeNode* node root_; while (!node-isLeaf) { if (features[node-featureIdx] node-threshold) { node node-left; } else { node node-right; } } return node-label; }这里使用和训练时的filterData保持一致。如果训练时把等于阈值的样本分到右子树这里就要同步改成否则等于阈值的那条样本会走错分支。这类边界问题在连续特征上尤其重要因为鸢尾花数据里特征值可能重复出现。4.2 分类评估指标准确率只是起点训练完成后需要一个量化标准判断树的分类效果。最直接的是准确率但准确率会把每个类别的错误混在一起所以还要看混淆矩阵。std::vectorstd::vectorint confusion(3, std::vectorint(3, 0)); int correct 0; for (const auto item : testData) { int pred tree.predict(item.features); confusion[item.label][pred]; if (pred item.label) correct; } double accuracy static_castdouble(correct) / testData.size();评估角度指标通过 confusion 计算方式整体正确率Accuracy对角线之和 / 总样本数单类覆盖率Recall第 i 行对角线 / 第 i 行总和单类可信度Precision第 i 列对角线 / 第 i 列总和综合指标Macro F1各类 F1 的算术平均混淆矩阵的行是真实类别列是预测类别。对角线越亮说明该类别正确率高。鸢尾花数据里 Versicolour 和 Virginica 在花瓣长度上存在重叠错误通常集中在它们之间Setosa 几乎不会错。只看准确率看不出这种倾向打印出混淆矩阵能快速定位是哪些特征区间产生了歧义。对这个小数据集还可以补一补宏平均 F1先分别计算每个类别的 precision 和 recall再求算术平均。宏平均 F1 对类别不平衡更敏感在鸢尾花这种三类均衡的数据上准确率和宏平均 F1 差异不大但实验报告里写上这个指标会显得完整。4.3 训练集/测试集划分与交叉验证把 150 条样本全部用于训练再报准确率没有意义因为决策树可以记住训练数据。常见的做法是按 7:3 分成训练集和测试集并保证每个类别在两个集合中的比例和原始数据一致。150 条样本足够做简单留出但结果依赖划分的随机性所以更稳的做法是 k 折交叉验证。k 折交叉验证的大致步骤是把所有样本随机打乱等分为 k 份每一轮取其中 1 份做测试、剩下 k-1 份做训练训练出一个树并记录测试准确率最终取 k 轮准确率的平均值。鸢尾花数据上我一般用 5 折因为类别已经均衡分布5 折每折 30 条测试样本误差信号足够稳定。如果样本量再小可以用 10 折但训练时间会变长意义不大。注意每次折数据划分的随机种子要固定否则复现实验时结果对不上。5. 决策树调参验证鸢尾花实验里最值得试的四个参数5.1 从 maxDepth 到 minGain 的参数组合手写决策树的参数就是预剪枝条件最大深度 maxDepth、叶子节点最小样本数 minSamples、信息增益最小阈值 minGain再加上分裂标准的选择。这四个参数不是孤立存在maxDepth 设得大时minSamples 要相应调大否则树会在深度限制内拼命把样本切碎。我建议按顺序搜索先固定 minSamples 为 5把 maxDepth 从 1 调到 6然后固定 maxDepth4把 minSamples 从 2 调到 20。这样能分别看到深度和叶子大小对准确率的影响。参数推荐范围观察信号maxDepth1~6深度为 3 时树高约 8 个节点通常已经能到 90% 以上minSamples2~20过小时训练集接近 100%测试集回落明显minGain0~0.1从小到大依次增长观察叶子数量变化splitCriterionentropy / gini两者结果基本一致gini 计算略快实际在鸢尾花数据集上maxDepth4、minSamples8 的组合通常能稳定在 95% 左右的测试准确率。如果 minSamples 降到 2训练集准确率会接近 100%测试集可能掉到 92% 以下典型的过拟合。minGain 调高到 0.05 以上后第二个分裂点可能被剪掉树就只剩一个根节点和两个叶子虽然准确率还能维持在 80% 左右但丢失了对 Versicolour 和 Virginica 的细分能力。5.2 验证树结构是否合理的小技巧调参完成后不要只看准确率。把根节点的分裂特征打印出来你会发现无论怎么调参数第一个分裂点几乎总是花瓣长度第 3 个特征阈值集中在 2.45 附近。这符合鸢尾花数据本身的分布Setosa 的花瓣长度明显小于另外两类。如果根节点分裂特征变成花萼宽度说明数据顺序没有洗牌或者阈值搜索代码有 bug。另一个快速验证方法是做一次决策树和随机森林的对比。用同一份训练集随机森林在 50 棵树上测试准确率会更高但单棵树的优势在于结构可读。你可以把每层选中的特征和阈值记录到 log 文件逐层检查哪些样本被错误分类、错误集中在哪个分支。这样即使不对着图也能定位到具体是特征重叠导致的天然误差还是预剪枝把该继续分割的节点提前截断了。最后一个小技巧在 build 函数里加一行统计depth和data.size()的输出训练完看一眼叶子节点平均样本数如果小于 3说明树切得太碎下一步优先调 minSamples 而不是 maxDepth。本文还有配套的精品资源点击获取
分享:

看完干货,该让你的企业上线了

免费需求沟通 · 48 小时内出具建站方案 · 河南本地可上门