回归决策树和线性回归对比
1 课程概览
本课对比回归决策树和线性回归。两者都能做回归任务,但原理不同。线性回归拟合直线,回归决策树分段拟合。重点讲解两者的构建原理、损失函数和应用场景。
2 核心概念与定义
- 线性回归:拟合一条直线,预测连续值。
- 回归决策树:分段拟合,每段用均值预测。
- 平方损失:预测值减真实值的平方和。
- 最小二乘:误差的平方和最小化。
3 算法与模型详解
3.1 两种回归方法对比
| 特性 | 线性回归 | 回归决策树 |
|---|---|---|
| 拟合方式 | 直线 | 分段常数 |
| 损失函数 | 平方损失 | 平方损失 |
| 预测值 | $wx + b$ | 叶子节点均值 |
| 可解释性 | 强 | 中 |
| 非线性 | 不支持 | 支持 |
3.2 线性回归
原理:拟合一条直线 $y = wx + b$
损失函数:平方损失 $$\text{Loss} = \sum_{i=1}^{n} (y_i - (wx_i + b))^2$$
特点:
- 拟合直线
- 全局模型
- 不能拟合非线性关系
3.3 回归决策树
原理:分段拟合,每段用均值预测
损失函数:平方损失 $$\text{Loss} = \sum_{i=1}^{n} (y_i - \hat{y}_i)^2$$
其中 $\hat{y}_i$ 是叶子节点的均值
特点:
- 分段常数
- 局部模型
- 能拟合非线性关系
3.4 构建原理对比
线性回归:
- 假设 $y = wx + b$
- 最小化平方损失
- 求解 $w$ 和 $b$
回归决策树:
- 计算相邻特征值的中位数作为切分点
- 在每个切分点计算平方损失
- 选择平方损失最小的切分点
- 递归构建子树
- 叶子节点用均值预测
3.5 平方损失计算
线性回归: $$\text{Loss} = \sum_{i=1}^{n} (y_i - (wx_i + b))^2$$
回归决策树: $$\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$$
相同点:都是预测值减真实值的平方和
不同点:
- 线性回归:预测值 = $wx + b$
- 回归决策树:预测值 = 叶子节点均值
3.6 切分点选择
中位数计算:
- 特征值:1, 2, 3, 4, 5
- 中位数:1.5, 2.5, 3.5, 4.5
切分示例:
- 在1.5切分:
- 左侧:1个样本
- 右侧:4个样本
- 左侧预测值 = 左侧标签均值
- 右侧预测值 = 右侧标签均值
- 平方损失 = Σ(预测值 - 真实值)²
3.7 应用场景
线性回归适用:
- 线性关系
- 需要强可解释性
- 数据量小
回归决策树适用:
- 非线性关系
- 需要分段拟合
- 数据量大
4 数学原理与推导
4.1 线性回归损失
$$\text{Loss}{LR} = \sum{i=1}^{n} (y_i - (wx_i + b))^2$$
4.2 回归决策树损失
$$\text{Loss}{DT}(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$$
4.3 最优切分点
$$s^* = \arg\min_s \text{Loss}_{DT}(s)$$
4.4 预测值
线性回归: $$\hat{y} = wx + b$$
回归决策树: $$\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
from sklearn.linear_model import LinearRegression
# 1. 生成数据
np.random.seed(42)
X = np.sort(5 * np.random.rand(100, 1), axis=0)
y = np.sin(X).ravel() + np.random.normal(0, 0.1, X.shape[0])
print("=== 数据信息 ===")
print(f"X形状: {X.shape}")
print(f"y形状: {y.shape}")
print(f"X范围: [{X.min():.2f}, {X.max():.2f}]")
print(f"y范围: [{y.min():.2f}, {y.max():.2f}]")
# 2. 线性回归
print("\n=== 线性回归 ===")
lr = LinearRegression()
lr.fit(X, y)
y_pred_lr = lr.predict(X)
print(f"系数w: {lr.coef_[0]:.4f}")
print(f"截距b: {lr.intercept_:.4f}")
print(f"R²: {lr.score(X, y):.4f}")
# 3. 回归决策树
print("\n=== 回归决策树 ===")
dt = DecisionTreeRegressor(max_depth=4, random_state=42)
dt.fit(X, y)
y_pred_dt = dt.predict(X)
print(f"R²: {dt.score(X, y):.4f}")
# 4. 对比平方损失
print("\n=== 平方损失对比 ===")
loss_lr = np.sum((y - y_pred_lr) ** 2)
loss_dt = np.sum((y - y_pred_dt) ** 2)
print(f"线性回归平方损失: {loss_lr:.4f}")
print(f"回归决策树平方损失: {loss_dt:.4f}")
# 5. 可视化对比
plt.figure(figsize=(15, 5))
# 训练数据
plt.subplot(1, 3, 1)
plt.scatter(X, y, color='blue', s=20, label='训练数据')
plt.xlabel('X')
plt.ylabel('y')
plt.title('原始数据')
plt.legend()
plt.grid(True)
# 线性回归
plt.subplot(1, 3, 2)
plt.scatter(X, y, color='blue', s=20, label='训练数据')
plt.plot(X, y_pred_lr, color='red', linewidth=2, label='线性回归')
plt.xlabel('X')
plt.ylabel('y')
plt.title(f'线性回归 (R²={lr.score(X, y):.4f})')
plt.legend()
plt.grid(True)
# 回归决策树
plt.subplot(1, 3, 3)
plt.scatter(X, y, color='blue', s=20, label='训练数据')
plt.plot(X, y_pred_dt, color='green', linewidth=2, label='回归决策树')
plt.xlabel('X')
plt.ylabel('y')
plt.title(f'回归决策树 (R²={dt.score(X, y):.4f})')
plt.legend()
plt.grid(True)
plt.tight_layout()
plt.show()
# 6. 预测新数据
print("\n=== 预测新数据 ===")
X_test = np.linspace(0, 5, 200).reshape(-1, 1)
y_pred_lr_new = lr.predict(X_test)
y_pred_dt_new = dt.predict(X_test)
# 7. 不同深度的回归决策树
print("\n=== 不同深度的回归决策树 ===")
depths = [1, 2, 4, 10]
plt.figure(figsize=(15, 10))
for i, depth in enumerate(depths, 1):
dt_depth = DecisionTreeRegressor(max_depth=depth, random_state=42)
dt_depth.fit(X, y)
y_pred = dt_depth.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', linewidth=2, label=f'深度={depth}')
plt.xlabel('X')
plt.ylabel('y')
plt.title(f'回归决策树 (深度={depth}, R²={dt_depth.score(X, y):.4f})')
plt.legend()
plt.grid(True)
plt.tight_layout()
plt.show()
# 8. 手动计算平方损失示例
print("\n=== 手动计算平方损失示例 ===")
# 简单数据
X_simple = np.array([1, 2, 3, 4, 5]).reshape(-1, 1)
y_simple = np.array([5.56, 5.70, 5.91, 6.40, 6.80])
print(f"X: {X_simple.ravel()}")
print(f"y: {y_simple}")
# 在2.5切分
split = 2.5
left_mask = X_simple.ravel() <= split
right_mask = X_simple.ravel() > split
left_y = y_simple[left_mask]
right_y = y_simple[right_mask]
left_pred = left_y.mean()
right_pred = right_y.mean()
left_loss = np.sum((left_y - left_pred) ** 2)
right_loss = np.sum((right_y - right_pred) ** 2)
total_loss = left_loss + right_loss
print(f"\n在{split}切分:")
print(f"左侧: {left_y}, 均值={left_pred:.4f}, 损失={left_loss:.4f}")
print(f"右侧: {right_y}, 均值={right_pred:.4f}, 损失={right_loss:.4f}")
print(f"总损失: {total_loss:.4f}")
输出示例:
=== 数据信息 ===
X形状: (100, 1)
y形状: (100,)
X范围: [0.01, 4.96]
y范围: [-0.19, 1.10]
=== 线性回归 ===
系数w: 0.1547
截距b: 0.2154
R²: 0.5217
=== 回归决策树 ===
R²: 0.9365
=== 平方损失对比 ===
线性回归平方损失: 12.8123
回归决策树平方损失: 1.6978
6 重难点与易错提醒
- ❗重点:线性回归拟合直线,回归决策树分段拟合。
- ❗重点:两者都用平方损失。
- ❗重点:回归决策树能拟合非线性关系。
- ❗重点:回归决策树预测值 = 叶子节点均值。
- ⚠️易错:混淆两种方法的预测值计算。
- ⚠️易错:回归决策树深度过大导致过拟合。
- 💡深入理解:回归决策树通过分段常数拟合非线性数据。
7 课堂问答精选
Q: 回归决策树和线性回归有什么区别?
A:
- 线性回归:拟合一条直线 $y = wx + b$,预测值 = $wx + b$
- 回归决策树:分段拟合,预测值 = 叶子节点均值
两者都用平方损失,但回归决策树能拟合非线性关系,而线性回归只能拟合线性关系。
Q: 什么时候用回归决策树,什么时候用线性回归?
A:
- 线性回归:数据呈线性关系,需要强可解释性
- 回归决策树:数据呈非线性关系,需要分段拟合
如果数据是非线性的(如正弦曲线),回归决策树效果更好;如果是线性的,线性回归更简单有效。
8 本课小结
- 线性回归:拟合直线,预测值 = $wx + b$。
- 回归决策树:分段拟合,预测值 = 叶子节点均值。
- 两者都用平方损失。
- 回归决策树能拟合非线性关系。
- 线性回归适合线性数据,回归决策树适合非线性数据。
9 延伸思考与实践
- 实践:对比线性回归和回归决策树。
- 预习:决策树剪枝介绍。
- 思考:如何选择回归决策树的深度?