跳到主要内容

训练的脚手架 — nn 工具箱与喂数据的流水线

这一章讲三件事: 一个神经网络由哪四件东西构成;PyTorch 的标准训练循环每一步在动什么; 数据从磁盘到 GPU 之间要过哪几道加工。 第 2 章已经能自动求导,这一章把「自动求导」装进一套可复用的工程结构里—— 后面 10 章的所有实验,跑的都是这同一套循环。

1. 顶层全景:四件套转成一个闭环

┌──────────────── 训练循环(重复 N 轮) ────────────────┐
│ │
数据 ──→ 模型(层叠层) ──→ 预测值 ──→ 损失函数 ──→ 损失值 │
标签 ────────────────────────────────↑ │ │
│ │
优化器 ←── 梯度(反向传播)
│ │
└── 更新参数 ────────┘

图说:层把输入张量变成输出张量;损失函数给「差多远」打一个数;
优化器拿着梯度改参数。四件套缺一,循环就转不起来。

书里对四件套的定义可以逐条对着上图核:「将输入张量转换为输出张量」;模型是层构成的网络;损失函数是参数学习的目标函数;优化器负责让损失变小1。这个闭环转到达阈值(预设的及格线)或轮数上限为止2

2. 主走查:MNIST 训练循环逐行走查

本章的主走查是一个完整的手写数字识别任务:两层隐藏层(300、100 个神经元)的全连接网络,输入 28×28=784 个像素,输出 10 个类别。全部超参数(开跑前手工预设的配置)在开跑前定死:训练份 64、测试份 128、步幅 0.01、冲量 0.5、20 遍3

网络三段:Linear→BatchNorm→ReLU 叠两层,再接输出层4;损失用 CrossEntropyLoss(分类任务的标准损失,第 04 章细讲它的坑),优化器用带动量的随机梯度下降(顺着「误差对每个参数的影响」小步走)5

下面把训练循环的每一行拆开看(这是全书所有实验共用的骨架):

model.train() ① 切到训练模式(Dropout/BN 行为随模式变)
for img, label in train_loader: ② 每次拿 64 张图
img = img.view(img.size(0), -1) ③ 784 个像素拉平成一行
out = model(img) ④ 前向:64×784 进,64×10 出
loss = criterion(out, label) ⑤ 一个数:这批预测离标签多远
optimizer.zero_grad() ⑥ 清梯度(梯度是累加的,见第 2 章)
loss.backward() ⑦ 反向:每个参数拿到梯度
optimizer.step() ⑧ 参数沿梯度挪一小步
_, pred = out.max(1) ⑨ 10 个分数里挑最大的当预测类别

跑 20 轮后的账本:训练准确率到 0.9996,测试准确率稳定在 0.9839 左右——训练集几乎全对、测试集 98.4%,书里如实指出中间的落差就是轻微过拟合(背题背过头),并说加卷积层和 Dropout 还能再涨6。另有一步小技巧藏在循环里:每逢轮数是 5 的倍数,学习率乘 0.1——后期小步微调7

主走查的三条读法,后面每章都会复用: ⑥⑦⑧ 三行顺序是铁律(乱序梯度就错);①③ 两行是最容易漏的模式切换与形状变化;⑤ 的损失值是判断「训练有没有在进步」的唯一仪表。

3. Module 与 functional:同一层的两副面孔

nn 工具箱里,几乎每层都有两个版本:nn.Linear(类)和 F.linear(函数)。功能相同、性能相当,区别在谁管参数:

nn.Xxx(如 nn.Linear)F.xxx(如 F.relu)
本质继承 Module 的层纯函数
参数 weight/bias自动创建、自动注册自己定义、每次手动传入
与 Sequential 搭配可以不行
Dropout 的 train/eval 切换自动完成不具备8

书里给出的选型规则很实用:卷积、全连接、Dropout 这类含可学习参数的用 nn.Xxx;激活函数、池化(在窗口内取最大/平均来缩图)这类无参数的,两可,书里习惯用 F.relu9。这条规则背后是一个工程判断:参数的「所有权」必须落在模型对象上,优化器才能找到它们——函数版参数散落在调用处,优化器无从注册。

4. 优化器与学习率:换起来像换零件

优化器统一住在 torch.optim 里,用法五步:实例化(把 model.parameters() 交给它)→ 前向 → 清梯度 → 反向 → step()10。第 3 章顺带做了一个四优化器对比实验:同一个两层网络、同一份 y=x²+噪声 数据,SGD、动量版 SGD、RMSProp、Adam 各跑一份,同轮数下自适应优化器收敛更快11——具体机制留到第 05 章拆。

学习率调整有一个容易踩的细节:新建一个优化器开销很小,但会把动量等状态清零,对带动量的优化器可能引起收敛震荡;所以训练中途调学习率,惯用做法是直接改 optimizer.param_groups[0]['lr']——主走查里那个「每 5 轮乘 0.1」就是这么写的12

5. 喂数据的流水线:Dataset、DataLoader、transforms

深度学习项目里,数据预处理往往占掉大部分时间13。PyTorch 把这段流水线拆成三级:

第一级:Dataset(定义「一条数据长什么样」)。 自定义数据集继承 Dataset,实现两个方法:__len__ 报总数,__getitem__(i) 返回第 i 条(数据+标签)。注意它一次只返回一条14

第二级:DataLoader(把一条条变成一批批)。 组成批次、打乱、多进程预取都在这层:batch_size 定批、shuffle=True 每轮打乱、num_workers 开子进程读数据、pin_memory 把数据放进对 GPU 友好的锁页内存、drop_last 丢弃不足一批的零头15

第三级:transforms(每条数据进门前的加工)。 对 PIL 图片的常见操作有裁剪、翻转、转张量——其中 ToTensor 做两件事:形状从「高×宽×通道」转成「通道×高×宽」,取值从 [0,255] 压到 [0,1];Normalize 再做减均值、除标准差(数据散得多开);多个操作用 Compose 像管道一样串起来16。另一件常被忽略的事:目录结构本身就是标签——ImageFolder 直接把「每个类一个文件夹」的组织变成整数标签的数据集17

喂数据的流水线还有个出口:tensorboardX 把训练过程写进日志(训练记录文件)、在浏览器里画出来——add_scalar 记损失曲线,add_graph 把网络结构存成图18。第 5 章的优化器对比、第 8 章的 GAN 损失监控,都靠它肉眼看。

6. 作者的判断与证据

给了证据的: 主走查的 20 轮账本(0.9996/0.9839)是实际运行输出;四优化器对比实验有完整代码与图像。层参数统计(第 6 章会给数字)显示全连接层吃掉绝大多数参数,本网的 Linear 层 16.6 万参数对卷积层 1216+5220——此账在第 06 章细算。

给了理由的选型建议: nn.Xxx 与 functional 的分工,书里引了「PyTorch 官方推荐」作依据(「第3章 PyTorch神经网络工具箱」第 375 段,text/04-ch03-3-pytorch.txt:375,搜「官方推荐」),是转述官方立场而非作者发明。

作者的经验之谈: 「新建优化器会重置动量状态导致震荡」12属于实战经验,书里没给实验数据,但与 PyTorch 文档的行为一致。

7. 边界与局限

  • 本章的 BatchNorm 只当「一层」用, 它在训练/测试两种模式下的行为差别(用批统计还是全局统计)书里只点了「training 属性」一句,机制留到第 04 章补。
  • num_workers 多进程在 Windows 上的坑、worker_init_fn 等参数, 书未覆盖;多进程读数据的调试是实战重灾区。
  • CrossEntropyLoss 内置 softmax(把输出压成总和为 1 的一组比例)这个关键细节只在第 5 章提一句。 主走查网络输出层之后没有 softmax——新手照着别的教材加一个,精度反而掉。此坑在第 04 章 §7 正面拆解。
  • 自动调学习率(调度)只有「手动乘 0.1」一种,torch.optim.lr_scheduler 的成套调度器不在书内。

8. 可带走的

  1. 神经网络=层+损失函数+优化器三选一都不许缺;训练循环的顺序 train→forward→zero_grad→backward→step 是全书通用骨架;
  2. model.train() / model.eval() 不是仪式:Dropout 和 BN 在两种模式下行为不同,忘了切换测试结果会不可复现;
  3. 含参数的层用 nn.Xxx(参数自动注册给优化器),无参数的用 F.xxx;Dropout 必须用 nn.Xxx 才能自动切模式;
  4. 梯度是累加的,zero_grad() 漏写是最常见的静默 bug;
  5. 中途调学习率改 param_groups,别重建优化器——动量状态会被清零;
  6. 数据三级流水线:Dataset 管单条、DataLoader 管批量与并行、transforms 管进门加工;
  7. ToTensor 顺手做了 HWC→CHW 和 [0,1] 归一化(把像素值缩到 0 到 1 之间);目录结构可以直接当标签(ImageFolder);
  8. 训练过程写进 tensorboardX,损失曲线是判断训练健康度的第一仪表。

9. 原文地图

主题原书章原文位置
四个核心组件第3章 PyTorch神经网络工具箱text/04-ch03-3-pytorch.txt:24(搜「将输入张量转换为输出张量」) · text/04-ch03-3-pytorch.txt:28(搜「目标函数」)
训练闭环与停止条件第3章 PyTorch神经网络工具箱text/04-ch03-3-pytorch.txt:40(搜「循环过程」)
Module 与 functional 区别第3章 PyTorch神经网络工具箱text/04-ch03-3-pytorch.txt:51(搜「纯函数」) · text/04-ch03-3-pytorch.txt:375(搜「官方推荐」)
Dropout 状态自动切换第3章 PyTorch神经网络工具箱text/04-ch03-3-pytorch.txt:373(搜「自动实现状态的转换」)
MNIST 超参数第3章 PyTorch神经网络工具箱text/04-ch03-3-pytorch.txt:104(搜「train_batch_size」) · text/04-ch03-3-pytorch.txt:109(搜「momentum」)
网络结构定义第3章 PyTorch神经网络工具箱text/04-ch03-3-pytorch.txt:167(搜「Sequential」) · text/04-ch03-3-pytorch.txt:182(搜「28 * 28」)
损失与优化器实例化第3章 PyTorch神经网络工具箱text/04-ch03-3-pytorch.txt:186(搜「CrossEntropyLoss」) · text/04-ch03-3-pytorch.txt:187(搜「momentum=momentum」)
动态修改学习率第3章 PyTorch神经网络工具箱text/04-ch03-3-pytorch.txt:205(搜「param_groups」)
训练循环三件套第3章 PyTorch神经网络工具箱text/04-ch03-3-pytorch.txt:214(搜「zero_grad」) · text/04-ch03-3-pytorch.txt:216(搜「optimizer.step」)
训练/测试结果数字第3章 PyTorch神经网络工具箱text/04-ch03-3-pytorch.txt:253(搜「0.9839」) · text/04-ch03-3-pytorch.txt:260(搜「cnn、Dropout」)
优化器使用五步第3章 PyTorch神经网络工具箱text/04-ch03-3-pytorch.txt:408(搜「清空梯度」) · text/04-ch03-3-pytorch.txt:423(搜「optimizer.step()」)
重建优化器的问题第3章 PyTorch神经网络工具箱text/04-ch03-3-pytorch.txt:426(搜「初始化动量」)
四优化器对比实验第3章 PyTorch神经网络工具箱text/04-ch03-3-pytorch.txt:498(搜「opt_SGD」) · text/04-ch03-3-pytorch.txt:501(搜「opt_Adam」)
Dataset 抽象类第4章 PyTorch数据处理工具箱text/05-ch04-4-pytorch.txt:23(搜「抽象类」)
Dataset 一次返回一个第4章 PyTorch数据处理工具箱text/05-ch04-4-pytorch.txt:93(搜「每次只返回一个样本」)
DataLoader 主要参数第4章 PyTorch数据处理工具箱text/05-ch04-4-pytorch.txt:121(搜「num_workers」) · text/05-ch04-4-pytorch.txt:128(搜「drop_last」)
预处理占大部分时间第4章 PyTorch数据处理工具箱text/05-ch04-4-pytorch.txt:4(搜「大部分时间」)
ToTensor 的两件事第4章 PyTorch数据处理工具箱text/05-ch04-4-pytorch.txt:180(搜「[0,255]」)
Normalize 与 Compose第4章 PyTorch数据处理工具箱text/05-ch04-4-pytorch.txt:191(搜「减均值」) · text/05-ch04-4-pytorch.txt:195(搜「管道」)
ImageFolder 目录当标签第4章 PyTorch数据处理工具箱text/05-ch04-4-pytorch.txt:234(搜「自动转化成序列」)
tensorboardX 与 add_scalar第4章 PyTorch数据处理工具箱text/05-ch04-4-pytorch.txt:289(搜「支持scalar」) · text/05-ch04-4-pytorch.txt:426(搜「训练损失值」)

Footnotes

  1. 出处:「第3章 PyTorch神经网络工具箱」第 24 段(text/04-ch03-3-pytorch.txt:24,搜「将输入张量转换为输出张量」)与第 28 段(text/04-ch03-3-pytorch.txt:28,搜「目标函数」)。

  2. 出处:「第3章 PyTorch神经网络工具箱」第 40 段(text/04-ch03-3-pytorch.txt:40,搜「循环过程」)。

  3. 出处:「第3章 PyTorch神经网络工具箱」第 104 段(text/04-ch03-3-pytorch.txt:104,搜「train_batch_size」)与第 109 段(text/04-ch03-3-pytorch.txt:109,搜「momentum」)。

  4. 出处:「第3章 PyTorch神经网络工具箱」第 167 段(text/04-ch03-3-pytorch.txt:167,搜「Sequential」)。隐藏层宽度 300/100 见第 182 段(text/04-ch03-3-pytorch.txt:182,搜「28 * 28」)。

  5. 出处:「第3章 PyTorch神经网络工具箱」第 186 段(text/04-ch03-3-pytorch.txt:186,搜「CrossEntropyLoss」)与第 187 段(text/04-ch03-3-pytorch.txt:187,搜「momentum=momentum」)。

  6. 出处:「第3章 PyTorch神经网络工具箱」第 253 段(text/04-ch03-3-pytorch.txt:253,搜「0.9839」)与第 260 段(text/04-ch03-3-pytorch.txt:260,搜「cnn、Dropout」)。

  7. 出处:「第3章 PyTorch神经网络工具箱」第 205 段(text/04-ch03-3-pytorch.txt:205,搜「param_groups」)。

  8. 出处:「第3章 PyTorch神经网络工具箱」第 373 段(text/04-ch03-3-pytorch.txt:373,搜「自动实现状态的转换」)。

  9. 出处:「第3章 PyTorch神经网络工具箱」第 375 段(text/04-ch03-3-pytorch.txt:375,搜「官方推荐」)与第 48 段(text/04-ch03-3-pytorch.txt:48,搜「nn.functional」)。

  10. 出处:「第3章 PyTorch神经网络工具箱」第 408 段(text/04-ch03-3-pytorch.txt:408,搜「清空梯度」)与第 423 段(text/04-ch03-3-pytorch.txt:423,搜「optimizer.step()」)。

  11. 出处:「第3章 PyTorch神经网络工具箱」第 498 段(text/04-ch03-3-pytorch.txt:498,搜「opt_SGD」)与第 501 段(text/04-ch03-3-pytorch.txt:501,搜「opt_Adam」)。

  12. 出处:「第3章 PyTorch神经网络工具箱」第 426 段(text/04-ch03-3-pytorch.txt:426,搜「初始化动量」)。 2

  13. 出处:「第4章 PyTorch数据处理工具箱」第 4 段(text/05-ch04-4-pytorch.txt:4,搜「大部分时间」)。

  14. 出处:「第4章 PyTorch数据处理工具箱」第 93 段(text/05-ch04-4-pytorch.txt:93,搜「每次只返回一个样本」)。

  15. 出处:「第4章 PyTorch数据处理工具箱」第 121 段(text/05-ch04-4-pytorch.txt:121,搜「num_workers」)与第 128 段(text/05-ch04-4-pytorch.txt:128,搜「drop_last」)。

  16. 出处:「第4章 PyTorch数据处理工具箱」第 180 段(text/05-ch04-4-pytorch.txt:180,搜「[0,255]」)与第 195 段(text/05-ch04-4-pytorch.txt:195,搜「管道」)。

  17. 出处:「第4章 PyTorch数据处理工具箱」第 234 段(text/05-ch04-4-pytorch.txt:234,搜「自动转化成序列」)。

  18. 出处:「第4章 PyTorch数据处理工具箱」第 289 段(text/05-ch04-4-pytorch.txt:289,搜「支持scalar」)与第 426 段(text/05-ch04-4-pytorch.txt:426,搜「训练损失值」)。