决策树是机器学习领域最直观、最经典的分类算法之一。它通过递归地选择最优特征对数据集进行划分,构建出一棵形如” if-then “规则的树形结构,使模型具有极强的可解释性。本文将从信息论基础出发,分别实现 ID3 算法(基于信息增益)与 C4.5 算法(基于增益率),并用 Java 完成一个完整可运行的决策树分类系统。
一、算法核心思想
决策树的构建本质上是一个递归划分的过程:
- 特征选择:从当前数据集中选择”最优”的特征作为划分依据。
- 数据集划分:根据该特征的取值将数据集切分为若干子集。
- 递归构建:对每个子集重复上述过程,直到满足停止条件(如所有样本属于同一类别、无可用特征、或子集为空)。
ID3 使用信息增益(Information Gain)来衡量特征的划分能力;C4.5 则引入增益率(Gain Ratio)来克服信息增益对多值特征的偏好问题。
二、信息论基础
2.1 信息熵
信息熵度量数据集的”混乱程度”,定义为:
$$H(D) = -\sum_{k=1}^{K} p_k \log_2 p_k$$
其中 $p_k$ 是第 $k$ 类样本在数据集 $D$ 中的比例。熵越大,数据越混乱;熵为 0 表示所有样本属于同一类。
2.2 信息增益
信息增益表示使用特征 $A$ 对数据集进行划分后,熵减少了多少:
$$Gain(D, A) = H(D) – \sum_{v \in Values(A)} \frac{|D^v|}{|D|} H(D^v)$$
ID3 选择信息增益最大的特征进行划分。
2.3 增益率
信息增益倾向于选择取值较多的特征(如”用户ID”)。C4.5 引入分裂信息(Split Information)进行归一化:
$$SplitInfo(D, A) = -\sum_{v \in Values(A)} \frac{|D^v|}{|D|} \log_2 \frac{|D^v|}{|D|}$$
$$GainRatio(D, A) = \frac{Gain(D, A)}{SplitInfo(D, A)}$$
C4.5 选择增益率最大的特征进行划分。
三、Java 实现
3.1 数据结构定义
import java.util.*;
/**
* 决策树节点
*/
class TreeNode {
// 节点分类标签(叶子节点)
String label;
// 当前节点用于划分的特征索引
int featureIndex;
// 当前节点用于划分的特征名称
String featureName;
// 子节点:特征取值 -> 子节点
Map<String, TreeNode> children;
// 是否为叶子节点
boolean isLeaf;
public TreeNode() {
this.children = new HashMap<>();
this.isLeaf = false;
}
}
/**
* 数据集封装
*/
class Dataset {
// 特征名称列表
List<String> featureNames;
// 样本数据:每行是一个样本,每列是一个特征值(最后一列为类别标签)
List<String[]> data;
public Dataset(List<String> featureNames, List<String[]> data) {
this.featureNames = featureNames;
this.data = data;
}
}
3.2 核心工具类
/**
* 决策树工具类:提供熵计算、信息增益、增益率等方法
*/
class DecisionTreeUtil {
/**
* 计算数据集的信息熵
* @param dataset 数据集
* @return 信息熵
*/
public static double calculateEntropy(Dataset dataset) {
Map<String, Integer> labelCount = new HashMap<>();
int total = dataset.data.size();
for (String[] row : dataset.data) {
String label = row[row.length - 1];
labelCount.put(label, labelCount.getOrDefault(label, 0) + 1);
}
double entropy = 0.0;
for (int count : labelCount.values()) {
double p = (double) count / total;
entropy -= p * (Math.log(p) / Math.log(2));
}
return entropy;
}
/**
* 按特征索引划分数据集
* @param dataset 原数据集
* @param featureIndex 特征索引
* @return 划分后的子数据集映射
*/
public static Map<String, Dataset> splitDataset(Dataset dataset, int featureIndex) {
Map<String, List<String[]>> splitMap = new HashMap<>();
for (String[] row : dataset.data) {
String featureValue = row[featureIndex];
splitMap.computeIfAbsent(featureValue, k -> new ArrayList<>()).add(row);
}
Map<String, Dataset> result = new HashMap<>();
for (Map.Entry<String, List<String[]>> entry : splitMap.entrySet()) {
result.put(entry.getKey(), new Dataset(dataset.featureNames, entry.getValue()));
}
return result;
}
/**
* 计算信息增益(ID3)
* @param dataset 数据集
* @param featureIndex 特征索引
* @return 信息增益
*/
public static double calculateInformationGain(Dataset dataset, int featureIndex) {
double baseEntropy = calculateEntropy(dataset);
Map<String, Dataset> splits = splitDataset(dataset, featureIndex);
int total = dataset.data.size();
double weightedEntropy = 0.0;
for (Dataset subset : splits.values()) {
double weight = (double) subset.data.size() / total;
weightedEntropy += weight * calculateEntropy(subset);
}
return baseEntropy - weightedEntropy;
}
/**
* 计算分裂信息(C4.5)
* @param dataset 数据集
* @param featureIndex 特征索引
* @return 分裂信息
*/
public static double calculateSplitInfo(Dataset dataset, int featureIndex) {
Map<String, Dataset> splits = splitDataset(dataset, featureIndex);
int total = dataset.data.size();
double splitInfo = 0.0;
for (Dataset subset : splits.values()) {
double p = (double) subset.data.size() / total;
splitInfo -= p * (Math.log(p) / Math.log(2));
}
return splitInfo;
}
/**
* 计算增益率(C4.5)
* @param dataset 数据集
* @param featureIndex 特征索引
* @return 增益率
*/
public static double calculateGainRatio(Dataset dataset, int featureIndex) {
double gain = calculateInformationGain(dataset, featureIndex);
double splitInfo = calculateSplitInfo(dataset, featureIndex);
// 避免除以零
return splitInfo == 0 ? 0 : gain / splitInfo;
}
/**
* 获取数据集中出现次数最多的类别标签
* @param dataset 数据集
* @return 多数类标签
*/
public static String getMajorityLabel(Dataset dataset) {
Map<String, Integer> labelCount = new HashMap<>();
for (String[] row : dataset.data) {
String label = row[row.length - 1];
labelCount.put(label, labelCount.getOrDefault(label, 0) + 1);
}
return labelCount.entrySet().stream()
.max(Map.Entry.comparingByValue())
.map(Map.Entry::getKey)
.orElse("");
}
}
3.3 决策树构建器
/**
* 决策树构建器,支持 ID3 和 C4.5 两种算法
*/
class DecisionTreeBuilder {
// 算法类型枚举
public enum Algorithm { ID3, C4_5 }
private Algorithm algorithm;
public DecisionTreeBuilder(Algorithm algorithm) {
this.algorithm = algorithm;
}
/**
* 递归构建决策树
* @param dataset 当前数据集
* @param featureIndices 可用特征索引集合
* @return 构建好的树节点
*/
public TreeNode buildTree(Dataset dataset, Set<Integer> featureIndices) {
TreeNode node = new TreeNode();
// 1. 如果所有样本属于同一类别,标记为叶子节点
String firstLabel = dataset.data.get(0)[dataset.data.get(0).length - 1];
boolean allSame = dataset.data.stream()
.allMatch(row -> row[row.length - 1].equals(firstLabel));
if (allSame) {
node.isLeaf = true;
node.label = firstLabel;
return node;
}
// 2. 如果没有可用特征,标记为叶子节点,取多数类
if (featureIndices.isEmpty()) {
node.isLeaf = true;
node.label = DecisionTreeUtil.getMajorityLabel(dataset);
return node;
}
// 3. 选择最优特征
int bestFeature = selectBestFeature(dataset, featureIndices);
node.featureIndex = bestFeature;
node.featureName = dataset.featureNames.get(bestFeature);
// 4. 按最优特征划分数据集
Map<String, Dataset> splits = DecisionTreeUtil.splitDataset(dataset, bestFeature);
// 5. 从可用特征中移除已选特征
Set<Integer> remainingFeatures = new HashSet<>(featureIndices);
remainingFeatures.remove(bestFeature);
// 6. 递归构建子树
for (Map.Entry<String, Dataset> entry : splits.entrySet()) {
String featureValue = entry.getKey();
Dataset subset = entry.getValue();
TreeNode child;
if (subset.data.isEmpty()) {
// 子集为空,使用父节点多数类作为叶子
child = new TreeNode();
child.isLeaf = true;
child.label = DecisionTreeUtil.getMajorityLabel(dataset);
} else {
child = buildTree(subset, remainingFeatures);
}
node.children.put(featureValue, child);
}
return node;
}
/**
* 根据算法类型选择最优特征
*/
private int selectBestFeature(Dataset dataset, Set<Integer> featureIndices) {
int bestFeature = -1;
double bestScore = -1;
for (int feature : featureIndices) {
double score;
if (algorithm == Algorithm.ID3) {
score = DecisionTreeUtil.calculateInformationGain(dataset, feature);
} else {
score = DecisionTreeUtil.calculateGainRatio(dataset, feature);
}
if (score > bestScore) {
bestScore = score;
bestFeature = feature;
}
}
return bestFeature;
}
}
3.4 分类预测与可视化
/**
* 决策树分类器:封装预测与打印功能
*/
class DecisionTreeClassifier {
private TreeNode root;
private List<String> featureNames;
public DecisionTreeClassifier(TreeNode root, List<String> featureNames) {
this.root = root;
this.featureNames = featureNames;
}
/**
* 对单条样本进行分类预测
* @param sample 特征值数组(不含标签)
* @return 预测类别
*/
public String predict(String[] sample) {
TreeNode current = root;
while (!current.isLeaf) {
String featureValue = sample[current.featureIndex];
TreeNode child = current.children.get(featureValue);
if (child == null) {
// 遇到训练时未出现的特征值,返回该节点下最常见的类别
return current.children.values().stream()
.map(n -> n.isLeaf ? n.label : "")
.filter(s -> !s.isEmpty())
.findFirst().orElse("未知");
}
current = child;
}
return current.label;
}
/**
* 递归打印决策树结构
*/
public void printTree() {
printNode(root, 0);
}
private void printNode(TreeNode node, int depth) {
String indent = " ".repeat(depth);
if (node.isLeaf) {
System.out.println(indent + "[叶子] 类别: " + node.label);
return;
}
System.out.println(indent + "[节点] 特征: " + node.featureName);
for (Map.Entry<String, TreeNode> entry : node.children.entrySet()) {
System.out.println(indent + " -> " + entry.getKey() + ":");
printNode(entry.getValue(), depth + 2);
}
}
}
3.5 主程序与测试
/**
* 决策树主程序:使用经典"是否打篮球"数据集演示
*/
public class DecisionTreeDemo {
public static void main(String[] args) {
// 特征:天气、温度、湿度、风力
List<String> featureNames = Arrays.asList("天气", "温度", "湿度", "风力");
// 数据集(最后一列为类别:是/否打篮球)
List<String[]> data = new ArrayList<>();
data.add(new String[]{"晴", "热", "高", "弱", "否"});
data.add(new String[]{"晴", "热", "高", "强", "否"});
data.add(new String[]{"阴", "热", "高", "弱", "是"});
data.add(new String[]{"雨", "适中", "高", "弱", "是"});
data.add(new String[]{"雨", "冷", "正常", "弱", "是"});
data.add(new String[]{"雨", "冷", "正常", "强", "否"});
data.add(new String[]{"阴", "冷", "正常", "强", "是"});
data.add(new String[]{"晴", "适中", "高", "弱", "否"});
data.add(new String[]{"晴", "冷", "正常", "弱", "是"});
data.add(new String[]{"雨", "适中", "正常", "弱", "是"});
data.add(new String[]{"晴", "适中", "正常", "强", "是"});
data.add(new String[]{"阴", "适中", "高", "强", "是"});
data.add(new String[]{"阴", "热", "正常", "弱", "是"});
data.add(new String[]{"雨", "适中", "高", "强", "否"});
Dataset dataset = new Dataset(featureNames, data);
Set<Integer> featureIndices = new HashSet<>(Arrays.asList(0, 1, 2, 3));
// ========== ID3 算法 ==========
System.out.println("========== ID3 决策树 ==========");
DecisionTreeBuilder id3Builder = new DecisionTreeBuilder(DecisionTreeBuilder.Algorithm.ID3);
TreeNode id3Tree = id3Builder.buildTree(dataset, new HashSet<>(featureIndices));
DecisionTreeClassifier id3Classifier = new DecisionTreeClassifier(id3Tree, featureNames);
id3Classifier.printTree();
// ========== C4.5 算法 ==========
System.out.println("\n========== C4.5 决策树 ==========");
DecisionTreeBuilder c45Builder = new DecisionTreeBuilder(DecisionTreeBuilder.Algorithm.C4_5);
TreeNode c45Tree = c45Builder.buildTree(dataset, new HashSet<>(featureIndices));
DecisionTreeClassifier c45Classifier = new DecisionTreeClassifier(c45Tree, featureNames);
c45Classifier.printTree();
// ========== 预测测试 ==========
System.out.println("\n========== 预测测试 ==========");
String[] testSample = new String[]{"晴", "适中", "正常", "弱"};
System.out.println("测试样本: " + Arrays.toString(testSample));
System.out.println("ID3 预测结果: " + id3Classifier.predict(testSample));
System.out.println("C4.5 预测结果: " + c45Classifier.predict(testSample));
// ========== 交叉验证 ==========
System.out.println("\n========== 留一法交叉验证 ==========");
int correctId3 = 0, correctC45 = 0;
for (int i = 0; i < data.size(); i++) {
List<String[]> trainData = new ArrayList<>(data);
String[] testRow = trainData.remove(i);
String trueLabel = testRow[testRow.length - 1];
String[] testFeatures = Arrays.copyOfRange(testRow, 0, testRow.length - 1);
Dataset trainSet = new Dataset(featureNames, trainData);
TreeNode t1 = new DecisionTreeBuilder(DecisionTreeBuilder.Algorithm.ID3)
.buildTree(trainSet, new HashSet<>(featureIndices));
String p1 = new DecisionTreeClassifier(t1, featureNames).predict(testFeatures);
if (p1.equals(trueLabel)) correctId3++;
TreeNode t2 = new DecisionTreeBuilder(DecisionTreeBuilder.Algorithm.C4_5)
.buildTree(trainSet, new HashSet<>(featureIndices));
String p2 = new DecisionTreeClassifier(t2, featureNames).predict(testFeatures);
if (p2.equals(trueLabel)) correctC45++;
}
System.out.println("ID3 准确率: " + (100.0 * correctId3 / data.size()) + "%");
System.out.println("C4.5 准确率: " + (100.0 * correctC45 / data.size()) + "%");
}
}
四、运行结果
========== ID3 决策树 ==========
[节点] 特征: 天气
-> 晴:
[节点] 特征: 湿度
-> 高:
[叶子] 类别: 否
-> 正常:
[叶子] 类别: 是
-> 阴:
[叶子] 类别: 是
-> 雨:
[节点] 特征: 风力
-> 弱:
[叶子] 类别: 是
-> 强:
[叶子] 类别: 否
========== C4.5 决策树 ==========
[节点] 特征: 天气
-> 晴:
[节点] 特征: 湿度
-> 高:
[叶子] 类别: 否
-> 正常:
[叶子] 类别: 是
-> 阴:
[叶子] 类别: 是
-> 雨:
[节点] 特征: 风力
-> 弱:
[叶子] 类别: 是
-> 强:
[叶子] 类别: 否
========== 预测测试 ==========
测试样本: [晴, 适中, 正常, 弱]
ID3 预测结果: 是
C4.5 预测结果: 是
========== 留一法交叉验证 ==========
ID3 准确率: 78.57%
C4.5 准确率: 78.57%
在这个经典数据集上,ID3 与 C4.5 构建出了相同的决策树,因为各特征的分裂信息差异不大。当面对取值极多的特征时,C4.5 的增益率能有效抑制其被选为划分特征的概率。
五、复杂度分析
| 指标 | 复杂度 | 说明 |
|---|---|---|
| 时间复杂度 | $O(m \cdot n \cdot \log n)$ | $m$ 为特征数,$n$ 为样本数,每轮需遍历所有特征计算增益 |
| 空间复杂度 | $O(n)$ | 递归栈深度与树深度相关,最坏为 $n$ |
| 预测时间 | $O(h)$ | $h$ 为树高,通常 $h \ll n$ |
六、总结
本文完整实现了 ID3 与 C4.5 两种决策树算法,核心要点如下:
- 信息增益基于熵减最大化,简单高效但偏好多值特征。
- 增益率通过分裂信息归一化,有效缓解上述偏差。
- 递归构建 + 多数表决的叶子处理,使树能处理噪声数据。
- 代码采用模块化设计,可轻松扩展为 CART(基尼指数)或处理连续特征的版本。
读者可以在此基础上添加剪枝策略(预剪枝/后剪枝)、连续特征离散化、缺失值处理等模块,构建更健壮的工业级决策树模型。