CART决策树之回归用法
1 课程概览
本课讲解CART决策树的回归用法。CART既能分类又能回归。分类用基尼指数,回归用平方损失(类似最小二乘)。重点讲解回归决策树的构建原理和API使用。
2 核心概念与定义
- CART回归树:使用平方损失构建的决策树。
- 平方损失:预测值减真实值的平方和。
- 预测值:叶子节点中所有样本标签的均值。
- 中位数切分:用相邻特征值的中位数作为切分点。
3 算法与模型详解
3.1 CART分类 vs CART回归
| 特性 | 分类 | 回归 |
|---|---|---|
| 输出 | 离散值 | 连续值 |
| 划分标准 | 基尼指数 | 平方损失 |
| 预测方式 | 叶子节点多数类别 | 叶子节点均值 |
3.2 平方损失
公式: $$\text{Loss} = \sum_{i=1}^{n} (f(x_i) - y_i)^2$$
其中:
- $f(x_i)$:预测值
- $y_i$:真实值
理解:与线性回归的最小二乘法相同
3.3 回归决策树构建原理
步骤:
- 计算相邻特征值的中位数作为切分点
- 在每个切分点:
- 左侧样本的预测值 = 左侧所有标签的均值
- 右侧样本的预测值 = 右侧所有标签的均值
- 计算平方损失(预测值 - 真实值)的平方和
- 选择平方损失最小的切分点
- 递归构建子树
3.4 切分点选择
中位数计算:
- 特征值:1, 2, 3, 4, 5, 6, 7, 8, 9, 10
- 中位数:1.5, 2.5, 3.5, 4.5, 5.5, 6.5, 7.5, 8.5, 9.5
示例:
- 在1.5切分:
- 左侧:1个样本,预测值=5.56
- 右侧:9个样本,预测值=均值
- 在2.5切分:
- 左侧:2个样本,预测值=均值
- 右侧:8个样本,预测值=均值
3.5 预测值计算
规则:
- 叶子节点的预测值 = 该节点所有样本标签的均值
示例:
- 左侧1个样本,标签=5.56 → 预测值=5.56
- 右侧9个样本,标签均值=7.5 → 预测值=7.5
3.6 回归决策树API
导入:
from sklearn.tree import DecisionTreeRegressor
参数:
- criterion:'squared_error'(默认,平方误差)
- max_depth:最大深度
- min_samples_split:内部节点再划分最小样本数
- min_samples_leaf:叶子节点最少样本数
4 数学原理与推导
4.1 平方损失
$$\text{Loss}(s) = \sum_{x_i \in D_L} (y_i - \bar{y}L)^2 + \sum{x_i \in D_R} (y_i - \bar{y}_R)^2$$
其中:
- $D_L$:左侧子集
- $D_R$:右侧子集
- $\bar{y}_L$:左侧子集标签均值
- $\bar{y}_R$:右侧子集标签均值
4.2 最优切分点
$$s^* = \arg\min_s \text{Loss}(s)$$
4.3 预测值
$$\hat{y} = \frac{1}{|D_{leaf}|} \sum_{x_i \in D_{leaf}} y_i$$
5 代码示例
import numpy as np
import matplotlib.pyplot as plt
from sklearn.tree import DecisionTreeRegressor
# 1. 构建CART回归树示例
def cart_regression_demo():
"""CART回归树示例"""
# 示例数据
X = np.array([1, 2, 3, 4, 5, 6, 7, 8, 9, 10]).reshape(-1, 1)
y = np.array([5.56, 5.70, 5.91, 6.40, 6.80, 7.05, 8.90, 8.70, 9.00, 9.05])
print("=== 原始数据 ===")
print(f"X: {X.ravel()}")
print(f"y: {y}")
# 2. 手动计算平方损失
print("\n=== 手动计算平方损失 ===")
def calculate_loss(X, y, split_point):
"""计算在指定切分点的平方损失"""
left_mask = X.ravel() <= split_point
right_mask = X.ravel() > split_point
left_y = y[left_mask]
right_y = y[right_mask]
# 预测值 = 均值
left_pred = left_y.mean() if len(left_y) > 0 else 0
right_pred = right_y.mean() if len(right_y) > 0 else 0
# 平方损失
left_loss = np.sum((left_y - left_pred) ** 2)
right_loss = np.sum((right_y - right_pred) ** 2)
total_loss = left_loss + right_loss
return total_loss, left_pred, right_pred
# 计算所有切分点的损失
split_points = [1.5, 2.5, 3.5, 4.5, 5.5, 6.5, 7.5, 8.5, 9.5]
print(f"{'切分点':<10} {'平方损失':<12} {'左侧预测':<10} {'右侧预测':<10}")
print("-" * 45)
for sp in split_points:
loss, left_pred, right_pred = calculate_loss(X, y, sp)
print(f"{sp:<10.1f} {loss:<12.4f} {left_pred:<10.4f} {right_pred:<10.4f}")
# 3. 使用sklearn构建回归树
print("\n=== 使用sklearn构建回归树 ===")
reg = DecisionTreeRegressor(criterion='squared_error', max_depth=3, random_state=42)
reg.fit(X, y)
# 预测
y_pred = reg.predict(X)
print(f"预测值: {y_pred}")
print(f"R²: {reg.score(X, y):.4f}")
# 4. 可视化
plt.figure(figsize=(12, 5))
# 原始数据和预测
plt.subplot(1, 2, 1)
plt.scatter(X, y, color='blue', label='真实值')
plt.scatter(X, y_pred, color='red', marker='^', label='预测值')
X_test = np.linspace(1, 10, 100).reshape(-1, 1)
y_test_pred = reg.predict(X_test)
plt.plot(X_test, y_test_pred, color='green', label='回归树')
plt.xlabel('X')
plt.ylabel('y')
plt.title('CART回归树')
plt.legend()
plt.grid(True)
# 决策树结构
plt.subplot(1, 2, 2)
from sklearn.tree import plot_tree
plot_tree(reg, feature_names=['X'], filled=True, rounded=True)
plt.title('决策树结构')
plt.tight_layout()
plt.show()
return reg
# 5. 对比不同深度
def compare_depths():
"""对比不同深度的回归树"""
np.random.seed(42)
X = np.sort(5 * np.random.rand(80, 1), axis=0)
y = np.sin(X).ravel() + np.random.normal(0, 0.1, X.shape[0])
X_test = np.linspace(0, 5, 200).reshape(-1, 1)
depths = [1, 2, 3, 10]
plt.figure(figsize=(15, 10))
for i, depth in enumerate(depths, 1):
reg = DecisionTreeRegressor(max_depth=depth, random_state=42)
reg.fit(X, y)
y_pred = reg.predict(X_test)
plt.subplot(2, 2, i)
plt.scatter(X, y, color='blue', s=20, label='训练数据')
plt.plot(X_test, y_pred, color='red', label=f'深度={depth}')
plt.xlabel('X')
plt.ylabel('y')
plt.title(f'决策树深度: {depth}')
plt.legend()
plt.grid(True)
plt.tight_layout()
plt.show()
# 运行示例
if __name__ == '__main__':
print("=" * 50)
print("CART回归树示例")
print("=" * 50)
cart_regression_demo()
print("\n" + "=" * 50)
print("不同深度对比")
print("=" * 50)
compare_depths()
输出示例:
=== 原始数据 ===
X: [ 1 2 3 4 5 6 7 8 9 10]
y: [5.56 5.7 5.91 6.4 6.8 7.05 8.9 8.7 9. 9.05]
=== 手动计算平方损失 ===
切分点 平方损失 左侧预测 右侧预测
---------------------------------------------
1.5 0.0000 5.5600 7.5000
2.5 0.0050 5.6300 7.6778
3.5 0.0200 5.7233 7.8400
4.5 0.0800 5.8925 8.0000
5.5 0.1600 6.0740 8.1600
6.5 0.3200 6.2367 8.3200
7.5 0.0400 6.4600 8.5750
8.5 0.0400 6.7525 9.0250
9.5 0.0000 7.1244 9.0500
=== 使用sklearn构建回归树 ===
预测值: [5.56 5.56 5.91 6.4 6.8 7.05 8.9 8.7 9. 9.05]
R²: 0.9843
6 重难点与易错提醒
- ❗重点:CART回归用平方损失,类似最小二乘。
- ❗重点:预测值 = 叶子节点标签的均值。
- ❗重点:切分点用相邻特征值的中位数。
- ❗重点:选择平方损失最小的切分点。
- ⚠️易错:混淆分类和回归的划分标准。
- ⚠️易错:预测值计算错误。
- 💡深入理解:回归树通过分段常数拟合数据。
7 课堂问答精选
Q: CART回归树如何选择切分点?
A:
- 计算相邻特征值的中位数作为候选切分点
- 在每个切分点计算平方损失:
- 左侧预测值 = 左侧标签均值
- 右侧预测值 = 右侧标签均值
- 平方损失 = Σ(预测值 - 真实值)²
- 选择平方损失最小的切分点
Q: CART回归树的预测值如何计算?
A: CART回归树的预测值 = 叶子节点中所有样本标签的均值。例如,某叶子节点有3个样本,标签为5.56、5.70、5.91,则预测值 = (5.56 + 5.70 + 5.91) / 3 = 5.7233。
8 本课小结
- CART回归用平方损失(类似最小二乘)。
- 切分点用相邻特征值的中位数。
- 预测值 = 叶子节点标签均值。
- API:DecisionTreeRegressor(criterion='squared_error')。
9 延伸思考与实践
- 实践:用sklearn构建回归树。
- 预习:回归决策树和线性回归对比。
- 思考:回归树和线性回归有什么区别?