跳到主要内容

训练循环与 99.7% 的假象:第一仗打败了,而报表说赢了

这一章讲三件事: 一个真实项目的训练代码长什么样(不再是 notebook 里的十几行);分类模型怎么搭、怎么记账;以及全书最值钱的一次失败——99.7% 的正确率,实际战绩是阳性样本一个都没答对。 位置:三步流水线的第 ③ 步(分类)第一次跑起来。教训直接逼出第 14 章的精确率/召回率。

1. 从 notebook 到命令行应用

之前的训练循环都写在 notebook 里,十几行。肺癌项目的复杂度(事情的规模与牵扯面)上了台阶,书里把它重构成一个带命令行参数的 Python 应用 LunaTrainingApp:启动时解析参数(--epochs=20 --balanced 这类),再进主流程1

书里有句方法论值得单独记住:项目的基础设施要跟得上项目的复杂度——但「等基础设施完美了再动手」是一种拖延症1。脚手架够用就好,这章的脚手架就是:一个解析参数的入口、一个模型、两个数据加载器、一个训练循环、一套指标。

三个工程细节,每条都有坑:

  • 多 GPU 透明包装:单机多卡时用 nn.DataParallel(model) 把模型包一层,数据自动切到各卡再合并结果,训练代码一行不改。跨机器的 DistributedDataParallel 留到第 16 章2
  • 优化器的创建顺序:model.to(device) 之后才能建优化器——优化器握着的是参数张量的引用,先建优化器再搬家,它更新的就是 CPU 上的旧副本。这条第 08 章上 GPU 那一节踩过,这里重申2
  • 默认超参数:SGD、lr=0.001momentum=0.99。动量(momentum)是给梯度下降加的「惯性」:每次更新不光看当前梯度,还带上一步更新的九成方向,小坑小洼直接滚过去。书里的说法是这组值「是安全的起跑点」,不是调出来的最优3

数据加载侧两个提速旋钮:num_workers 开后台进程提前取下一批(GPU 算当前批时,CPU 在备下一批);pin_memory 把数据锁在「直通内存」里,CPU 到 GPU 的拷贝少走一道。开起来的效果:训练 495958 个样本、验证 55107 个,一轮约 1900 多个批次4

2. 模型:tail、backbone、head

LunaModel 是第 08 章那个 2D 卷积网的 3D 版,输入是第 12 章切好的 32×48×48 小方块。书里给分类网络的通用三段式命了名5:

输入 [1, 32, 48, 48]

tail BatchNorm3d ← 先把输入归一化到骨干适应的范围

backbone 4 × LunaBlock ← 每块:Conv3d→ReLU→Conv3d→ReLU→MaxPool3d
│ 每块后分辨率减半:32×48×48 → … → 2×3×3

head Linear(1152, 2) ← 拉直成 1152 个数,判两类
Softmax

图说:数据从「尾」流向「头」(和成语「从头到尾」相反——书里特地吐槽了这点)。

为什么每层卷积用两个 3×3×3 叠着,而不是一个 5×5×5?感受野(receptive field,第 08 章:一个输出单元能「看见」的输入范围)账:两个 3×3×3 叠起来,最里层那个输出体素受 5×5×5=125 个输入体素影响,和单个 5×5×5 一样;但参数量是 2×27=54 对 125(不计通道),同样的视野,一半不到的参数——还多赚一次中间的非线性6

模型还有两个不显眼但讲究的设计:

  • forward 返回两样东西:未归一化的原始分数(logits,第 07 章讲过的概念)和 softmax 后的概率。训练时用 logits 配 CrossEntropyLoss(一步算完,数值上更稳);给人看、做指标时用概率。训练时的输出和展示时的输出不同,是正常的7
  • 初始化自己动手:按第 09 章讲的 Kaiming 初始化给每个卷积层赋初值(扇出模式、配 ReLU),而不是用 PyTorch 的默认8

3. 指标记账:每样本损失与三行账

这章最重要的工程决策不在模型里,在记账方式里。

computeBatchLoss 建损失函数时传了 reduction='none':不让 CrossEntropyLoss 把一批 32 个样本的损失压成一个平均数,而是保留每个样本各自的损失,再自己求平均。为什么多此一举?因为要按类别分开算账——平均损失一个数,看不出模型对哪一类在学9

配套的是一块三行的指标数组 metrics_g,每列对应一个样本,三行分别是:真实标签、模型预测(概率 >0.5 判阳性)、该样本的损失。一个 epoch 结束后,拿 0.5 做阈值切出「阴性样本掩码」「阳性样本掩码」,就能分别报:全体/阴性/阳性三组各自的损失和正确率9

验证循环则是第 05 章的标准动作收进工程里:model.eval()(关掉 dropout、让批归一化用运行统计)加 torch.no_grad()(不建计算图,快一截),只记账,不更新10

这套记账是后面「大揭露」的全部前提。 只报总正确率,这章的事故根本不会被发现。

4. 主走查:第一个 epoch 的日志,逐行读

训练跑起来了。走查就是读第一个 epoch(E1)打出来的日志——每个数字都是书里真实的运行输出11:

E1 trn 2.4576 loss, 99.7% correct
E1 trn_neg 0.1936 loss, 99.9% correct (494289 of 494743)
E1 trn_pos 924.34 loss, 0.2% correct (3 of 1215)
E1 val 0.0172 loss, 99.8% correct
E1 val_pos ..., 0.0% correct (0 of 136)

图说:第一行是喜报,下面四行是验尸报告。

逐行拆:

  1. 总账 99.7%:494743 个阴性 + 1215 个阳性,全答「阴性」也有 494743/495958 ≈ 99.75%。这个分数不训练也拿得到
  2. trn_neg 99.9%:阴性样本里几乎全对——模型确实在学,学的是「反正说不是」。
  3. trn_pos 0.2%(3/1215):训练集 1215 个真结节,只认出 3 个。924 的平均损失说明它答「不是」时还非常自信。
  4. val_pos 0.0%(0/136):验证集 136 个真结节,一个没认出来。

书里给这个失败模式配了一个考试比喻:100 道判断题,99 道答案是「错」,全答「错」得 99 分;另一个学生答对了唯一那道「对」,反而得分更低——但谁都知道谁更懂12。放到癌症筛查上更直白:一个永远说「没事」的模型,漏掉的就是肿瘤。书里原话:这种失败模式在真实世界是最危险的一种11

为什么会塌成这样?先记住现象,机制留给下一章:495958 个样本里阳性只有 1215 个(约 400:1),梯度几乎全被阴性样本攥着,模型最快降低损失的方式就是对所有输入输出「阴性」——它找到了一个我们没打算让它找的解。

顺便:两个看得见进度的工具

训练中怎么知道还要多久?书里用了 tqdm(一个进度条库,把数据加载器包一层就显示「16/1938,预计还要 X」),作者们的梗是用预估时间决定「来不来得及去巴黎吃晚饭」13

训练完怎么复盘?TensorBoard(Google 出的训练可视化面板,PyTorch 用 SummaryWriter 往里写):每行指标 add_scalar(标签, 值, global_step) 记一个点,浏览器里看趋势线。两个细节:横轴用累计样本数而不是 epoch 号——epoch 长度一变,两条曲线就没法对齐;斜杠分组标签(loss/trainloss/val)让曲线自动归堆14

5. 作者的判断与证据

说法性质
SGD lr=0.001 momentum=0.99 是安全起点作者经验,明说「不是最优,是起跑点」3
两个 3×3×3 比单个 5×5×5 省参数有据,感受野与参数账可复算6
训练用 logits、展示用概率有据(数值稳定性),第 07 章已论证7
99.7% 是假象、该模式最危险有据,日志数字在书里;「最危险」是作者的强调,但医学语境下成立11
「基础设施拖延症是陷阱」作者经验论,无证据,但便宜1
横轴用样本数不用 epoch工程判断,理由是可比性;无人反对过14

6. 边界与局限

  • DataParallel 是单机多卡的省事方案,书上自己也说它效率不如分布式方案——真要快,看第 16 章。
  • 指标阈值取 0.5 是默认,不是调出来的;下一章会看到阈值本身就是个旋钮。
  • tqdm 的进度估计在批次耗时不均时会跳变,书里没提——别拿它当精确倒计时。
  • 这章的「失败」其实是设计好的教学时刻:作者们早知道会这样(下一章开头就揭)。真实项目里,这种时刻没人给你预告。

7. 可带走的

  1. 结构跟着复杂度走:脚本变大就升级成带命令行参数的应用;但别等脚手架完美才开工。
  2. 优化器在 .to(device) 之后建;多卡先 DataParallel 顶着。
  3. 模型分 tail/backbone/head:尾管适配输入,骨干管提取,头管输出格式。
  4. 同视野叠小核比单大核省参数——两个 3×3×3(54 参数)顶一个 5×5×5(125 参数)。
  5. 损失留每样本一份(reduction='none'),按类别分开记账;平均数会藏起少数类的灾难。
  6. 类别不平衡时,总正确率是说谎的指标——99.7% 可以等于「一个都没认出来」。报指标必须分阳/阴两列。
  7. 训练用 logits,展示用概率,一个 forward 两个返回值很正常。
  8. 曲线横轴用累计样本数,标签用斜杠分组。

8. 原文地图

主题原书章原文位置
CLI 应用与基础设施观ch13text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:133(搜「infrastructure」) · :192(搜「LunaTrainingApp」)
DataParallel 与 DDPch13text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:293(搜「nn.DataParallel(model)」) · :311(搜「DistributedDataParallel」)
SGD 默认值与动量ch13text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:299(搜「momentum=0.99」) · :323(搜「safe place to start」)
num_workers/pin_memory 与样本数ch13text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:398(搜「num_workers」) · :416(搜「495958 training samples」)
tail/backbone/headch13text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:486(搜「tail, a backbone」)
感受野 5×5×5 省参数ch13text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:575(搜「effective receptive field of 5 × 5 × 5」)
logits 与概率双输出ch13text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:681(搜「raw logits」)
Kaiming 初始化ch13text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:710(搜「kaiming_normal_」)
每样本损失与 metrics_gch13text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:859(搜「reduction='none'」) · :894(搜「METRICS_LABEL_NDX」)
验证循环 eval/no_gradch13text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:982(搜「torch.no_grad()」)
99.7% 大揭露ch13text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:1337(搜「99.7% correct」) · :1349(搜「3 of 1215」)
最危险失败模式ch13text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:1374(搜「most dangerous in the real world」)
考试比喻ch13text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:1581(搜「100 True/False」)
tqdm 与巴黎晚饭ch13text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:1292(搜「dinner in Paris」)
TensorBoard 与 global_stepch13text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:1504(搜「SummaryWriter」) · :1568(搜「epoch number」)

Footnotes

  1. 出处:「13 Training a classification model to detect suspected tumors」第 133 段(text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:133,搜「infrastructure」)、第 187 段(:187,搜「argparse」)与第 192 段(:192,搜「LunaTrainingApp」)。章首方法论:「应用的规模要和项目匹配;陷入『先把基础设施做完美』是一种拖延症」;参数在 __init__ 里用标准库 argparse 解析。 2 3

  2. 出处:「13 Training a classification model to detect suspected tumors」第 293 段(text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:293,搜「nn.DataParallel(model)」)与第 311 段(:311,搜「DistributedDataParallel」)。多 GPU 时用 DataParallel 透明包装;PyTorch 另提供跨机器的 DistributedDataParallel(第 16 章讲)。 2

  3. 出处:「13 Training a classification model to detect suspected tumors」第 299 段(text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:299,搜「momentum=0.99」)与第 323 段(:323,搜「safe place to start」)。「用带动量的 SGD 通常被认为是安全的起跑点」;学习率、动量、网络规模这类配置统称超参数(hyperparameter),调它们的过程叫超参数搜索。 2

  4. 出处:「13 Training a classification model to detect suspected tumors」第 398 段(text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:398,搜「num_workers」)、第 399 段(:399,搜「pin_memory」)与第 416 段(:416,搜「495958 training samples」)。num_workers 起后台进程预取;pin_memory 加速 CPU→GPU 传输;训练集 495958、验证集 55107 个样本,batch_size=256 时一轮 1938 个批次(:1198,搜「E1 Training」)。

  5. 出处:「13 Training a classification model to detect suspected tumors」第 486 段(text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:486,搜「tail, a backbone」)与第 733 段(:733,搜「BatchNorm3d」)。分类模型常由 tail/backbone/head 组成,数据从尾流向头(和「head to tail」的说法相反);LunaModel 的 tail 是一个 BatchNorm3d,backbone 是四个 LunaBlock,head 是 Linear(1152, 2) 加 softmax。

  6. 出处:「13 Training a classification model to detect suspected tumors」第 567 段(text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:567,搜「receptive field of 3 × 3 × 3」)与第 575 段(:575,搜「effective receptive field of 5 × 5 × 5」)。单个 3×3×3 卷积:27 个体素进、1 个出;两个叠起来等效感受野 5×5×5,而「叠两层 3×3×3 比单层 5×5×5 参数更少」。 2

  7. 出处:「13 Training a classification model to detect suspected tumors」第 681 段(text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:681,搜「raw logits」)与第 685 段(:685,搜「CrossEntropyLoss」)。forward 同时返回原始 logits 和 softmax 概率;训练时用 logits 算 CrossEntropyLoss,展示与指标用概率。 2

  8. 出处:「13 Training a classification model to detect suspected tumors」第 710 段(text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:710,搜「kaiming_normal_」)。_init_weightskaiming_normal_(a=0, mode='fan_out', nonlinearity='relu') 初始化卷积权重,偏置按 1/sqrt(fan_out) 的正态分布给。

  9. 出处:「13 Training a classification model to detect suspected tumors」第 859 段(text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:859,搜「reduction='none'」)与第 894 段(:894,搜「METRICS_LABEL_NDX」)。CrossEntropyLoss(reduction='none') 保留每样本损失再求均值;metrics_g 数组按列存每个样本的标签/预测/损失,detach 防止梯度泄漏。 2

  10. 出处:「13 Training a classification model to detect suspected tumors」第 982 段(text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:982,搜「torch.no_grad()」)与第 984 段(:984,搜「model.eval()」)。验证循环用 eval() 模式加 no_grad 上下文,只读不写,也不需要返回值。

  11. 出处:「13 Training a classification model to detect suspected tumors」第 1330 段(text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:1330,搜「Getting 99.7% correct means」)、第 1337 段(:1337,搜「99.7% correct」)、第 1349 段(:1349,搜「3 of 1215」)与第 1374 段(:1374,搜「most dangerous in the real world」)。E1:trn 99.7%、trn_pos 0.2%(3/1215)、val 99.8%、val_pos 0.0%;「这个特定的失败模式在真实世界是最危险的」——漏诊。 2 3

  12. 出处:「13 Training a classification model to detect suspected tumors」第 1581 段(text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:1581,搜「100 True/False」)。考试比喻:100 道判断题 99 道答案为假,全答「假」得 99 分;答对那道「真」的学生其实懂得更多。

  13. 出处:「13 Training a classification model to detect suspected tumors」第 827 段(text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:827,搜「tqdm」)与第 1292 段(:1292,搜「dinner in Paris」)。用 tqdm 包数据加载器显示进度与预估完成时间;作者的玩笑是按预估时间决定要不要「去巴黎吃晚饭」。

  14. 出处:「13 Training a classification model to detect suspected tumors」第 1504 段(text/21-ch13-13-training-a-classification-model-to-detect-sus.txt:1504,搜「SummaryWriter」)、第 1550 段(:1550,搜「add_scalar」)与第 1568 段(:1568,搜「epoch number」)。torch.utils.tensorboard 的 SummaryWriter 写 runs/ 目录;add_scalar(tag, value, global_step) 记点,斜杠分组;文档建议用 epoch 号做横轴,书里改用累计样本数——epoch 长度可变时曲线才可比;smoothing 0.6,注意别留下 junk runs。 2