章节摘要:训练循环的第一步是「取出一个批次」。TensorFlow 把这件事交给
tf.data.Dataset:从张量切片或文件构建,再map/shuffle/batch/prefetch。PyTorch 拆成Dataset(怎么取一条)和DataLoader(怎么打成批次、是否多进程)。本章先分别按原文把两条管线写清楚,再专节对照:预取对num_workers、打乱窗口、与fit/ 手写循环的接口。读完你应能把同一份内存里的特征-标签对,写成两边都能喂给第 2.5 节五步的迭代器。
阅读完本章,你应当能够:
from_tensor_slices 建一条可 batch 的 TF 管线并接到 model.fitDataset 的 __len__ 与 __getitem__,并用 DataLoader 取出形状正确的批次num_workers 的原文提示金句:模型再快,也会在等下一批数据时睡着;管线对照比层名字对照更常决定你能不能跑满 GPU。
从「为什么需要」到 map/batch/shuffle/prefetch/repeat,以及与 Keras fit 的对接。
__getitem__ 随机访问、DataLoader 的 batch_size=64、shuffle、num_workers,以及自定义 collate。
同一意图的 API 对照表、预处理放哪、何时不必上复杂管线。
3.1 读完应能画出链条顺序:切片、map、shuffle、batch、prefetch。顺序写反是本章唯一的「硬错误」。3.2 读完应能默写 __len__/__getitem__ 和「getitem 没有批次维」。3.3 把两套默写对上表。不要在 3.1 里幻想 DataLoader,也不要在 3.2 里把 Dataset 叫成 tf.data。名字混用一天,排查时就会把 num_workers 的报错当成 prefetch 的问题。
小数据上允许你在 Keras 里偷懒直接 fit(x, y),但偷懒前请至少用 take(1) 或 next(iter(loader)) 打印过一批。打印是管线的单元测试。没有这一个单元格,第 5 章的 min/max 检查会变成考古。
先会一条 TF 管线、再会一条 PT 管线,对照才有实物。反过来先看对照表容易变成背单词。
| 节 | 实物 |
|---|---|
| 3.1 | 可 take 的 Dataset 链条 |
| 3.2 | 可 next 的 DataLoader |
| 3.3 | 六行验收单 |
3.1 tf.data ──► 3.2 DataLoader ──► 3.3 对照取舍
fit 或 for data in loader;第 5 章用 Keras 内置与 torchvision 两套 Fashion MNIST 入口本章故意先把两条管线分开写,而不是一上来对照表。原因是两边的对象个数不同:TF 是一条 Dataset 链条,PT 是 Dataset 类加 DataLoader 类。没摸过链条上的 prefetch,对照表里的「重叠准备」就是空词;没写过 __getitem__,就不知道为什么循环里第一维突然变成 64。3.1 与 3.2 的任务是各留下一个能迭代的实物,3.3 才允许你说「原来是翻译」。
Fashion MNIST 很小,严格来说可以直接把数组丢给 fit,PT 用 TensorDataset 包一下也行。我们仍走完整管线,是为了把打乱位置、预处理位置、Windows 上 num_workers 这些会在真项目里咬人的点,提前放进肌肉记忆。小数据上 prefetch 的收益可以忽略,但顺序错了(先 batch 再 shuffle)在大数据上会变成隐蔽的数据泄漏或打乱不充分。本章把顺序当纪律,不把速度当成绩。
读完三节,你应能回答:训练集要打乱、测试集不要;预处理放 map/transform 不放训练循环;批次大小两侧必须写成同一个数才能比曲线。答得上来,第 5 章的加载才不会各玩各的。答不上来,先别实现 CNN。
管线对照还有一层「速度观感」要提前拆穿。Fashion MNIST 的 28×28 拷贝便宜,GPU 上的 MLP 也小,你几乎感觉不到 prefetch 或 workers 的好处,有时开多进程反而更慢。不要因此下结论说「tf.data 没必要」或「DataLoader 很重」。换成解码 JPEG、做随机裁剪、批次大到占满总线时,同一套顺序才会显示价值。本章用小数据练的是顺序与接口,不是吞吐量竞赛。把这个预期写进笔记,第 5 章才不会有人因为「开了 8 个 worker 更慢」而把环境判为损坏。
和第 2 章的接缝是设备:管线产出的批次默认在 CPU。Keras fit 会帮你搬;PT 循环必须 .to(device)。接缝写错时,报错发生在训练第一步,看起来像模型问题。排查顺序应是:先打印批次 shape 与 device,再打印模型参数 device,两者一致才查层定义。本章结束时,你应能独立完成这个三行打印,而不是等到 4.3 的循环模板里才第一次见到 .to。
管线章的交付物是两条能迭代的实物,加上一张翻译表。实物意味着你已经用 take 或 next iter 打印过一批的形状、类型、最小最大值和标签。没有打印记录,第 5 章的对照数字不受理。翻译表意味着你能指出 shuffle 对 shuffle、batch 对 batch_size、prefetch 对 workers。口诀「切变洗叠等」只属于 tf.data,不要套到 DataLoader 上生造第五个对象。小数据上感觉不到加速是正常的,本章练的是接口与顺序,不是吞吐量。
和第 4 章交接时,请把「第一批已经打印」当作门票,而不是把「我写了 Dataset 类」当作门票。类可以写错返回值,打印不会说谎。Keras 一侧门票是 take(1),PyTorch 一侧是 next(iter(loader))。两张门票都要有 shape、dtype、min、max、labels。缺 min max,第 5 章会把 Normalize 差异说成框架差异。缺 labels,会把 one-hot 误喂给稀疏交叉熵。本章篇幅不长,纪律却最长:顺序、元素无批次维、训练打乱测试不打乱。把纪律执行到打印为止,管线对照才算结束。
本章不引入第三种数据库。不要把表格型读取、数据库游标、云存储列表写进入门对照,那会把翻译表从六行变成三十行。六行够用:样本、预处理、打乱、组批、重叠、接入训练。多出来的行通常是特定存储的方言,留给你自己的项目。对照课把方言挡在门外,是为了让 tf.data 与 DataLoader 的对应关系能一次看清。看清之后,方言只是 map 或 getitem 内部的实现替换,不会再动摇对象模型。