图解机器学习:线性回归模型
前言
最近有些好奇,大模型到底是如何训练的。每天都在跟 ChatGPT、Claude 一起工作,作为一个计算机行业的工作者,难免对它们的实现产生好奇心。
大模型与传统软件的区别
如今的大模型和我们过去认知中的软件并不是相同的东西。
过去开发传统软件时,我们会预先定义一系列规则。用户输入数据后,软件按照这些既定规则进行处理,并输出相应的答案。

而大模型的工作方式恰好相反。训练时,我们向模型提供大量输入及其对应答案,让模型从数据中学习两者之间的规律。传统软件是“先定义规则,再根据输入得到答案”;大模型则是“先提供输入和答案,再从中归纳规则”。

线性回归模型
如今的大模型本质上都是预测型模型,而线性回归作为最简单的预测模型,虽然体量微不足道,却已经浓缩了大模型的核心:模型 → 损失函数 → 优化。
所谓训练模型,本质上就是不断优化、让模型能力逐步升级。今天,我们就借线性回归这个最简单的模型,学习线性回归模型的三种求解方法:
- 穷举法:把参数的候选取值一个个代入试,谁让损失最小就用谁。
- 最小二乘法:直接对损失函数求导、令导数为 0,一步解出最优参数。
- 梯度下降法:沿损失函数梯度反方向一步步迭代,逼近最低点。
举个例子:小明在电脑城打工,我们不知道他每组装一台电脑具体能赚多少钱,只知道他每天组装了多少台电脑、当天拿到了多少工资。
| 电脑数量 x | 工资 y |
|---|---|
| 5 | 50 |
| 10 | 100 |
| 20 | 200 |
| 15 | 150 |
| …… | …… |
| x | y |
这时,聪明的你肯定已经想到设一个一元一次方程,只要找出 x 和 y 的关系,就能知道装一台电脑赚多少钱。
损失函数
损失函数就是衡量“模型预测”的记分器,分数越小越好。
还是沿用上面的例子,我们需要知道小明装一台机器赚多少钱。假设装一台机器赚 w 元,此时我们的方程就是 y = w * x。因为我们并不知道 w 到底是多少,所以就需要猜 w 到底是多少,然后代入 x,计算预测结果与实际结果之间有多少差别。
- 假设 w = 10:预测 50、100、200、150,和真实工资一分不差 → 没有任何损失。
- 假设 w = 5:预测 25、50、100、75,每个都差一截 → 损失很大。
- 假设 w = 15:预测 75、150、300、225,每个都超一截 → 损失也很大。
注意,损失函数和
w不是一回事。损失函数是我们自己定义的“记分器”,而w才是真正要找的参数。
我们可以发现,在不同的 w 下损失各不相同。我们需要算出误差的平均值,所以就需要用到下面这个公式:
L(w, b) = (1/n) · Σ (yᵢ − ŷᵢ)²
给不懂数学的小伙伴解释一下这段“鬼画符”的含义:
- L(w, b):损失函数,输入参数
w、b,输出一个“差多少”的分数。 - n:数据点个数。
- Σ:求和,把每个点算一遍再加起来。
- yᵢ:第 i 个点的真实值。
- ŷᵢ:第 i 个点的预测值,ŷᵢ = w * xᵢ + b。
由上面的例子可以得出,损失越小,直线越贴合;损失越大,直线越离谱。于是问题就统一转换为寻找让 L(w, b) 最小的参数。
如果把 b 固定,那么 L(w, b) 就是一个开口向上的抛物线(梦回高中数学),最低点就是最优解。
穷举法
穷举法,就是把 w 可能的范围切成很多份,得到一串候选值,然后一个一个代入计算损失,挑最小的那个。
- 优点:直观、几乎不需要数学;给定候选网格内一定能找到最小值。
- 缺点:计算量爆炸。
w试 100 个、b试 100 个,就是 100 × 100 = 1 万个组合,参数一多,维度相乘,直接算不动了。
穷举法适合建立直觉,不适合实战。
最小二乘法(解析解)
既然最低点处曲线是“平的”,也就是导数为 0,那就直接求导、令导数等于 0,把 w、b 解出来,一步到位。
对损失函数分别求偏导并令其为 0:
∂L/∂b = (2/n) Σ (w·xᵢ + b − yᵢ) = 0 … ①
∂L/∂w = (2/n) Σ (w·xᵢ + b − yᵢ)·xᵢ = 0 … ②
先解 ①,求 b:
Σ (w·xᵢ + b − yᵢ) = 0
n·b = Σ yᵢ − w·Σ xᵢ
b = ȳ − w·x̄ … ③
其中 x̄、ȳ 是均值。
把 ③ 代入 ②,求 w:
Σ (w·xᵢ + ȳ − w·x̄ − yᵢ)·xᵢ = 0
Σ [ w·(xᵢ − x̄) − (yᵢ − ȳ) ]·xᵢ = 0
w·Σ (xᵢ − x̄)·xᵢ = Σ (yᵢ − ȳ)·xᵢ
利用均值性质 Σ (xᵢ − x̄) = 0,可得 Σ (xᵢ − x̄)·xᵢ = Σ (xᵢ − x̄)²;同理,右边 Σ (yᵢ − ȳ)·xᵢ = Σ (yᵢ − ȳ)(xᵢ − x̄)。于是:
w = Σ (xᵢ − x̄)(yᵢ − ȳ) / Σ (xᵢ − x̄)² … ④
③④ 就是一元线性回归的闭式解。写成多元矩阵形式更简洁:
θ = (XᵀX)⁻¹ Xᵀ y
优点:一步算出、结果精确,不用迭代、不用选超参数。
缺点:要求 XᵀX 可逆(不可逆时要用伪逆或加正则项);特征很多时矩阵求逆代价高(约 O(d³)),内存也吃紧。
梯度下降法(逐步逼近)
最小二乘法“一步到位”虽爽,但特征很多、或损失函数复杂到解不出解析解时就不灵了。于是换个思路:不指望一步到位,而是沿梯度反方向,一步步“滑”到最低点。
先算损失对参数的偏导:
∂L/∂w = (2/n) Σ (w·xᵢ + b − yᵢ)·xᵢ
∂L/∂b = (2/n) Σ (w·xᵢ + b − yᵢ)
梯度方向是“上升最快”的方向,取反方向就是“下降最快”,反复迭代:
w ← w − α · ∂L/∂w
b ← b − α · ∂L/∂b
其中 α 是学习率,控制每步迈多大。
优点:能处理海量数据和高维参数;几乎适用于任何可求导的损失函数;配合小批量(mini-batch)可扩展到深度学习。
缺点:要调学习率 α——太大时震荡甚至发散,太小时收敛太慢;可能陷入局部最优(对线性回归这种凸损失则不会);结果是近似解。
三种方法对比
| 方法 | 求解方式 | 精度 | 适用场景 |
|---|---|---|---|
| 穷举法 | 暴力遍历候选值 | 网格内精确、网格外受限 | 教学演示、小参数空间 |
| 最小二乘法 | 求导 = 0 得解析解 | 精确 | 特征不多、XᵀX 可逆 |
| 梯度下降法 | 沿梯度反方向迭代 | 近似(可控制) | 大规模、高维、复杂损失 |
小结
三者的底层逻辑一致:把“拟合数据”转化为“最小化损失函数”,再去求这个最小值。
- 穷举法:最笨、最直观,帮你“看见”问题。
- 最小二乘法:数学上一步到位,适合小规模精确求解。
- 梯度下降法:工程上最通用,是大规模机器学习和深度学习的基础。
搞懂这三条路,线性回归乃至大部分机器学习的“训练”本质,就都通了。
💬 评论区