跳到主要内容

训练工程学:调参、回调(挂在训练流程上、到点自动触发的钩子)、可视化、部署与数据管线

这一章讲三件事: 第 04 章那 11 个手工实验怎么自动化;训练过程中怎么「埋事件」、怎么看曲线; 以及模型训练完之后的事——存盘、上线、喂数据。 一句话:第 04 章教你训练一个模型,这一章教你把它当项目做。

1. 主走查:一个不写代码就能调用的模型

本章的主走查是书里部署环节的演示,它回答一个实战问题:模型练好了,别人怎么用? 不用给他 Python 环境,给他一个网址:

TF Serving 官方案例:部署一个叫 half_plus_two 的模型(顾名思义:输入除以 2 再加 2)。任何人向服务地址 POST 一段 JSON——{"instances":[1.0, 2.0, 5.0]}——服务立刻返回 {"predictions":[2.5, 3.0, 4.5]};把输入换成 [10.0, 2.0, 5.0],返回 [7.0, 3.0, 4.5]1

三进三出,输入减半加二,一步不差——而服务端一行模型代码都不用写。真实项目里,half_plus_two 换成你训练的 MNIST 模型即可,「程序是一样的」2。这是「训练完成」与「产生价值」之间那座桥,本节前后的一切(存盘格式、部署平台)都是这座桥的部件。

2. 自动调参:把第 04 章的 11 个实验交给机器

第 04 章的实验有个局限:一次只动一个旋钮。旋钮一多,组合数爆炸,手工跑不完。工具叫 Keras Tuner:声明每个旋钮的候选范围(学习率取 0.01/0.001/0.0001,第一层神经元取 32~512),剩下的组合让机器跑3

书里的实战:用 Hyperband 策略(要点是不再每组合都训到底,先海选再重点培养——这是我们的补充解释,机制细节不在书里)自动调出最佳组合:第一层 Dense 输出 160、学习率 0.0014。策略可换:组合太多跑不完就 RandomSearch(随机抽样);想更聪明就用 BayesianOptimization——根据上一轮的结果决定下一轮试什么(搭配高斯过程)5

作者对「为什么要调参」的总结值得原样记住:深度学习是黑箱科学,加上我们对高维(成千上万个维度)数据的联合分布不熟,唯有通过大量实验才能获得较佳模型;而模型训练非常耗时,「如何缩短调校时间,是工程师建构 AI 模型时须思考如何改善的课题」6。调参工具存在的理由不是偷懒,是省下工程师等结果的时间。

3. 回调函数:往训练循环里埋事件

训练动辄几十轮,人不可能盯全程。Callback(回调函数)=在训练的固定时点自动触发的钩子:每轮开始/结束、每批开始/结束,都能挂上你自己的逻辑7。三个最常用的:

  • EarlyStopping(提前停止):设定「连续 3 轮验证准确率没改善就停」。书里的实测:原计划训练 20 轮,实际第 12 轮就停了——省了 8 轮的电,效果还不打折8;

  • ModelCheckpoint:把检查点(训练中途的存档)存盘,训练中断不用重来。书里先训到准确率 0.9859 停掉,再加载检查点(训练中途的存档)续训 3 轮,涨到 0.9902——断点续训不是理论,是这几行输出9;

  • TensorBoard 回调:训练实时写入日志(训练过程写下的记录文件),浏览器打开就能看曲线(下一节)。

Callback 还有个隐藏身份:调试器。自定义一个「每批结束记录损失」的回调,能看见优化过程的真面目——书里的曲线显示,损失不是一路下降,而是起起伏伏,整体趋势向下10。如果你以为训练应该平滑收敛,看到抖动就慌,这张图是解药。出了 NaN 或不收敛,回调可以逐批检查中间状态,书里明说这是「便捷的除错功能」11

4. TensorBoard:训练过程的仪表盘

TensorBoard 是可视化的诊断工具,读模型日志、画曲线。功能清单12:

  • 效果指标曲线(损失/准确率随轮次变化);
  • 运算图:模型的计算流程,一层层数据怎么流;
  • 权重直方图:加一个参数 histogram_freq=1,每轮画一张权重分布——权重怎么从随机初始化慢慢「各就各位」,肉眼可见13;
  • 词嵌入三维投影:输入单词 king,空间里跟它最近的词就亮起来(第 11 章讲完词向量再回来看,会更明白)14;

有意思的一条花絮:PyTorch 团队自己做的 TensorBoardX 与它兼容——竞争对手也不得不承认这个工具优秀,书里把这当成 TensorBoard 地位的注脚15。它还有个「hparams」面板,能并排比较多组调参实验的结果;书里跑出的最佳组合是 dropout(随机丢弃神经元)=0.2、神经元 32、优化器 adam,准确度 0.977516

5. 存盘与部署:主走查的两块基石

存盘格式二选一:官方推荐的 SavedModel(一个目录,结构、权重、compile 配置、优化器状态各归其位——优化器状态在,断点就能续训)和单文件的 H5 格式(存不了自定义神经层)17。存的是「整套可复现的训练现场」,不只是权重。

部署三条路:

  1. 自己写网页:Streamlit 只用 Python 就能搭出上传图片→显示辨识结果的页面(书里几十行代码);坑只有一个:训练数据 MNIST 的白是 0,而普通图片 RGB 的白是 255,上传的图要先反转颜色,否则模型看的是「负片」18;
  2. TF Serving:主走查的方案——高效服务系统,自动暴露 REST/gRPC 接口(别的程序调用的入口),只支持 Linux(书里用 Docker 或 WSL 跑)19;
  3. 云端或边缘设备(物联网网关做初步过滤再回传数据中心)20

6. 数据管线:Dataset,让数据「流」进模型

最后一个零件常被初学者跳过,却决定项目能不能规模化:数据加载。TensorFlow Dataset 类似 Python 的生成器——逐批读数据,不把全库塞进内存(数据集一大,内存就爆了),还自带 map(逐条转换)、filter(筛选)、repeat(复制,数据太少时扩充)、shard(分片给多机)等操作21

两个性能开关,书里用时间轴图讲得极清楚22:

  • prefetch(预取):GPU 在训练第 N 批时,CPU 提前读好第 N+1 批——读数据和训练并行,不再互相等;
  • cache(缓存):第一轮读过的数据留在内存/磁盘,后面每轮不再重复读硬盘。

还有一个格式:TFRecord——TensorFlow 自家的二进制序列化(转成可存储传输的二进制)格式(遵循 Google Protocol Buffer),跨平台跨语言,图像语音等二进制数据都能装23。练手用 NumPy 数组即可,上规模就换它。

7. 可带走的

  1. 调参交给工具:Keras Tuner 声明范围、机器跑组合;Hyperband/RandomSearch/BayesianOptimization 按「组合多少、想多聪明」选;
  2. 为什么必须调参:黑箱+高维联合分布不熟→唯有实验;瓶颈是训练耗时,工具省的是等待;
  3. EarlyStopping 实测省 8 轮(20→12),ModelCheckpoint 断点续训(0.9859→0.9902)——两个回调是长训练的保险绳;
  4. 损失天生是振荡的:逐批记录可见「起伏但整体向下」,这不是故障;
  5. 回调可当调试器:NaN、不收敛,逐批检查;
  6. TensorBoard:曲线/运算图/权重直方图(histogram_freq=1);连 PyTorch 都兼容它;
  7. 存盘存全套:SavedModel 含优化器状态,断点续训的前提;H5 存不了自定义层;
  8. 部署:不写代码暴露 API(主走查:[1.0,2.0,5.0]→[2.5,3.0,4.5]);自建网页注意「反转颜色」这类训练/上线数据规格差异;
  9. 数据管线:Dataset 逐批流式;prefetch 让读和训并行,cache 省重复读盘;上规模用 TFRecord。

8. 原文地图

主题原书章原文位置
Keras Tuner 与最佳组合4-8 超参数调校text/32-ch04-4-8.txt:23(搜「Hyperband」) · text/32-ch04-4-8.txt:45(搜「最佳参数组合」)
三种调参策略4-8 超参数调校text/32-ch04-4-8.txt:71(搜「Hyperband:测试所有组合」) · text/32-ch04-4-8.txt:75(搜「高斯过程」)
为什么要调参4-8 超参数调校text/32-ch04-4-8.txt:91(搜「模型训练的执行非常耗时」)
EarlyStopping 实测5-4 回调函数text/36-ch05-5-4.txt:29(搜「EarlyStopping用于设定训练提前结束」) · text/36-ch05-5-4.txt:39(搜「预计训练20次」)
检查点续训 0.9859→0.99025-4 回调函数text/36-ch05-5-4.txt:59(搜「最后的准确率等于0.9859」) · text/36-ch05-5-4.txt:63(搜「0.9902」)
损失振荡、回调调试5-4 回调函数text/36-ch05-5-4.txt:123(搜「起起伏伏」) · text/36-ch05-5-4.txt:135(搜「除错(Debug)」)
TensorBoard 功能与直方图5-5 TensorBoardtext/37-ch05-5-5-tensorboard.txt:5(搜「TensorBoardX」) · text/37-ch05-5-5-tensorboard.txt:77(搜「histogram_freq」)
词嵌入投影、XAI 伏笔5-5 TensorBoardtext/37-ch05-5-5-tensorboard.txt:23(搜「词嵌入(Word Embedding)展示」) · text/37-ch05-5-5-tensorboard.txt:73(搜「可解释的AI」)
TensorBoard 调参面板5-5 TensorBoardtext/37-ch05-5-5-tensorboard.txt:107(搜「dropout rate=0.2」)
存盘两种格式5-2 模型存盘与加载text/34-ch05-5-2.txt:9(搜「SavedModel格式」) · text/34-ch05-5-2.txt:11(搜「无法存储自定义的神经层」)
TF Serving 主走查5-6 模型部署与TensorFlow Servingtext/38-ch05-5-6-tensorflow-serving.txt:61(搜「half_plus_two」) · text/38-ch05-5-6-tensorflow-serving.txt:67(搜「2.5, 3.0, 4.5」)
Streamlit 与反转颜色5-6 模型部署与TensorFlow Servingtext/38-ch05-5-6-tensorflow-serving.txt:17(搜「Streamlit最为简单」) · text/38-ch05-5-6-tensorflow-serving.txt:41(搜「反转颜色」)
Dataset 逐批读取5-7 TensorFlow Datasettext/39-ch05-5-7-tensorflow-dataset.txt:5(搜「支持缓存(Cache)」)
prefetch/cache5-7 TensorFlow Datasettext/39-ch05-5-7-tensorflow-dataset.txt:195(搜「prefetch」) · text/39-ch05-5-7-tensorflow-dataset.txt:197(搜「空档先读取下一批数据」)
TFRecord5-7 TensorFlow Datasettext/39-ch05-5-7-tensorflow-dataset.txt:109(搜「Protocol Buffer」)
部署三级(本地/云端/边缘)5-6 模型部署与TensorFlow Servingtext/38-ch05-5-6-tensorflow-serving.txt:11(搜「IoT Hub」)

Footnotes

  1. 出处:「5-6 模型部署与TensorFlow Serving」第 61 段(text/38-ch05-5-6-tensorflow-serving.txt:61,搜「half_plus_two」)、第 65 段(text/38-ch05-5-6-tensorflow-serving.txt:65,搜「1.0, 2.0, 5.0」)与第 67 段(text/38-ch05-5-6-tensorflow-serving.txt:67,搜「2.5, 3.0, 4.5」)。修改输入得 [7.0, 3.0, 4.5] 在第 69 段。

  2. 出处:「5-6 模型部署与TensorFlow Serving」第 71 段(text/38-ch05-5-6-tensorflow-serving.txt:71,搜「如法炮制」)。原文:「换上用户自己的模型,程序是一样的」。

  3. 出处:「4-8 超参数调校」第 17 段(text/32-ch04-4-8.txt:17,搜「0.01, 0.001, 0.0001」)。

  4. 出处:「4-8 超参数调校」第 45 段(text/32-ch04-4-8.txt:45,搜「最佳参数组合」)。

  5. 出处:「4-8 超参数调校」第 71 段(text/32-ch04-4-8.txt:71,搜「Hyperband:测试所有组合」)与第 75 段(text/32-ch04-4-8.txt:75,搜「高斯过程」)。

  6. 出处:「4-8 超参数调校」第 91 段(text/32-ch04-4-8.txt:91,搜「模型训练的执行非常耗时」)。

  7. 出处:「5-4 回调函数」第 5 段(text/36-ch05-5-4.txt:5,搜「在每一个周期执行之前与之后」)。

  8. 出处:「5-4 回调函数」第 29 段(text/36-ch05-5-4.txt:29,搜「EarlyStopping用于设定训练提前结束」)与第 39 段(text/36-ch05-5-4.txt:39,搜「预计训练20次」)。

  9. 出处:「5-4 回调函数」第 59 段(text/36-ch05-5-4.txt:59,搜「最后的准确率等于0.9859」)与第 63 段(text/36-ch05-5-4.txt:63,搜「0.9902」)。

  10. 出处:「5-4 回调函数」第 123 段(text/36-ch05-5-4.txt:123,搜「起起伏伏」)。

  11. 出处:「5-4 回调函数」第 135 段(text/36-ch05-5-4.txt:135,搜「除错(Debug)」)。

  12. 出处:「5-5 TensorBoard」第 11 段(text/37-ch05-5-5-tensorboard.txt:11,搜「追踪损失和准确率」)。

  13. 出处:「5-5 TensorBoard」第 77 段(text/37-ch05-5-5-tensorboard.txt:77,搜「histogram_freq」)。

  14. 出处:「5-5 TensorBoard」第 23 段(text/37-ch05-5-5-tensorboard.txt:23,搜「词嵌入(Word Embedding)展示」)。

  15. 出处:「5-5 TensorBoard」第 5 段(text/37-ch05-5-5-tensorboard.txt:5,搜「TensorBoardX」)。

  16. 出处:「5-5 TensorBoard」第 113 段(text/37-ch05-5-5-tensorboard.txt:113,搜「最佳准确度0.9775」)与第 107 段(text/37-ch05-5-5-tensorboard.txt:107,搜「dropout rate=0.2」)。

  17. 出处:「5-2 模型存盘与加载」第 9 段(text/34-ch05-5-2.txt:9,搜「SavedModel格式」)与第 11 段(text/34-ch05-5-2.txt:11,搜「无法存储自定义的神经层」)。「优化器状态使训练可由断点继续」见第 5 段(搜「断点处继续执行」)。

  18. 出处:「5-6 模型部署与TensorFlow Serving」第 17 段(text/38-ch05-5-6-tensorflow-serving.txt:17,搜「Streamlit最为简单」)与第 41 段(text/38-ch05-5-6-tensorflow-serving.txt:41,搜「反转颜色」)。

  19. 出处:「5-6 模型部署与TensorFlow Serving」第 47 段(text/38-ch05-5-6-tensorflow-serving.txt:47,搜「不需撰写程序」)。gRPC/REST 与 Linux 限制在同段。

  20. 出处:「5-6 模型部署与TensorFlow Serving」第 11 段(text/38-ch05-5-6-tensorflow-serving.txt:11,搜「IoT Hub」)。

  21. 出处:「5-7 TensorFlow Dataset」第 5 段(text/39-ch05-5-7-tensorflow-dataset.txt:5,搜「支持缓存(Cache)」)。

  22. 出处:「5-7 TensorFlow Dataset」第 195 段(text/39-ch05-5-7-tensorflow-dataset.txt:195,搜「prefetch」)与第 197 段(text/39-ch05-5-7-tensorflow-dataset.txt:197,搜「空档先读取下一批数据」)。

  23. 出处:「5-7 TensorFlow Dataset」第 109 段(text/39-ch05-5-7-tensorflow-dataset.txt:109,搜「Protocol Buffer」)。