跳到主要内容

三个入门任务 — 分类与回归的手感

这一章讲三件事: 入门级的三类任务各长什么样;每个任务各藏一条 以后会反复用的经验;以及「怎么知道该训几轮」这个问题的第一个答案。 前面四章把工具和原理都备齐了,这一章是第一次完整干活—— 三个任务,三套数据,三个坑,一个比一个真实。

1. 顶层全景

书里给深度学习入门任务分了三类,全书的任务几乎都落进这三类1:

任务数据输出最后一层损失函数
二分类IMDB 影评 5 万条正面还是负面1 个 sigmoid 单元binary_crossentropy
多分类路透社新闻 11228 条46 个主题之一46 个 softmax 单元categorical_crossentropy
标量回归波士顿房价 506 条房价(一个连续数)1 个无激活单元mse(均方误差)

「标量回归」就是预测一个连续取值的数(房价),而不是预测类别。 注意最后一行的「无激活」:回归任务的最后一层故意不加激活函数, 这样输出可以取任意范围的值——加了 sigmoid 反而把房价锁死在 0 到 1 之间2

本章的主走查是第一个任务:一条 IMDB 影评,从一段英文文本变成 10000 维的 0/1 向量, 穿过两个 Dense(16),变成 sigmoid 的一个数,最后停在 88% 测试精度—— 以及验证损失曲线在第 4 轮的那个掉头。

2. 任务一:IMDB 影评二分类

2.1 文本怎么变成向量

IMDB 数据集是 5 万条两极分化的影评,2.5 万训练、2.5 万测试,正负各半3。 书里用的版本已经把每条影评转成了整数序列(每个整数代表一个单词), 而且只保留训练数据里最常出现的前 10000 个词——不设这个限,词表会有 88585 个词, 绝大多数只出现一两次,没有信息量4

回到那条影评:主走查第 ① 步。 一条影评此刻是一串整数,比如 [8, 5, 234, 17, ...]。 这串整数不能直接喂网络——第 02 章说过,Dense 层吃的是定长向量。 书里用的办法叫 multi-hot 编码:造一个 10000 维的全零向量, 序列里出现过的每个整数,把对应位置拨成 1—— 序列 [8, 5] 编码后,就是只有第 8 位和第 5 位是 1、其余 9998 位全是 0 的向量5:

影评: [8, 5, 234, 17, ...] ← 每个数是一个词的编号
编码: [0,0,0,0,0,1,0,0,1,...,1,...] ← 10000 维,出现过的词对应位置为 1

▼ Dense(16, relu) → Dense(16, relu) → Dense(1, sigmoid)
输出: 0.89 ← 「这条影评是正面」的程度

图说:multi-hot 丢掉了词序和次数,只保留「哪些词出现过」。
这个简化丢掉了词序,为什么还打得赢?第 11 章会正面回答(词袋模型)。

2.2 为什么是 16 个单元,以及交叉熵

中间层为什么是 16 个单元?书里给了一个理解方式:这个维数是「模型学习内部表示时所拥有的自由度」—— 把 10000 维的输入压到 16 维空间里重新表示。单元越多,自由度越大; 但自由度过大,模型会学到只和训练数据有关的模式6

损失函数用 binary_crossentropy(二元交叉熵)。交叉熵来自信息论, 衡量的是两个概率分布之间的距离——你的模型输出一个概率,真实标签是 0 或 1 的概率分布, 交叉熵把「这两个分布差多远」算成一个可求导的数。对输出概率的模型,它通常是最佳选择7。 优化器用 rmsprop,书里评价:对于几乎所有问题,它都是很好的默认选择,无须为此费神8

2.3 主走查第 ② 步:第 4 轮的那个掉头

留出 10000 条做验证,训练 20 轮,然后看两条曲线——这是全章最重要的一个画面9:

训练损失: 一路下降 ↘↘↘↘↘↘↘↘ (20 轮一直在降)
验证损失: ↘↘↘↗↗↗↗↗ (第 4 轮左右到最低,然后掉头向上)

图说:两条曲线在第 4 轮附近分道扬镳。这个分叉点,就是过拟合的发生点。

书里的话很准:在第 4 轮之后,你是在对训练数据过度优化,学到的表示只针对训练数据, 无法泛化到新数据10。对策也直接:照分叉点重新训练——从头训一个新模型,只训 4 轮, 测试精度约 88%。作为参照,用今天最先进的方法,这个任务大约能到 95%11

3. 任务二:路透社新闻 46 分类

第二个任务把类别数从 2 加到 46:路透社 1986 年的新闻专线数据, 8982 条训练、2246 条测试,每条新闻属于 46 个互斥主题之一12

输入照旧 multi-hot(10000 维),但两个地方要跟着变:

最后一层变成 46 个 softmax 单元——输出一个 46 维向量,和为 1, 第 i 个数读作「属于第 i 个主题的可能性」。损失函数换成 categorical_crossentropy(分类交叉熵), 它衡量的还是两个分布的距离,只是分布从 2 个类别变成了 46 个13

中间层从 16 加到 64——这就是这一节的硬知识:信息瓶颈。 46 个类别的分离信息,必须完整流过中间层那 16 维(或 64 维)的窄通道; 通道太窄,相关信息就被永久挤丢了14。书里做了对照实验:把中间层故意压到 4 维, 验证精度从约 79% 掉到约 71%,跌了 8 个点——每一层都可能成为信息瓶颈15

结果:约 80% 精度。这个数要和基准比着读——46 个类别、样本数还不均衡, 完全随机猜只有约 19%。先建一个「常识基准」再谈模型好坏,是第 06 章的正式规矩16

顺带一个工程细节:标签也可以不编码成 46 维 one-hot 向量,而直接当整数(0~45)用, 这时损失函数换成 sparse_categorical_crossentropy——数学上完全一样,只是接口不同17

4. 任务三:波士顿房价回归

4.1 只有 506 个样本的数据集

第三个任务换了一种输出:预测 70 年代中期波士顿郊区房价中位数(单位千美元)。 数据只有 506 条,404 训练、102 测试,每条 13 个特征(犯罪率、房间数等)18

13 个特征的取值范围天差地别(有的是 0~1 的比例,有的是几百),直接喂网络会让大数值特征主导梯度。 做法是逐特征标准化:每个特征减去自己这一列的平均值,再除以这一列的标准差 ——标准差就是「这一列数散得有多开」:全挤在平均值附近就小,忽大忽小拉得很远就大19

所以这两步做的事是:先把这一列的中心挪到 0,再按它自己散开的程度缩一遍。 结果是 13 列数被统一到同一个尺度上,谁也不再仗着数值大压过别人(全部变成平均值 0、标准差 1 的数)。

这里有一条铁律,书里特意加粗了语气:测试集的标准化,必须用训练集算出来的均值和标准差—— 不能使用在测试数据上计算得到的任何结果,哪怕是标准化这么简单的事。 测试数据的一切统计量都属于「偷看答案」20

4.2 K 折交叉验证:506 条数据怎么验证

模型很小(两个 64 单元层),最后一层无激活(房价可以是任意正数), 损失用 mse(均方误差),监控指标用 MAE(平均绝对误差——预测和实际差多少,单位同为千美元, MAE = 0.5 就是平均差 500 美元)21

问题来了:只有 404 条训练数据,再留 100 条做验证,验证分数会随「恰好留下哪 100 条」剧烈波动。 对策是 K 折交叉验证:把数据切成 K 份(书里 K=4),轮流拿每一份当验证集、其余当训练集, 训 K 个模型,把 K 个验证分数取平均22:

K=4 的结果:[2.11, 3.08, 2.65, 2.43] ← 四个分数差得肉眼可见
平均: 2.6 ← 平均差 2600 美元

图说:单看任何一折都会误判模型;平均值才是可靠指标。
参照物:房价本身在 1 万到 5 万美元之间,差 2600 美元算不小。

4.3 主走查第 ③ 步:500 轮曲线与 130 轮终点

K 折还有第二个用途:回答「该训几轮」。书里把每个模型训 500 轮, 把 K 个模型每轮的验证 MAE 取平均,画出一条平滑曲线—— 验证 MAE 在 120~140 轮后不再显著下降,之后开始过拟合23

于是最终模型照 130 轮重训(这次用全部非测试数据),测试 MAE 停在 2.46—— 平均误差约 2500 美元,和 K 折估计的 2.6 对得上,说明整个评估流程是自洽的24

5. 作者的判断与证据

有实测证据的: 88% / 80% / 2.46 三个终点数字,以及信息瓶颈的 8 个点、 K 折分数的波动范围,全是代码跑出来的;

作者的经验规则: 「rmsprop 几乎总是好的默认」;「分类问题几乎总该用交叉熵」; 「数据越少,越该用小模型」——这些是经验法则,书里明说是模式匹配,不是定理25;

贯穿全章的方法论: 每个任务都走同一条流水线——向量化 → 搭模型 → 留验证 → 看曲线找过拟合点 → 照点重训 → 测试一次。这条流水线在第 07 章会被正式写成「通用工作流程」, 并且前后各接上一段:前面接「定义任务、收集数据」,后面接「部署、监控、再训练」。

6. 边界与局限

  • multi-hot 丢掉了词序和词频——第 11 章会系统比较「词袋」和「序列模型」,并给出选择判据;
  • IMDB 的 88% 离最先进的 95% 有距离,原书的意思是「这只是教学模型的起点」,不是让你止步于此;
  • 波士顿房价数据集本身有历史伦理问题(其中一个特征来自种族构成相关的历史变量)—— 这条不在书里,来自通用知识:该数据集后来被 scikit-learn 从官方源移除, 教学可用,别拿它做任何真实决策;
  • 「看曲线找掉头点」在数据少时好用;数据大、训练贵时,更标准的做法是第 04 章的 EarlyStopping 回调——第 13 章还会把它接进超参数(就是那些不由训练调、只能由人事先定死的数: 训几轮、学习率多大、每层几个单元)的自动搜索里。

7. 可带走的

  1. 入门任务就三类:二分类(sigmoid+二元交叉熵)、多分类(softmax+分类交叉熵)、回归(无激活+mse);
  2. multi-hot 编码:10000 维、出现过的词为 1——文本进 Dense 层的最简方案;
  3. 中间层是信息的窄通道,太窄就是信息瓶颈(46 类配 4 维,直接丢 8 个点);
  4. 标准化铁律:测试集用训练集的均值和标准差;测试数据的任何统计量都不能用;
  5. 数据少就用 K 折交叉验证——单次的验证分数可能是噪声,平均值才是信号;
  6. 验证损失掉头的那一轮 = 过拟合发生点,照它重训;
  7. 结果永远要和基准比着读:80% 看着一般,但随机猜只有 19%;
  8. 每个任务同一条流水线:向量化 → 搭模型 → 留验证 → 找过拟合点 → 重训 → 测试一次。

8. 原文地图

主题原书章原文位置
三类任务与术语表神经网络入门:分类与回归text/11-ch04.txt:10(搜「3 种使用场景」) · text/11-ch04.txt:11(搜「标量回归」)
IMDB 数据与词表截断神经网络入门:分类与回归text/11-ch04.txt:53(搜「50000」) · text/11-ch04.txt:69(搜「88585」)
multi-hot 编码神经网络入门:分类与回归text/11-ch04.txt:101(搜「multi-hot 编码」)
16 单元与自由度神经网络入门:分类与回归text/11-ch04.txt:171(搜「自由度」)
交叉熵与 rmsprop神经网络入门:分类与回归text/11-ch04.txt:200(搜「交叉熵」) · text/11-ch04.txt:205(搜「默认选择」)
第 4 轮掉头与 88%神经网络入门:分类与回归text/11-ch04.txt:279(搜「第 4 轮」) · text/11-ch04.txt:304(搜「88%」)
路透社 46 类与信息瓶颈神经网络入门:分类与回归text/11-ch04.txt:354(搜「46 个主题」) · text/11-ch04.txt:421(搜「信息瓶颈」)
80% 与 19% 基准神经网络入门:分类与回归text/11-ch04.txt:523(搜「80%」) · text/11-ch04.txt:532(搜「19%」)
4 维瓶颈实验神经网络入门:分类与回归text/11-ch04.txt:566(搜「4 维」) · text/11-ch04.txt:586(搜「71%」)
波士顿房价与标准化铁律神经网络入门:分类与回归text/11-ch04.txt:617(搜「506」) · text/11-ch04.txt:659(搜「测试数据上计算」)
K 折与 2600 美元神经网络入门:分类与回归text/11-ch04.txt:692(搜「K 折交叉验证」) · text/11-ch04.txt:748(搜「2.1 到 3.1」) · text/11-ch04.txt:750(搜「2600 美元」)
120~140 轮与 MAE 2.46神经网络入门:分类与回归text/11-ch04.txt:803(搜「120 ~ 140 轮」) · text/11-ch04.txt:821(搜「2.46」)

Footnotes

  1. 出处:「神经网络入门:分类与回归」第 10 段(text/11-ch04.txt:10,搜「3 种使用场景」)。

  2. 出处:「神经网络入门:分类与回归」第 190 段(text/11-ch04.txt:190,搜「线性层」)。原文:「最后一层是纯线性层……模型可以自由地预测任意范围的值」。

  3. 出处:「神经网络入门:分类与回归」第 53 段(text/11-ch04.txt:53,搜「50000」)。

  4. 出处:「神经网络入门:分类与回归」第 69 段(text/11-ch04.txt:69,搜「88585」)与第 91 段(text/11-ch04.txt:91,搜「padding」)。索引 0、1、2 分别保留给填充、序列开始、未知词。

  5. 出处:「神经网络入门:分类与回归」第 101 段(text/11-ch04.txt:101,搜「multi-hot 编码」)。

  6. 出处:「神经网络入门:分类与回归」第 75 段(text/11-ch04.txt:75,搜「16」)与第 171 段(text/11-ch04.txt:171,搜「自由度」)。

  7. 出处:「神经网络入门:分类与回归」第 200 段(text/11-ch04.txt:200,搜「交叉熵」)。

  8. 出处:「神经网络入门:分类与回归」第 205 段(text/11-ch04.txt:205,搜「默认选择」)与第 338 段(text/11-ch04.txt:338,搜「足够好的选择」)。

  9. 出处:「神经网络入门:分类与回归」第 64 段(text/11-ch04.txt:64,搜「10000」)与第 279 段(text/11-ch04.txt:279,搜「第 4 轮」)。

  10. 出处:「神经网络入门:分类与回归」第 282 段(text/11-ch04.txt:282,搜「过度优化」)。

  11. 出处:「神经网络入门:分类与回归」第 304 段(text/11-ch04.txt:304,搜「88%」)。

  12. 出处:「神经网络入门:分类与回归」第 351 段(text/11-ch04.txt:351,搜「路透社」)、第 354 段(text/11-ch04.txt:354,搜「46 个主题」)与第 365 段(text/11-ch04.txt:365,搜「8982」)。

  13. 出处:「神经网络入门:分类与回归」第 436 段(text/11-ch04.txt:436,搜「概率分布」)与第 438 段(text/11-ch04.txt:438,搜「categorical_crossentropy」)。

  14. 出处:「神经网络入门:分类与回归」第 417 段(text/11-ch04.txt:417,搜「输出空间的维度」)与第 421 段(text/11-ch04.txt:421,搜「信息瓶颈」)。

  15. 出处:「神经网络入门:分类与回归」第 586 段(text/11-ch04.txt:586,搜「71%」)。

  16. 出处:「神经网络入门:分类与回归」第 523 段(text/11-ch04.txt:523,搜「80%」)与第 532 段(text/11-ch04.txt:532,搜「19%」)。

  17. 出处:「神经网络入门:分类与回归」第 557 段(text/11-ch04.txt:557,搜「sparse_categorical_crossentropy」)与第 562 段(text/11-ch04.txt:562,搜「数学上」)。

  18. 出处:「神经网络入门:分类与回归」第 614 段(text/11-ch04.txt:614,搜「波士顿房价」)、第 617 段(text/11-ch04.txt:617,搜「506」)与第 634 段(text/11-ch04.txt:634,搜「千美元」)。

  19. 出处:「神经网络入门:分类与回归」第 618 段(text/11-ch04.txt:618,搜「取值范围」)与第 644 段(text/11-ch04.txt:644,搜「标准化」)。

  20. 出处:「神经网络入门:分类与回归」第 659 段(text/11-ch04.txt:659,搜「测试数据上计算」)。

  21. 出处:「神经网络入门:分类与回归」第 673 段(text/11-ch04.txt:673,搜「mse」)与第 684 段(text/11-ch04.txt:684,搜「500 美元」)。

  22. 出处:「神经网络入门:分类与回归」第 686 段(text/11-ch04.txt:686,搜「K 折交叉验证」)与第 689 段(text/11-ch04.txt:689,搜「很大波动」)。

  23. 出处:「神经网络入门:分类与回归」第 803 段(text/11-ch04.txt:803,搜「120 ~ 140 轮」)。

  24. 出处:「神经网络入门:分类与回归」第 821 段(text/11-ch04.txt:821,搜「2.46」)。

  25. 出处:「神经网络入门:分类与回归」第 598 段(text/11-ch04.txt:598,搜「几乎总是」)与第 663 段(text/11-ch04.txt:663,搜「较小的模型」)。