一元线性回归之正规方程法
1 课程概览
本课讲解一元线性回归的正规方程法求解过程。通过损失函数对K和B分别求偏导,令偏导等于0,联立方程组求解K和B的最优值。详细推导了正规方程的数学过程,说明该方法为一步到位求解方法,适合小数据量,大数据量时可能内存溢出。
2 核心概念与定义
- 正规方程法:通过对损失函数求偏导,令偏导等于0,直接求解最优参数的方法。
- 偏导数:对某一变量求导,其他变量视为常数。
- h(Xi):第i个样本的预测值。
- Yi:第i个样本的真实值。
- m:样本总数。
3 算法与模型详解
3.1 正规方程法思路
步骤:
- 写出损失函数
- 对K求偏导,令其等于0
- 对B求偏导,令其等于0
- 联立方程组求解K和B
3.2 损失函数
公式: $$L(K, B) = \sum_{i=1}^{m} (h(X_i) - Y_i)^2 = \sum_{i=1}^{m} (KX_i + B - Y_i)^2$$
说明:
- $h(X_i) = KX_i + B$:第i个样本的预测值
- $Y_i$:第i个样本的真实值
- $m$:样本总数
- 预测值 - 真实值 = 误差
3.3 求解过程
步骤1:对K求偏导 $$\frac{\partial L}{\partial K} = \sum_{i=1}^{m} 2(KX_i + B - Y_i) \cdot X_i = 0$$
步骤2:对B求偏导 $$\frac{\partial L}{\partial B} = \sum_{i=1}^{m} 2(KX_i + B - Y_i) = 0$$
步骤3:联立方程组求解
由步骤2可得: $$\sum_{i=1}^{m} (KX_i + B - Y_i) = 0$$
$$K \sum_{i=1}^{m} X_i + mB - \sum_{i=1}^{m} Y_i = 0$$
$$B = \frac{1}{m} \sum_{i=1}^{m} Y_i - K \cdot \frac{1}{m} \sum_{i=1}^{m} X_i = \bar{Y} - K\bar{X}$$
其中 $\bar{X}$ 和 $\bar{Y}$ 分别是X和Y的均值。
步骤4:代入求K
将B代入步骤1的方程,求解K: $$K = \frac{\sum_{i=1}^{m} X_i(Y_i - \bar{Y})}{\sum_{i=1}^{m} X_i(X_i - \bar{X})}$$
3.4 正规方程法特点
| 特点 | 说明 |
|---|---|
| 求解方式 | 一步到位,直接求解 |
| 优点 | 精确解,无需迭代 |
| 缺点 | 数据量大时可能内存溢出 |
| 适用场景 | 小数据量 |
3.5 与梯度下降法对比
| 对比维度 | 正规方程法 | 梯度下降法 |
|---|---|---|
| 求解方式 | 一步到位 | 逐步逼近 |
| 计算量 | 大(矩阵求逆) | 小(迭代) |
| 内存 | 数据量大时溢出 | 适合大数据量 |
| 精度 | 精确解 | 近似解 |
| 适用场景 | 小数据量 | 大数据量 |
4 数学原理与推导
4.1 损失函数
$$L(K, B) = \sum_{i=1}^{m} (KX_i + B - Y_i)^2$$
4.2 对B求偏导
$$\frac{\partial L}{\partial B} = 2\sum_{i=1}^{m} (KX_i + B - Y_i) = 0$$
$$\Rightarrow B = \bar{Y} - K\bar{X}$$
4.3 对K求偏导
$$\frac{\partial L}{\partial K} = 2\sum_{i=1}^{m} (KX_i + B - Y_i)X_i = 0$$
代入B后求解K: $$K = \frac{\sum_{i=1}^{m} X_i(Y_i - \bar{Y})}{\sum_{i=1}^{m} X_i(X_i - \bar{X})}$$
4.4 最终结果
$$K = \frac{\sum_{i=1}^{m} (X_i - \bar{X})(Y_i - \bar{Y})}{\sum_{i=1}^{m} (X_i - \bar{X})^2}$$
$$B = \bar{Y} - K\bar{X}$$
5 代码示例
import numpy as np
# 训练数据
X = np.array([160, 166, 172, 174, 180])
Y = np.array([56.3, 60.6, 65.1, 68.5, 75.0])
# 计算均值
X_mean = np.mean(X)
Y_mean = np.mean(Y)
# 正规方程法求解K和B
numerator = np.sum((X - X_mean) * (Y - Y_mean))
denominator = np.sum((X - X_mean) ** 2)
K = numerator / denominator
B = Y_mean - K * X_mean
print("权重K(斜率):", K)
print("偏置B(截距):", B)
# 预测
X_test = 176
Y_predict = K * X_test + B
print("预测值:", Y_predict)
输出示例:
权重K(斜率): 0.8582142857142856
偏置B(截距): -81.04857142857142
预测值: 70.30476190476191
6 重难点与易错提醒
- ❗重点:正规方程法通过对K和B求偏导,令偏导等于0求解。
- ❗重点:B = Ȳ - K·X̄(截距公式)
- ❗重点:K的求解公式涉及均值和协方差。
- ❗重点:正规方程法适合小数据量,大数据量用梯度下降。
- ⚠️易错:求偏导时忘记链式法则。
- ⚠️易错:联立方程时计算错误。
- 💡深入理解:正规方程法是精确解,梯度下降法是近似解。
7 课堂问答精选
Q: 正规方程法和梯度下降法有什么区别?
A:
- 正规方程法:一步到位直接求解,精确解,适合小数据量,大数据量可能内存溢出。
- 梯度下降法:逐步逼近最优解,近似解,适合大数据量。 选择依据:数据量小用正规方程,数据量大用梯度下降。
Q: 为什么正规方程法大数据量会内存溢出?
A: 正规方程法需要计算矩阵的逆,当数据量很大时,矩阵维度很大,计算逆矩阵需要大量内存和计算资源,可能导致内存溢出。梯度下降法通过迭代逐步逼近,不需要计算逆矩阵,适合大数据量。
8 本课小结
- 正规方程法:对损失函数求偏导,令偏导等于0,联立方程组求解。
- 损失函数:$L(K, B) = \sum (KX_i + B - Y_i)^2$
- B = Ȳ - K·X̄
- K = $\frac{\sum (X_i - \bar{X})(Y_i - \bar{Y})}{\sum (X_i - \bar{X})^2}$
- 特点:一步到位,精确解,适合小数据量。
- 对比:梯度下降法逐步逼近,适合大数据量。
9 延伸思考与实践
- 实践:用NumPy手动实现正规方程法。
- 预习:多元线性回归的正规方程法。
- 思考:为什么正规方程法需要求矩阵的逆?