跳到主要内容

最小训练循环 — 「训练」这两个字具体是哪五个动作

这一章讲三件事: 训练循环长什么样、为什么这个最小版本不能直接拿去跑真数据、 以及训练完之后怎么把它存下来再原样取回来。

它在全书链条里的位置: 第 01 章给了容器,第 02 章给了「往哪边改」。 这一章把它们接成一个能转起来的圈——后面二十章换的只是圈里那台机器和喂进去的数据。

需要的基础: 第 02 章的梯度清零。那一行就出现在这个循环里。

1. 先看现象:模型不会自己变好

你搭好了一个网络,喂进去一条数据,它吐出一个数。这个数是错的。 然后呢? 网络不会因为「答错了」就自动变好——必须有人把「错了多少」翻译成 「每个数该往哪边挪」,再真的去挪。

书里把这件事拆成四个准备动作加一个循环:用现成的层定义模型;准备好数据; 指定那把尺子和那个执行者;然后写循环——前向、算损失、反向、更新1

这里有两个词先各交代一句。 损失函数就是那把尺子: 量「答得离正确答案差多远」,输出一个越小越好的数。

优化器就是那个执行者:拿着每个参数的梯度,决定这一步实际挪多少。

全章主走查:256 条二维样本,拟合 y = x₁ + x₂ + 噪声
─────────────────────────────────────────────────────────────
准备 X 是 256 行 2 列的随机数;y 是每行两个数之和,再加一点噪声
模型 TinyMLP:2 → 16 → 1,中间夹一道非线性
损失取均方误差,优化器取 Adam,每步挪的幅度设为 0.05
循环 跑 200 轮,每 50 轮打印一次
─────────────────────────────────────────────────────────────
§3 把一轮里的五步逐个走一遍,每步旁边写出这一步手上拿着什么形状的东西

这个例子的规模值得先记一下:256 条数据、一个只有两层的网络。 后面第 04 章那个加州房价有 20 640 条,第 08 章的图像数据集有 6 万张。 规模一变,这个骨架就得改——改在哪儿,下一节说。

2. 先记住:这个骨架是把 256 条一口全塞的

结论先行:书里这个最小循环,每一轮都把全部 256 条数据一次性喂进去。

看一眼循环体就清楚:里面没有任何「取一小撮」的动作, model(X) 里那个 X 就是完整的 256 行2

在这个玩具上这么干完全没问题——256 行 2 列一共 512 个数, 放哪儿都装得下,一次算完还更快。

但真数据一大就撑不住,原因有两条:

  • 装不下。 6 万张 28×28 的图一次全搬进显卡内存,再加上前向过程中每一层的中间结果, 内存直接爆掉;
  • 太慢。 每挪一步都要把整个数据集过一遍,一轮下来只挪了一次。

所以工程上的标准做法是拆开:每一步只随机取出一小撮样本来算梯度。 这一小撮就叫一(第 01 章讲广播时提过这个词的另一层意思, 那里指的是张量里「一次算多少份」那一维,和这里是同一个概念); 一批里放多少条,叫批大小。

这件事怎么做、批大小取多少合适、为什么这么做反而有额外的好处,第 07 章讲。 在那之前,你读到的所有「训练循环」都是全量版本——这不是简化,是书里的原样。

3. 主走查:一轮里的五个动作

这一节把一轮循环逐步拆开,每一步旁边写出手上拿着什么。

一轮循环(五步)
═══════════════════════════════════════════════════════════════
① 前向 pred = model(X)
X 的形状 [256, 2] ──→ 过第一层 ──→ [256, 16]
──→ 过非线性 ──→ [256, 16]
──→ 过第二层 ──→ [256, 1]
手上:256 个预测值

② 算损失 loss = loss_fn(pred, y)
把 256 个预测和 256 个真值逐个相减、平方、取平均
手上:一个数(标量)

③ 清零 opt.zero_grad()
把上一轮留在每个参数上的梯度抹掉
手上:所有参数的梯度格子都变回 0

④ 反向 loss.backward()
从那一个数出发,按第 02 章的账本倒着回放
手上:每个参数各自的梯度,形状与参数本身相同

⑤ 更新 opt.step()
优化器按梯度把每个参数挪一小步
⚠ 这一步用的是全部 256 条一起算出来的那一个梯度
═══════════════════════════════════════════════════════════════

第 ③ 步为什么排在第 ④ 步前面,第 02 章已经答过:梯度是累加不是覆盖3顺序写反了不会报错,只会让参数越挪越偏。

第 ② 步那个「一个数」值得多看一眼。 反向传播只能从一个孤零零的数出发 (这种不带任何形状的数就叫标量)——你不能对 256 个数同时求「往哪边改」。所以损失函数的最后一步永远是把一批数压成一个数, 通常是取平均。这条约定在后面每一章都成立。

书里对这套骨架的总结是:前向—损失—反向—更新这一固定节拍, 几乎所有基于梯度的训练任务都共享4

4. 换个任务,这五步几乎原样保留

结论先行:后面二十章里变的是模型和数据,不是循环。

书里明写:这套四步骨架在切换任务时只需局部替换—— 模型部分换成卷积、循环或者 Transformer,尺子换成另一把,数据切到图像或文本, 循环主体几乎原样保留5

能换的 不能换的
───────────────────────── ─────────────────────────
model = 什么结构都行 ① 前向
loss_fn = 什么尺子都行 ② 算成一个数
X, y = 什么数据都行 ③ 清零
opt = 什么走法都行 ④ 反向
⑤ 更新

这就是为什么这本书能用一套代码讲完十种模型。 它把不变的那部分抽出来打包了——打包成什么,第 7 节说。

5. 存模型:存参数字典,别存整个对象

结论先行:训练完要落盘,而落盘有两种存法,书里明确推荐其中一种。

书里的说法是:框架推荐保存参数字典而非整个模型对象—— 前者只记录参数张量,与代码结构解耦;后者依赖类定义和模块路径, 一旦改名或重构就难以读回来6

两种存法的区别
────────────────────────────────────────────────────────
存整个对象 文件里记着「这是 TinyMLP 类的一个实例」
→ 你把 TinyMLP 改名成 SmallMLP,文件就读不回来了

存参数字典 文件里只有一张表:参数路径 → 对应的一摞数
例如 net.0.weight → 形状 [16, 2] 的那些数
→ 只要结构不变,改名、搬文件都不影响
────────────────────────────────────────────────────────

这张表的键长什么样,书里给了例子:net.0.weight7—— 点号左边是层在模型里的路径,右边是这一层的哪一部分。 读别人的模型文件时,这串路径就是你判断结构的唯一线索。

还有一个安全选项要一起记: 读取时传一个「只要参数」的开关, 它会禁止把文件里任意的东西还原成 Python 对象, 避免加载来源不明的参数文件时执行到恶意代码8

6. 加载完必须切到推断模式,否则输出对不上

这是这一章第二个「不报错但结果不对」的坑。

书里把它单独点了出来:加载完一定要调用一次切换, 在推断模式下,随机关闭神经元的那个零件会停止随机置零, 而按批统计的那个零件会改用训练时累积的统计量; 漏掉这一步,同一份输入在推断阶段会因为内部仍存在随机或可变的行为而给出不一致的输出9

这两个零件是什么,后面才会讲——随机关闭那一个在第 14 章, 按批统计那一个在第 13 章。你现在只需要记住这条因果: 它们在训练时和推断时的行为不一样,而切换靠的就是这一个调用。

同一份输入,两种模式下的差别
────────────────────────────────────────────
训练模式 随机关掉一部分神经元 → 每次跑结果都不同
按当前这一批算统计量 → 结果受同批其他样本影响

推断模式 全部神经元都开着 → 每次跑结果一致
用训练时攒下的统计量 → 结果只取决于这一条输入
────────────────────────────────────────────

书里的验证方式很干脆: 存下来、新建一个同结构的模型、读回参数、 切到推断模式,然后比较两份模型在同样输入下的输出——差值应该是 010它还专门把「去掉那一行会怎样」留成了动手练习11

训练里数「跑了第几遍」的那个单位就叫轮——一轮就是把训练数据完整过一遍。 想留完整快照的话,把优化器状态和当前轮数一起打包成一个字典存12;断点续训靠的就是这个。

7. Runner:把不变的那部分收进一个类

结论先行:第 4 节说循环主体不变,书里就把它抽出来了。

一次完整实验要贯穿数据准备、模型搭建、损失与优化器选择、量成绩的尺子怎么定、 训练、评估、预测、模型保存与加载。这些环节每章都重复出现, 唯一不同的只是模型结构与数据集13

书里因此维护了一个叫 Runner 的类,一共三版,一版比一版能干14:

版本什么时候出场新增了什么
第一版第 04 章面向有现成公式能直接解出来的模型,没有训练循环
第二版第 05 章引入一步步逼近的训练,记录训练与验证两条曲线,保留验证集上最好的那份
第三版第 07 章改成一小撮一小撮地喂,量成绩的尺子可以任意换,按框架约定存取参数

第三版能覆盖书中绝大多数任务,后续章节都以它为起点14这份拆解里凡是提到「跑了多少轮、保留了哪一份」,背后都是它。

判断(我们的,不是书里的):这三个版本的编号会成为读者的负担,而它的收益是真的。 好处很实在:每一章的正文只剩下模型定义那几行,读者的注意力不会被样板代码稀释。 代价是:你在第 12 章看到一个 Runner,得回头想它是哪一版、有没有那个能力。 如果错,会错在: 如果读者是照着配套代码仓库一路跑下来的, 版本差异由导入语句自动解决,根本不构成负担。 判据是:只读书不跑代码的人,能不能一口说出第三版比第二版多了什么。

8. 边界与局限

这一章没覆盖的:

没讲什么依据 / 为什么值得知道
怎么一小撮一小撮地喂书里到第 4.5.1 节才讲,这一章的循环是全量的(见第 2 节)
训练中途怎么调整每步挪多远第 12 章才出现
多台机器一起训全书没有
训练过程怎么记录、实验怎么归档书里只有 print
早停第 05 章才随第二版 Runner 出现

出门会撞见的名字:

  • state_dict 就是第 5 节那张「参数路径 → 一摞数」的表7;

  • load_state_dict 是把这张表填回模型的动作;

  • model.eval() / model.train() 是第 6 节那两种模式的切换开关9;

  • optimizer.zero_grad() 是第 3 步那个清零;

  • 检查点(checkpoint) 指的是第 6 节末尾那个「模型 + 优化器 + 轮数」的完整快照12;

  • 评价指标(metric) 就是量成绩的那把尺子,它和损失函数常常不是同一把——为什么不是,第 04 章讲。

补充(不在书里,来自通用知识): 「训练模式」和「推断模式」这对说法, 在别的框架里可能写成 training / evaluation,或者 fit / predict。 换个名字,管的还是同一件事:那两个零件的行为要不要切换。

9. 可带走的

  1. 一轮训练 = 前向 → 算成一个数 → 清零 → 反向 → 更新,五步一个都不能少;
  2. 清零必须排在反向前面,顺序反了不报错、只会越挪越偏;
  3. 损失函数的最后一步永远是把一批数压成一个数,因为反向只能从标量出发;
  4. 书里这个最小循环是把 256 条一口全塞的,真数据大了会装不下也太慢;
  5. 一小撮叫一批,一批里放多少条叫批大小;具体怎么做第 07 章讲;
  6. 换任务只换模型、损失、数据、优化器,循环主体原样保留;
  7. 存参数字典,不存整个对象,前者与类名和文件路径解耦;
  8. 读取时开「只要参数」那个开关,防止来源不明的文件执行到恶意代码;
  9. 加载完必须切推断模式,否则同一份输入会给出不一致的输出;
  10. 要断点续训就把优化器状态和轮数一起打包存下来。

10. 原文地图

主题原书章原文位置
训练循环的四步拆分第1章 实践基础text/02-ch01.txt:1196(搜「一个标准训练循环可以拆成」)
最小例子的数据与超参第1章 实践基础text/02-ch01.txt:1216(搜「manual_seed」) · text/02-ch01.txt:1223(搜「for epoch in range」)
固定节拍与换任务第1章 实践基础text/02-ch01.txt:1232(搜「前向—损失—反向—更新」) · text/02-ch01.txt:1240(搜「切换任务时只需局部替换」)
忘记清零的后果第1章 实践基础text/02-ch01.txt:1236(搜「如果忘记清零」)
存参数字典而非整个模型第1章 实践基础text/02-ch01.txt:1263(搜「推荐保存」) · text/02-ch01.txt:1266(搜「键为参数路径」)
只要权重那个开关第1章 实践基础text/02-ch01.txt:1272(搜「安全选项」)
推断模式那一行第1章 实践基础text/02-ch01.txt:1296(搜「切到推断模式」) · text/02-ch01.txt:987(搜「故意」)
完整快照第1章 实践基础text/02-ch01.txt:1307(搜「断点续训」)
Runner 的三个版本第1章 实践基础text/02-ch01.txt:1246(搜「一次完整的实验流程」) · text/02-ch01.txt:1253(搜「面向有解析解的简单模型」)

Footnotes

  1. 出处:「第1章 实践基础」第 1196 段(text/02-ch01.txt:1196,搜「一个标准训练循环可以拆成」)。原文:用 nn.Module 定义模型;准备好数据;指定损失函数与优化器;写循环——前向、算损失、反向、更新。

  2. 出处:「第1章 实践基础」第 1223 段(text/02-ch01.txt:1223,搜「for epoch in range」)。循环体内 pred = model(X) 里的 X 就是第 1218 段(text/02-ch01.txt:1218,搜「torch.randn(256, 2)」)生成的完整 256 行,循环里没有任何取子集的动作。

  3. 出处:「第1章 实践基础」第 1236 段(text/02-ch01.txt:1236,搜「如果忘记清零」)。这是书里紧跟在最小训练循环后面的一条笔记。

  4. 出处:「第1章 实践基础」第 1232 段(text/02-ch01.txt:1232,搜「前向—损失—反向—更新」)。原话:循环的核心是前向—损失—反向—更新这一固定节拍,几乎所有基于梯度的训练任务都共享这套骨架。

  5. 出处:「第1章 实践基础」第 1240 段(text/02-ch01.txt:1240,搜「切换任务时只需局部替换」)。原文举的替换项是:模型部分换成 CNN/RNN/Transformer、损失改成交叉熵、数据切到图像或文本。

  6. 出处:「第1章 实践基础」第 1263 段(text/02-ch01.txt:1263,搜「推荐保存」)。原话:框架推荐保存 state_dict 而非整个 model 对象;前者只记录参数张量,与代码结构解耦,后者依赖类定义和模块路径,一旦改名或重构就难以反序列化。

  7. 出处:「第1章 实践基础」第 1266 段(text/02-ch01.txt:1266,搜「键为参数路径」)。原文给的例子正是 net.0.weight 2

  8. 出处:「第1章 实践基础」第 1272 段(text/02-ch01.txt:1272,搜「安全选项」)。原话:weights_only=True 是 PyTorch 2.4 以来推荐的安全选项,会禁止反序列化任意 Python 对象,避免加载来源不明的权重时执行恶意代码。

  9. 出处:「第1章 实践基础」第 1296 段(text/02-ch01.txt:1296,搜「切到推断模式」)。原话:在推断模式下 Dropout 停止随机置零、BatchNorm 改用训练时累积的统计量;漏掉这一步,同一份输入在推断阶段会因为内部仍存在随机/可变行为而给出不一致的输出。 2

  10. 出处:「第1章 实践基础」第 1288 段(text/02-ch01.txt:1288,搜「两份模型在同样输入下输出应当一致」)。书里的做法是打印两份模型输出的最大绝对差。

  11. 出处:「第1章 实践基础」第 987 段(text/02-ch01.txt:987,搜「故意」)。这是书里的动手练习 1.6,第一问就是把切换那一行去掉、重跑比较代码,观察差值是否仍为 0。

  12. 出处:「第1章 实践基础」第 1307 段(text/02-ch01.txt:1307,搜「断点续训」)。原文把轮数、模型参数、优化器状态、当前损失打包成一个字典统一保存。 2

  13. 出处:「第1章 实践基础」第 1246 段(text/02-ch01.txt:1246,搜「一次完整的实验流程」)。原文列出的环节包括数据准备、模型搭建、损失与优化器选择、评价指标定义、训练、评估、预测、模型保存与加载。

  14. 出处:「第1章 实践基础」第 1253 段(text/02-ch01.txt:1253,搜「面向有解析解的简单模型」)。原文逐条说明三个版本各自新增了什么,并写明第三版已能覆盖书中绝大多数任务,后续章节都以它为起点。 2