共计 1666 个字符,预计需要花费 5 分钟才能阅读完成。
为什么需要决策树?
传统线性回归在遇到非线性数据时,表现往往不尽如人意。比如房价预测中,房屋面积和价格的关系可能在不同区间呈现不同趋势(小面积时单价高,大面积时单价低)。这时候线性模型就无法准确捕捉这种非线性关系。

决策树的优势在于:
- 天然适合处理非线性关系
- 对异常值不敏感
- 结果可解释性强
- 不需要对数据做太多预处理
算法原理详解
分类树 vs 回归树
虽然都叫决策树,但分类树和回归树有本质区别:
- 分类树预测的是离散类别
- 回归树预测的是连续数值
MSE 分裂准则
回归树的核心是找到最佳分裂点,使分裂后的两个子集的均方误差(MSE)最小。数学表达式为:
MSE = Σ(y_i - ŷ)^2 / n
其中 ŷ是子集中样本的均值。分裂时要遍历所有可能的分裂点,计算分裂后的 MSE,选择使 MSE 减少最多的特征和分裂点。
特征划分过程
- 从根节点开始
- 对每个特征,寻找最佳分裂点
- 选择使 MSE 减少最多的分裂
- 递归地对子节点重复上述过程
代码实战
环境准备
import numpy as np
import matplotlib.pyplot as plt
from sklearn.tree import DecisionTreeRegressor
from sklearn.model_selection import train_test_split
数据准备
# 生成非线性数据
np.random.seed(42)
X = np.sort(5 * np.random.rand(80, 1), axis=0)
y = np.sin(X).ravel() + np.random.randn(80) * 0.1
# 划分训练测试集
X_train, X_test, y_train, y_test = train_test_split(X, y, test_size=0.2)
模型训练
# 创建回归树模型
# max_depth 控制树的最大深度,防止过拟合
tree = DecisionTreeRegressor(max_depth=3)
# 训练模型
tree.fit(X_train, y_train)
# 评估
print("Train score:", tree.score(X_train, y_train))
print("Test score:", tree.score(X_test, y_test))
可视化决策路径
# 生成测试数据
X_test = np.arange(0.0, 5.0, 0.01)[:, np.newaxis]
y_pred = tree.predict(X_test)
# 绘图
plt.figure()
plt.scatter(X, y, s=20, edgecolor="black", c="darkorange", label="data")
plt.plot(X_test, y_pred, color="cornflowerblue", label="prediction")
plt.xlabel("data")
plt.ylabel("target")
plt.title("Decision Tree Regression")
plt.legend()
plt.show()
生产建议
解决过拟合
- 预剪枝:设置 max_depth、min_samples_split 等参数
- 后剪枝:训练完成后剪去不重要的分支
- Early Stopping:监控验证集表现,提前停止训练
重要参数调优
# 常用参数说明
params = {
'max_depth': 3, # 树的最大深度
'min_samples_split': 2, # 分裂所需最小样本数
'min_samples_leaf': 1, # 叶节点最少样本数
'max_features': None, # 考虑的特征数量
'random_state': 42 # 随机种子
}
性能考量
时间复杂度
- 训练:O(n_features × n_samples × log(n_samples))
- 预测:O(depth)
随机森林对比
- 单棵决策树容易过拟合
- 随机森林通过集成多棵树提高泛化能力
- 但随机森林牺牲了部分可解释性
延伸思考
- 尝试修改 max_depth 参数,观察模型表现如何变化
- 用 GridSearchCV 进行超参数调优
- 尝试用其他分裂标准(如 MAE)替代 MSE
正文完
