跳到主要内容

第一个完整的训练循环 — 线性回归从零到一

这一章讲一件事: 把第 01 章那张四步循环,变成一套能跑的最小系统。 模型是线性回归——全书最简单的模型,也是唯一有解析解的; 但我们的目标不是这个模型本身, 而是把「模型→损失→优化→评估」这条流水线在脑子里走通一遍。 后面所有复杂的网络,都是这条流水线在四个环节上分别升级。

1. 顶层全景:一条流水线,四个工位

数据 (面积, 房龄) → 价格


① 模型:ŷ = w₁x₁ + w₂x₂ + b ← 一条带参数的直线


② 损失:l = ½(ŷ − y)² ← 预测离真值多远


③ 优化:minibatch SGD ← 参数往损失变小的方向挪


④ 评估:在没见过的数据上量误差 ← 这才是真正的目标

图说:本章主走查——用合成数据(真参数已知)把这条线跑通,
看学出来的参数能不能逼近真参数。

2. 模型:一条带参数的直线

第 01 章的水管工例子已经有了线性回归的雏形。 正式写法:给定特征(面积、房龄),预测价格——

价格 = w_面积 × 面积 + w_房龄 × 房龄 + b

w 叫权重,决定每个特征的影响力;b 叫偏置, 是「所有特征都为零」时的预测值。 没有 b,直线必须过原点——那样就表达不了「零面积也要收上门费」这类关系。 严格说,这个形式叫仿射变换:线性变换(加权和)加一次平移(加 b)1

把它写成线性代数的样子:特征拼成向量 x,权重拼成向量 w, 则 ŷ = wᵀx + b——一次点积加一个数,这就是第 02 章说的「每层一次矩阵-向量积」的最简形态。 整个数据集 n 条样本拼成设计矩阵 X(一行一个样本,一列一个特征), 所有预测一次算完:ŷ = Xw + b。

3. 损失:为什么是平方误差

第 01 章说损失是「越低越好的尺」。回归问题里最常用的那把尺是平方误差2:

l = ½(ŷ − y)²

系数 ½ 只是为了让求导后式子干净(导数里的 2 和 ½ 约掉),没有实质含义。 对全部 n 条样本取平均,就得到总损失 L(w, b)。

平方这个形状是把双刃剑:预测差两倍,损失差四倍—— 大误差被加倍惩罚,这逼着模型避开大错; 但反过来,一个离谱的异常数据也会把模型狠狠拽过去3

为什么偏偏是平方,而不是绝对值或四次方? 书里给了一个更根本的理由:如果数据是「真实线性关系 + 高斯噪声」生成的, 那么最小化平方误差,恰好等于挑出「让已观测到的数据最可能出现」的那组参数4

这种挑法有个名字:先定义似然(likelihood,「这组参数之下,已观测数据出现的概率」), 再取让似然最大的那组参数——所以叫最大似然估计。

把高斯噪声的概率公式取负对数(log,把连乘变连加的那个函数)展开,里面就躺着平方误差这一项—— 平方误差不是惯例,是「噪声是高斯的」这个假设的推论。 假设换了,尺就该换。

4. 优化:解析解存在,但别指望它

线性回归是全书唯一可以直接「解」出来的模型: 把损失对 w 求导、令其为零,直接写出最优参数 w* = (XᵀX)⁻¹Xᵀy(要求 XᵀX 可逆,即没有特征能被其他特征线性表出)5

作者对此的提醒值得记住:「你别习惯这种好运。」 解析解的要求苛刻到会排除深度学习里几乎所有有趣的东西6—— 层数一多、非线性一进,方程就解不开了。 所以真正通用的办法是迭代的数值优化:梯度下降

全量梯度下降每走一步要把全部数据过一遍——数据集一大就慢死, 而且数据里的冗余让全量更新浪费。 另一个极端是每次只看一条样本(随机梯度下降,SGD)—— 快是快,但有两个毛病: 一条一条算,硬件上比批量算低效一个数量级(处理器做矩阵乘比做一连串向量乘划算得多)7; 而且有些层(后面讲的批归一化)天生需要一次看到多条样本。

折中就是全书此后一直用的小批量随机梯度下降(minibatch SGD): 每次随机抽一小批(通常 32 到 256 条,取 2 的幂附近)8, 算这批样本上的平均损失和梯度,然后

参数 ← 参数 − 学习率 × 梯度

学习率(η)控制每步挪多大。像学习率、批大小(每步训练用多少条数据)这种 「由人设定、训练循环自己不更新」的数,叫超参数; 它们不能用训练数据调,要用另外留出的验证集来挑9

5. 主走查:用合成数据(人工造的假数据,真参数已知、方便验证实现对不对)把整个循环跑一遍

原书在这里做了一件教学上很漂亮的事:数据是自己造的,真参数事先知道, 这样训完可以直接对照「学出来的参数」和「真参数」10

造数据:1000 条样本,每条 2 个特征(标准正态), 标签按 y = Xw + b + ε 生成,真参数 w = [2, −3.4]、b = 4.2,噪声 ε 是 σ=0.01 的高斯噪声11

然后只用张量和自动微分(不用任何框架的高层功能), 把四个零件各写一遍:

① 初始化:w 从 N(0, 0.01) 随机抽(0.01 是「实践上常用」的经验值),b = 0
② 模型:ŷ = Xw + b
③ 损失:这批样本上 ½(ŷ−y)² 的平均
④ 优化:算梯度 → w ← w − η·∂L/∂w,b 同理(每步前记得把上次梯度清零)

训练循环:重复 3 轮(epoch——每轮把全部 1000 条过一遍)
每轮:打乱顺序,按批大小逐个取小批量,做 ③④

结果:学出来的 w ≈ [2, −3.4],b ≈ 4.2,和真值非常接近12

但作者立刻给了一句重要的解毒:别把「恢复出真参数」当成理所当然。 第一,深度模型根本没有唯一解; 第二,机器学习真正要的从来不是「恢复真参数」, 而是「找到预测足够准的参数」13—— 深度网络里有大量不同的参数组合,预测能力几乎一样好, 这反而是 SGD 在实践中好用的部分原因。

另外两个工程事实:

  • 向量化:用 Python 循环逐元素加两个一万维向量, 比直接调一次 + 慢得多——向量化常带来数量级的提速14;
  • 模型、数据、训练器在原书里被组织成三个类(Module/DataModule/Trainer), 从此全书所有模型都套这个骨架——我们拆解不展开代码,记住这个三分结构即可。

6. 线性回归是单层神经网络

最后一个视角转换:线性回归就是只有一个神经元、一层的神经网络15—— 每个特征是一个输入,全部连到唯一的输出,中间没有任何隐藏层(夹在输入与输出之间、不直接对外可见的中间计算层)。 「深度学习」的所有复杂结构,都是从这张小图上长出来的。

名字里的「神经」有出处:1943 年 McCulloch 和 Pitts 模仿生物神经元(树突收信号、突触权重加权、轴突输出)提出的数学模型16。 但作者特意引了 Russell 和 Norvig 的提醒: 飞机受鸟启发,但鸟类学早已不是航空学进步的主要驱动力17—— 今天深度学习的灵感同样更多来自数学、统计和计算机科学, 把神经网络当成对大脑的工程模拟,会读错这本书。

7. 作者的判断与证据

书里给了证据的: 解析解的推导;「矩阵乘比向量串乘高效一个数量级」; 向量化实验;合成数据上真参数被恢复(书里有运行结果)。

作者明说「这是经验不是理论」的: 0.01 的初始化尺度(「often works well in practice」); 批大小 32–256 的建议(「a good start」); 深层网络「训练集上找到好参数不难,难的是泛化」—— 这是作者对实践状态的概括,理论解释在第 06 章会被他亲口标为「未解决」。

判断(我们的,不是书里的): 这一章真正该刻进脑子里的是第 3 节的逻辑: 损失函数是从「数据怎么生成的」这个假设里推出来的,不是拍脑袋选的。 平方误差 ↔ 高斯噪声;后面第 05 章会看到交叉熵 ↔ 类别分布。 以后遇到任何新损失,先问「它隐含了什么生成假设」。 如果错,会错在: 有些广泛使用的损失(如各类 hinge/对比损失) 最初是为优化性质设计的,事后才补的概率解释——「先有假设再有损失」并非总是历史事实。

8. 边界与局限

  • 线性假设本身就是边界:真实世界大多不是加权和(第 06 章专门讲这个局限);
  • 平方误差对异常值敏感(第 3 节的双刃剑);
  • minibatch SGD 不保证找到全局最优——线性回归恰好有全局唯一最优, 深网则到处是鞍点;但第 5 节说了,实践者要的不是「那个」最优;
  • 「epoch 跑几轮、学习率定多少」这一章故意没讲怎么选——那是第 04、14 章的事。

9. 可带走的

  1. 线性回归 = 特征的加权和 + 偏置 = 一层、一个神经元的神经网络;
  2. 平方误差 ↔ 「数据 = 线性关系 + 高斯噪声」的假设;损失来自假设,不是惯例;
  3. 解析解是孤例,不是常态——迭代优化(梯度下降)才是通用工具;
  4. minibatch SGD 是「全量太慢、单样本太抖」的折中;批大小 32–256 起步;
  5. 超参数(学习率、批大小)不在训练循环里更新,要用验证集挑;
  6. 训练循环五行:初始化 → 取批 → 算损失 → 反传算梯度 → 更新;先清零再反传;
  7. 目标不是恢复真参数,是预测准——深网里等价的好参数多的是;
  8. 合成数据的价值:真参数已知,可以验证实现本身没写错。

10. 原文地图

主题原书章原文位置
仿射变换、bias 的必要Linear Regressiontext/12-linear-regression.txt:119(搜「affine transformation」)
平方误差与双刃剑Linear Regressiontext/12-linear-regression.txt:204(搜「squared error」) · text/12-linear-regression.txt:225(搜「double-edge sword」)
解析解与「别习惯好运」Linear Regressiontext/12-linear-regression.txt:277(搜「analytic solutions」) · text/12-linear-regression.txt:278(搜「good fortune」)
单样本低效一个量级Linear Regressiontext/12-linear-regression.txt:316(搜「order of magnitude more efficient」)
批大小 32–256、超参数、验证集Linear Regressiontext/12-linear-regression.txt:332(搜「32 and 256」) · text/12-linear-regression.txt:363(搜「hyperparameters」) · text/12-linear-regression.txt:366(搜「validation dataset」)
深网损失面多鞍点Linear Regressiontext/12-linear-regression.txt:384(搜「saddle points」)
向量化提速Linear Regressiontext/12-linear-regression.txt:483(搜「order-of-magnitude speedups」)
高斯噪声 ↔ 最大似然Linear Regressiontext/12-linear-regression.txt:609(搜「maximum likelihood estimation」)
单层神经网络Linear Regressiontext/12-linear-regression.txt:642(搜「single-layer fully connected」)
McCulloch-Pitts、飞机与鸟Linear Regressiontext/12-linear-regression.txt:654(搜「McCulloch」) · text/12-linear-regression.txt:690(搜「ornithology」)
合成数据真参数Synthetic Regression Datatext/14-synthetic-regression-data.txt:105(搜「true parameters」)
初始化 0.01、epoch、恢复真参数Linear Regression from Scratchtext/15-linear-regression-implementation-from-scratch.txt:74(搜「magic number」) · text/15-linear-regression-implementation-from-scratch.txt:500(搜「exactly recover」)
目标是预测准不是恢复Linear Regression from Scratchtext/15-linear-regression-implementation-from-scratch.txt:512(搜「highly accurate prediction」)

Footnotes

  1. 出处:「Linear Regression」第 119 段(text/12-linear-regression.txt:119,搜「affine transformation」)。原文:bias 让我们能表达所有线性函数,而不是只有过原点的那部分。

  2. 出处:「Linear Regression」第 204 段(text/12-linear-regression.txt:204,搜「squared error」)。

  3. 出处:「Linear Regression」第 225 段(text/12-linear-regression.txt:225,搜「double-edge sword」)。

  4. 出处:「Linear Regression」第 609 段(text/12-linear-regression.txt:609,搜「maximum likelihood estimation」)。原文:「minimizing the mean squared error is equivalent to the maximum likelihood estimation of a linear model under the assumption of additive Gaussian noise」。

  5. 出处:「Linear Regression」第 245 段(text/12-linear-regression.txt:245,搜「analytically」)与第 270 段(text/12-linear-regression.txt:270,搜「invertible」)。

  6. 出处:「Linear Regression」第 278 段(text/12-linear-regression.txt:278,搜「good fortune」)。原文:「you should not get used to such good fortune」;解析解的要求「so restrictive that it would exclude almost all exciting aspects of deep learning」。

  7. 出处:「Linear Regression」第 316 段(text/12-linear-regression.txt:316,搜「order of magnitude more efficient」)。原文:处理器乘加比从主存搬数据到缓存快得多,所以一次矩阵-向量乘比一串向量-向量运算高效约一个数量级。

  8. 出处:「Linear Regression」第 332 段(text/12-linear-regression.txt:332,搜「32 and 256」)。

  9. 出处:「Linear Regression」第 363 段(text/12-linear-regression.txt:363,搜「hyperparameters」)与第 366 段(text/12-linear-regression.txt:366,搜「validation dataset」)。

  10. 出处:「Synthetic Regression Data」第 17 段(text/14-synthetic-regression-data.txt:17,搜「correct parameters are known」)。原文:合成数据让我们可以检验模型「can in fact recover them」。

  11. 出处:「Synthetic Regression Data」第 61 段(text/14-synthetic-regression-data.txt:61,搜「1000 examples」)与第 105 段(text/14-synthetic-regression-data.txt:105,搜「true parameters」)。

  12. 出处:「Linear Regression Implementation from Scratch」第 476 段(text/15-linear-regression-implementation-from-scratch.txt:476,搜「comparing the true parameters」)。

  13. 出处:「Linear Regression Implementation from Scratch」第 500 段(text/15-linear-regression-implementation-from-scratch.txt:500,搜「exactly recover」)与第 512 段(text/15-linear-regression-implementation-from-scratch.txt:512,搜「highly accurate prediction」)。

  14. 出处:「Linear Regression」第 483 段(text/12-linear-regression.txt:483,搜「order-of-magnitude speedups」)。

  15. 出处:「Linear Regression」第 642 段(text/12-linear-regression.txt:642,搜「single-layer fully connected」)。

  16. 出处:「Linear Regression」第 654 段(text/12-linear-regression.txt:654,搜「McCulloch」)。

  17. 出处:「Linear Regression」第 690 段(text/12-linear-regression.txt:690,搜「ornithology」)。原文引 Russell & Norvig:「airplanes might have been inspired by birds, but ornithology has not been the primary driver of aeronautics innovation for some centuries」。