跳到主要内容

鸟与飞机:分类、softmax 与交叉熵(一种给分类错误打分的办法),以及全连接的极限

这一章讲三件事: 第一个真实任务——看图分辨鸟和飞机——从数据到模型到训练怎么走完;分类问题的输出和损失为什么要专门设计(softmax 与交叉熵);以及全连接结构撞上的那堵墙:它不知道「位置」是什么。 位置:零件全部就绪后第一次实战,结尾的战败直接引出第 08 章。

1. 任务与数据:Jane 的鸟

书里给了任务一个人情味的包装:Jane 在树林上架了相机拍鸟,但旁边机场的飞机不停触发快门,她要一个能自动删掉飞机照片的东西1

数据集用 CIFAR-10:6 万张 32×32 彩色小图,分 10 类(飞机、汽车、鸟、猫……),每类 6000 张2。它小到今天已不能用来验证新方法,但教学刚好。只取飞机和鸟两类,凑成二分类(把标签 0 和 2 重映射(重新对应)成 0 和 1),训练集约 1 万张。

拿数据的方式本身是一个要学的抽象3:

  • Dataset(数据集):一个只要求两个方法的对象——__len__(我有多少个样本)和 __getitem__(给我第 i 个)。它不必把数据装进内存,可以现取现算;这两个方法就是数据和 PyTorch 之间的全部桥梁。
  • DataLoader(数据加载器):包在 Dataset 外面,负责打乱顺序、按小批量分组、用后台进程预取。DataLoader(cifar2, batch_size=64, shuffle=True) 之后,训练循环每次拿到的就是 64 张图 + 64 个标签。

小批量(minibatch)在第 05 章出现过:一次用一小撮样本估计梯度再更新。100 轮训练,每轮 157 个小批量,每批 64 张——损失从 0.52 一路降到 0.024

2. 输出设计:让网络说「概率」

回归(第 04–06 章,预测一个连续数值)的输出是一个数;分类(classification:从有限个类别里挑一个)的输出该长什么样?

答案:一个类别一个数,并且这排数要满足概率的两条规矩——每个都在 0 到 1 之间、加起来等于 1。 这样,输出就能读作「模型认为各类别的可能性」5

把任意一排实数变成这样一排数的函数就是 softmax(第 01 章见过一次):对每个数取指数(全变正),再除以总和(归一)。书里的具体例子:[1.0, 2.0, 3.0] 过 softmax 得 [0.090, 0.245, 0.665]——大小次序保住了(3 最大,出来还是最大),而且它会放大差距:输入 2 是 1 的两倍,输出里 0.245 却是 0.090 的 2.7 倍6

未训练的网络上跑第一张鸟图,输出 [0.478, 0.522]——模型说是鸟。书里很诚实:「Pure luck」——随机初始化的权重,答对纯属蒙的7

3. 损失设计:只奖励「把正确答案的概率抬高」

第 04 章的均方误差搬过来行不行?不行,有两个层次的原因8

浅一层:我们不真的在乎输出是不是精确的 [1.0, 0.0];在乎的是正确类别的概率高于其他。均方误差会没完没了地惩罚「0.9 还不够像 1」,而分类只要求排序对。

深一层,书里给了一条实测曲线:均方误差配合 softmax,在预测快对时梯度就躺平了;而交叉熵在预测到 99.97% 正确时仍有坡度——「快对了」和「完全对」之间那一点差别,它还愿意追8

分类的标准损失是这么搭出来的。先给一个名词:模型把正确类别的概率赋得多高,叫这组参数的似然(likelihood——「现有参数之下,这批数据出现的可能性」);我们要最大化它。等价地,最小化它的负对数——负对数似然(NLL,negative log likelihood):-log(正确类别的概率),再对批量求和。概率越接近 0,负对数越冲向无穷——把正确答案说得越没可能,罚得越狠;过了 0.5 之后罚得平缓9

PyTorch 里两个零件拼一件事:nn.LogSoftmax(数值稳定版的 softmax+取对数)+ nn.NLLLoss;更常用的是合体版 nn.CrossEntropyLoss——交叉熵,直接吃 softmax 之前的那排裸分数(这排裸分数有个专门名字,叫 logits),内部一步完成9。日常代码都写 CrossEntropyLoss;代价是模型输出不再直接是概率,要概率就自己补一遍 softmax。

4. 战果与败因:81%,以及「它不认得挪了位置的飞机」

战果

一层隐藏层的全连接网络(3072→512→2),100 轮后:验证集准确率 81.1%10。换成更深的四层版(3072→1024→512→128→2):验证 81.3%,但训练集准确率 100%——教科书级的过拟合,第 05 章的 B 型曲线11

参数都花在哪了

数一下参数12:第一个模型 157 万个,四层版 374 万个。看明细 [3145728, 1024, 524288, …]——第一层一个就占 314 万(1024×3072+1024:每个隐藏神经元都要为全部 3072 个输入各配一个权重)。外推(把这里的算术推广到没试过的尺寸)到一张 1024×1024 的正常照片:输入 310 万维,第一层就要 30 多亿参数,12GB 显存只够装第一层。全连接结构在图片尺寸面前根本不扩展。12

真正的病根:它不知道「位置」

比参数更多的是结构盲区。把图片拉直成 3072 个数的那一瞬间,「哪两个像素相邻」这个信息就被扔了——对网络来说,第 176 个输入和第 208 个输入(其实上下相邻的两个像素)和第 3000 个输入没有任何区别13

后果用一个思想实验看清楚:飞机在位置 (4,4) 时,网络学到「(0,1) 暗 + (1,1) 暗 + … → 像飞机」;同一架飞机挪到 (8,8),所有像素关系全变,这套规律要整个重学13。这个性质叫平移不变性(translation invariance:同一个东西挪个位置,答案不该变)——全连接网络没有它。

位置 (4,4) 的飞机 位置 (8,8) 的同一架飞机
学到的规律:(0,1) 暗 ∧ 全部作废,要从头学:
(1,1) 暗 ∧ … (0,2) 暗 ∧ (1,2) 暗 ∧ …

图说:每个位置都要重新学一遍「什么是飞机」——参数量就是这样爆炸的。

补救办法不是没有:把训练图随机平移生成副本(数据增广,第 14 章细讲),让网络把每个位置都见一遍。但那要用更多参数去「存」每个位置的副本——治标,而且贵13

这堵墙就是下一章的起点:有没有一种结构,天生就把「位置无关」写进骨子里? 有,卷积。

5. 作者的判断与证据

说法性质
81.1% / 81.3% / 训练 100%有据,书里的真实运行输出1011
参数量(可以调的数的总数) 157 万 / 374 万 / 1024×1024 图要 30 亿有据,numel 逐项打印,外推的算术可以自验(1024×1024×3×1024+1024 ≈ 32 亿)12
「交叉熵在接近正确时仍有梯度,MSE 早早饱和」有据,书里印了两种损失随预测分数变化的实测曲线8
「这两个类好分,部分因为背景色等系统差异」作者的坦白式推测:81% 的成色里有多少是真本事,没拆开10
「minibatch 的噪声帮助跳出局部极小」领域通行解释,书里以直陈句给出;视为「有用的工作假设」4

6. 边界与局限

  • 全连接 + 表格数据、无位置关系的任务,本章这套就是成品——「全连接的极限」只在有空间结构的数据上成立13
  • 二分类的两输出写法有冗余(一个概率就够,另一个=1 减它);书里给了 nn.BCELoss 系作为替代,正文没展开5
  • 「模型输出当概率读」有个深坑书里只点了注:训练出来的分类器(做分类的模型)通常过度自信,0.9 的概率不保证 90% 的时候对;校准问题是另一个课题7
  • 本章没回答「同一图里既有鸟又有飞机怎么办」「鸟在哪个位置」——前者要多标签,后者就是第 15 章的分割。

7. 可带走的

  1. Dataset 只要 __len____getitem__;DataLoader 管打乱、分批、后台预取。
  2. 分类输出 = 一类一个数,过 softmax 当概率读;softmax 保序但放大差距。
  3. 分类的标准损失是交叉熵:只追「正确类别的概率」,接近正确时仍有梯度;CrossEntropyLoss 直接吃 logits,省去手写 LogSoftmax。
  4. 训练集 100% + 验证集 81% = 过拟合确诊,先想结构,别急着加技巧。
  5. 第一层线性层吃掉绝大多数参数;全连接对图片尺寸是平方级爆炸。
  6. 拉直图片 = 扔掉位置信息;没有平移不变性的模型,每个位置都要重学。
  7. 数据增广是拿参数换不变性,治标;治本的结构在下一章。

8. 原文地图

主题原书章原文位置
Jane 的鸟任务ch7text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:321(搜「bird-watching club」)
CIFAR-10ch7text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:31(搜「60,000 tiny 32 × 32」)
Dataset/DataLoaderch7text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:93(搜「len and getitem」) · :878(搜「sample minibatches」)
过滤成两类ch7text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:357(搜「label_map = {0: 0, 2: 1}」)
3072 输入全连接ch7text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:384(搜「3,072 input features」)
概率约束与 softmaxch7text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:478(搜「add up to 1.0」) · :530(搜「0.0900, 0.2447, 0.6652」)
未训练输出 Pure luckch7text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:603(搜「0.4784, 0.5216」) · :637(搜「Pure luck」)
似然与 NLLch7text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:655(搜「likelihood」) · :666(搜「negative log likelihood」)
交叉熵 vs MSE 曲线ch7text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:743(搜「99.97%」)
LogSoftmax+NLL=CrossEntropych7text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:981(搜「Combining nn.LogSoftmax and nn.NLLLoss」)
小批量训练与 81.1%ch7text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:926(搜「64 × 3 × 32 × 32」) · :959(搜「0.811000」)
深模型过拟合ch7text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:1009(搜「0.813000」) · :1010(搜「1.000000」)
参数量账ch7text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:1028(搜「3737474」) · :1055(搜「3.1 million input values」)
平移不变性ch7text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:1089(搜「dark, cross-like shape」) · :1096(搜「not translation invariant」)
增广治标ch7text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:1138(搜「augment the dataset」)

Footnotes

  1. 出处:「7 Telling birds from airplanes」第 321 段(text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:321,搜「bird-watching club」)。Jane 的相机要自动删掉误触发的飞机照片。

  2. 出处:「7 Telling birds from airplanes」第 31 段(text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:31,搜「60,000 tiny 32 × 32」)。CIFAR-10 由 Hinton 等人收集自 8000 万张小图;今天对研究太简单,教学正好。

  3. 出处:「7 Telling birds from airplanes」第 93 段(text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:93,搜「len and getitem」)与第 878 段(:878,搜「sample minibatches」)。Dataset 甚至可以不持有数据,只提供统一访问;DataLoader 从 Dataset 里采样小批量。

  4. 出处:「7 Telling birds from airplanes」第 845 段(text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:845,搜「partial estimation of the gradient」)与第 929 段(:929,搜「0.523478」)。原文:小批量梯度「是梯度的部分估计……跟着这种更差的估计走,反而帮助收敛、避免卡进局部极小」;训练 100 轮损失从 0.52 降到 0.02。 2

  5. 出处:「7 Telling birds from airplanes」第 478 段(text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:478,搜「add up to 1.0」)。概率两约束:每项 ∈[0,1],总和=1;二分类时两输出的冗余与 BCELoss 替代见同章注。 2

  6. 出处:「7 Telling birds from airplanes」第 530 段(text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:530,搜「0.0900, 0.2447, 0.6652」)与第 540 段(:540,搜「not scale invariant」)。softmax 单调保序但改变比例:[1,2,3] 的首两项之比 0.5,输出里变成约 0.37。

  7. 出处:「7 Telling birds from airplanes」第 603 段(text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:603,搜「0.4784, 0.5216」)与第 637 段(:637,搜「Pure luck」);过度自信的提醒见第 611 段(搜「overconfident」):训练出的分类器常过度自信,贝叶斯神经网络是一种补救但超出本书范围。 2

  8. 出处:「7 Telling birds from airplanes」第 743 段(text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:743,搜「99.97%」)。原文:正确类别预测概率达 99.97% 时,交叉熵仍保持坡度,而 MSE 早已饱和——MSE 的坡太缓,压不过 softmax 的平坦区。那条对照曲线是原书的图 7.11。 2 3

  9. 出处:「7 Telling birds from airplanes」第 655 段(text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:655,搜「likelihood」)、第 666 段(:666,搜「negative log likelihood」)与第 981 段(:981,搜「Combining nn.LogSoftmax and nn.NLLLoss」)。NLL=−Σlog(正确类概率);LogSoftmax+NLLLoss 与 CrossEntropyLoss 数值相同,后者直接吃 logits。 2

  10. 出处:「7 Telling birds from airplanes」第 959 段(text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:959,搜「0.811000」)与第 962 段(:962,搜「shallow classifier」)。验证准确率 81.1%;作者自评:模型很浅,「它能工作简直是个奇迹」——两个类的样本可能有系统性差异(比如背景色)帮了忙。 2 3

  11. 出处:「7 Telling birds from airplanes」第 1009 段(text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:1009,搜「0.813000」)与第 1010 段(:1010,搜「1.000000」)。四层模型验证 81.3%、训练 100%:「两个模型都在过拟合」。 2

  12. 出处:「7 Telling birds from airplanes」第 1028 段(text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:1028,搜「3737474」)与第 1055 段(:1055,搜「3.1 million input values」)。1024×3072+1024=3,146,752 个参数在第一层;1024×1024 RGB 图的第一层将超 30 亿参数、仅权重就 12GB。 2 3

  13. 出处:「7 Telling birds from airplanes」第 436 段(text/15-ch07-7-telling-birds-from-airplanes-learning-from-ima.txt:436,搜「structurally unaware」)、第 1089 段(:1089,搜「dark, cross-like shape」)与第 1138 段(:1138,搜「augment the dataset」)。原文:网络「结构上就不知道」两个输入相邻;平移后的飞机要「从头重学」;增广要为平移副本付出参数。 2 3 4