假设:
| 人 | 饭量 X (碗) | 体重 Y (斤) |
|---|---|---|
| 小明 | 1 | 62(比预期重了2斤) |
| 小红 | 2 | 68(比预期轻了2斤) |
| 小刚 | 3 | 82(比预期重了2斤) |
数据(1,62)、(2,68)、(3,82)来手算。
第一步:写出“总误差平方和”的式子
假设我们猜的直线是 Y=aX+b。
当 X=1,预测值是 a+b,误差是 62−(a+b)
当 X=2,预测值是 2a+b,误差是 68−(2a+b)
当 X=3,预测值是 3a+b,误差是 82−(3a+b)
总误差平方和(记作 S):
S=(62−a−b)2+(68−2a−b)2+(82−3a−b)2
最小二乘法的目标:找到合适的 a 和 b,让这个 S最小。
第二步:求偏导等于 0(变成两个“普通方程”)
“求偏导”的数学含义:把 S 分别对 a 和 b 求导,并令其等于 0。你可以理解为“当误差平方和降到谷底时,它再也下不去了”。
我们把上面那个复杂的平方式子分别求导,化简后神奇地变成了下面这两个简单的二元一次方程组:
方程①(对 b 求导化简):
6a+3b=212
方程②(对 a 求导化简):
14a+6b=444
(如果你好奇怎么变的:就是把括号里的常数相加,把 aa 和 bb 的系数归类,和中学的合并同类项完全一样。)
第三步:解这个“二元一次方程组”
我们现在有两个方程:
6a+3b=212
14a+6b=444
消元法:
把第一个方程两边同时乘以 2,让 b 的系数变成和第二个方程一样:
方程① × 2 得:12a+6b=424 ……(③)
用方程② 减去 ③:
(14a−12a)+(6b−6b)=444−424
2a=20
所以:a=10(斜率出来了!)
把 a=10a=10 代回最简单的方程①:
6×10+3b=212
60+3b=212
3b=152
所以:b=152/3≈50.67(截距出来了!)
最终结果验证
把 a=10,b=50.6 代回预测公式(Y=10X+50.67):
X=1 时,预测 60.67,真实 62,误差 +1.33
X=2 时,预测 70.67,真实 68,误差 -2.67
X=3 时,预测 80.67,真实 82,误差 +1.33
算一下误差平方和:1.332+(−2.67)2+1.332≈1.77+7.13+1.77=10.67
这个 10.67完美验证了“最小”二字。
我们用协方差可以更快地方式求的a和b的值:
(1,62)、(2,68)、(3,82)这三个数
第 1 步:算出 X 和 Y 的平均值(小学数学)
X 的平均值(记作 Xˉ):(1+2+3)÷3=2(1+2+3)÷3=2
Y 的平均值(记作 Yˉ):(62+68+82)÷3=70.67
第 2 步:列一张“三列”小表格(这是口算的关键)
我们算一下每个数离平均值有多远(这叫“中心化”):
| 数据点 | X 的离差 (X - 平均值2) | Y 的离差 (Y - 平均值70.67) | 乘积(判断步调) (① × ②) | X离差的平方 (① × ①) |
|---|---|---|---|---|
| 点1 (1,62) | 1 - 2 =-1 | 62 - 70.67 =-8.67 | (-1) × (-8.67) =+8.67 | (-1)² =1 |
| 点2 (2,68) | 2 - 2 =0 | 68 - 70.67 =-2.67 | 0 × (-2.67) =0 | 0² =0 |
| 点3 (3,82) | 3 - 2 =+1 | 82 - 70.67 =+11.33 | (+1) × 11.33 =+11.33 | 1² =1 |
| 求和 | 分子总和 = 20 | 分母总和 = 2 |
第 3 步:直接套公式,两秒出斜率
第 4 步:口算截距 b(附赠一个铁律)
统计学有个铁律:最优直线一定会穿过X的平均值和Y的平均值的那个交叉点(即点 (Xˉ,Yˉ))。
既然直线是 Y=aX+b,把平均值点(2, 70.67)和刚算的斜率 10 代进去:
70.67=10×2+b
b=70.67−20=50.67
一定要从最简单的方式去了解他的原理,方可举一反三。
最小二乘法 ≠ 线性回归,最小二乘法是线性回归的“金牌打工人”,专门负责帮它算出最优的斜率和截距。
最小二乘法编程:java实现,python更简单
方法一:
public static void main(String[] args) { double[] x = {4.6, 5.25, 8.9, 3.4, 1.6, 10.1, 4.8}; double[] y = {108, 114, 130, 114, 106, 150, 122}; int n = x.length; double sumX = 0, sumY = 0, sumXY = 0, sumX2 = 0; for (int i = 0; i < n; i++) { sumX += x[i]; sumY += y[i]; sumXY += x[i] * y[i]; sumX2 += x[i] * x[i]; } // 极简公式(仅适用于一元线性回归) double a = (n * sumXY - sumX * sumY) / (n * sumX2 - sumX * sumX); double b = (sumY - a * sumX) / n; System.out.printf("拟合直线为:Y = %.4f X + %.4f\n", a, b); // 预测 X = 6.2 时的值 double pred = a * 6.2 + b; System.out.printf("当 X=6.2 时,预测 Y = %.2f\n", pred); }方法二:
梯度下降法球a和b
public static void main(String[] args) { double[] x = {4.6, 5.25, 8.9, 3.4, 1.6, 10.1, 4.8}; double[] y = {108, 114, 130, 114, 106, 150, 122}; double a = 0, b = 0; double rate = 0.0001; // 学习率(步长) // 迭代 1000 次,让 a,b 自动逼近最优值 for (int iter = 0; iter < 2000; iter++) { double dS_da = 0, dS_db = 0; for (int i = 0; i < x.length; i++) { double pred = a * x[i] + b; double error = y[i] - pred; dS_da += -2 * x[i] * error; // 等价于 2*error*(-x[i]) dS_db += -2 * error; } // 更新参数(减去梯度) a -= rate * dS_da; b -= rate * dS_db; } System.out.printf("迭代后:a=%.4f, b=%.4f\n", a, b); System.out.printf("预测 X=6.2: Y=%.2f\n", a * 6.2 + b); }