本节摘要:PyTorch 把数据访问拆成两层。
torch.utils.data.Dataset负责「第 i 条样本是什么」,你要实现__len__和__getitem__。DataLoader负责组批、打乱、多进程拉取,原文 Fashion MNIST 示例使用batch_size=64、训练shuffle=True、测试shuffle=False。自定义数据集时,这一对抽象比把所有预处理写进训练循环更可测。Windows 上num_workers>0出问题,原文建议先改回 0。
阅读完本节,你应当能够:
ToTensor 与 Normalize((0.5,), (0.5,)) 在 torchvision 路径上的作用num_workers 与 pin_memory 的取舍,而不把它们当默认必开训练循环不想知道图片在磁盘哪、标签在表格哪一列。它只想 for inputs, labels in loader。Dataset 把随机访问藏起来:给整数索引,返回一条已经变成张量的样本。DataLoader 再用这些随机访问拼批次、打乱索引、可选地开子进程。原文用这段话定义分工:Dataset 访问样本与标签,DataLoader 批次、打乱、并行。
内置 FashionMNIST 已经是 Dataset。你仍要理解接口,因为换自己的数据时没有内置类。最小实现:在 __init__ 里持有数组或文件列表,__len__ 返回样本数,__getitem__ 读第 i 条、做转换、返回元组。返回值会被 collate 函数堆叠;默认 collate 能堆张量和数字标签。变长序列需要自定义 collate,入门图像尺寸一致,用默认即可。
__getitem__ 里做重 I/O 可以,但要接受:num_workers=0 时这发生在主进程,和前向串行。开 workers 后,每个子进程执行 getitem,主进程组批。这和 tf.data 的 map 并行是同一意图、不同实现:一边是图运行时拉数据,一边是 Python 多进程。Windows 的多进程启动方式与 Linux 不同,fork 不那么干净,所以原文点名:遇到问题把 num_workers 设为 0。这是 SOURCE 里的平台经验,不是过时谣言。
⚠️ 常见坑:在 Dataset 里返回 NumPy 却期望自动上 GPU。DataLoader 默认不帮你
.to(device)。设备迁移仍在循环里做,这是和第 2.1 节一致的纪律。
原文 PyTorch 入门用 transforms.Compose:ToTensor 把图像变成通道在前的张量且缩放到 0–1;Normalize((0.5,), (0.5,)) 对单通道做减均值除标准差,灰度只有一个通道所以元组长度为 1。然后 FashionMNIST(..., train=True/False, download=True, transform=transform)。download=True 表示缺失时获取数据;数据根目录由参数指定,本教程不写具体磁盘路径。
DataLoader 构造:训练 batch_size=64, shuffle=True;测试 64、shuffle=False。测试打乱通常无益,还破坏你想按固定顺序看错例的需求。drop_last 默认 False。迭代时 for i, data in enumerate(train_loader, 0): inputs, labels = data,这是原文训练循环的取数写法。enumerate 从 0 起只为了每 100 步打印。
| 参数 | 原文用法 | 含义 |
|---|---|---|
batch_size |
64 | 每批图像数 |
shuffle |
训练 True,测试 False | 是否打乱索引 |
num_workers |
示例可 2,Windows 先 0 | 子进程加载 |
pin_memory |
未作为入门必选项 | 锁页内存加速 CPU 到 GPU |
transform |
ToTensor + Normalize | 在 getitem 时执行 |
# 概念:Dataset 随机访问 + DataLoader 组批 # for i in range(len(dataset)): # x, y = dataset[i] # loader = DataLoader(dataset, batch_size=64, shuffle=True, num_workers=0) # inputs, labels = next(iter(loader)) # inputs.shape 预期类似 64, 1, 28, 28
类别名字元组原文列出十个英文服装名,仅用于打印预测,不进入 Dataset 的标签张量。标签仍是整数。自定义 Dataset 若返回字符串标签,默认 collate 会让你难受,在 getitem 里转成整数。
__getitem__ 应当无副作用地可重复:多次访问 i 在无随机增强时应相同。有随机翻转时,每次不同是故意的。把随机种子固定在 Dataset 内部要小心多进程:每个 worker 需要独立种子,否则增强相关性上升。入门关闭随机增强,对照更干净。
长度必须准确。__len__ 返回 60000 而实际只有 59999,最后一个 epoch 可能越界。反过来偏短,等于悄悄丢数据。用内置 FashionMNIST 时这不是你的锅;自己读列表时要用同一长度源。
pin_memory=True 在有 GPU 时常常值得开:主机内存锁页,复制到显卡更快。没有 GPU 时无意义。persistent_workers 减少每个 epoch 重启进程的开销,workers 为 0 时无关。不要一上来抄网上「八 workers + pin」的配置,先 0 worker 跑通形状,再加。
collate 默认把列表里的第 0 维叠起来。样本已是 (1,28,28),批次 (64,1,28,28)。若你返回已经带批次维的张量,默认 collate 会叠出五维,训练立刻形状爆炸。getitem 必须是单样本。这和 tf.data 先单样本再 batch 是同一纪律。
💡 关键直觉:Dataset 永远返回一条;只有 DataLoader 才出现批次维。谁在 getitem 里
unsqueeze出批次,谁就要自己写 collate。
与优化器循环的接缝:inputs, labels = data 之后立刻 .to(device),与模型同设备。不要在 Dataset 里写死 .cuda(),否则 CPU 机器无法跑,也无法把同一 Dataset 用于多设备实验。转换在 CPU 做、批次在循环里搬家,是更常见的分工;有人用 pin_memory + 非阻塞传输再抠延迟,那是模型已经很大之后的事。
IterableDataset 是另一种接口:不能按索引访问,只能顺序迭代,适合流式数据。入门用 map-style(len + getitem)即可。不要为了「更高级」把 Fashion MNIST 改成流。
调试手段:next(iter(loader)) 看一批。打印 inputs.min()、max() 验证 Normalize 后大约在 -1 到 1,而不是还在 0–255。若误把已经 ToTensor 的数据再除以 255,对比 TF 侧除以 255 的实验会不公平,第 5 章会再强调。
多进程下的随机性、共享文件句柄、在 Dataset 里再创建 DataLoader(套娃)都是事故源。Windows 上还要求自定义 Dataset 的代码处在可导入的模块路径,且入口保护好,否则子进程重新执行训练脚本会递归爆炸。原文「先设 0」是对这条事故链的最短止血。对照课把止血当成正确流程的一部分,不是示弱。
标签与图像必须按同一索引对齐。常见做法:一个列表存路径,一个列表存整数标签,__getitem__ 用同一个 i 去取。错位一整类时,损失会下降(网络仍能学那个错的对应),测试时你用正确名字去读,会觉得模型「蠢得离奇」。所以 5.1 要求用类别名肉眼看几张图。转换放在 __init__ 传入的 transform 上,而不是写死在类内部,便于训练用增强、测试用确定性转换,两个 Dataset 实例共享同一读盘逻辑。
Subset 可以切验证集而不复制数据。随机切分时固定生成器,否则每次启动验证样本都变,曲线不可比。random_split 返回的子集仍走原 Dataset 的 getitem,transform 也是同一套——若训练增强写在原 Dataset 里,验证子集也会被增强。这是把 transform 按实例传入的理由:训练子集与验证子集应是两个 Dataset 或两套 transform。
默认保留。丢掉会浪费样本,Fashion 最后一批大约 60000 除 64 的余数。有些分布式或 BN 场景要求每步形状固定,才 drop_last=True。入门保留。评估时更不要丢,否则准确率的分母变了,和 Keras evaluate 对不上。对照时两侧对最后一批的策略必须相同。
collate 自定义出现在变长或「一张图多个框」的任务。返回已经 stack 好的批次时,DataLoader 会再 stack 一次。getitem 保持单样本,是避免五维张量的最简单规则。与 tf.data「先元素后 batch」对齐记忆。
ToTensor 不只做缩放,还把通道维放到前面,并且把 PIL 图像变成浮点张量。若你的自定义 Dataset 读进来已经是 NumPy 的 (28,28) 且值在 0–255,不要再套 ToTensor 一次,否则类型和维度会错。应自己除以 255、unsqueeze(0) 加通道。torchvision 的 FashionMNIST 返回 PIL,所以原文的 Compose 是对的。两种入口不要混抄。Keras load_data 给的是 NumPy,对应除以 255,没有 PIL 这一步。对照时「加载入口」必须写在卡片上,因为它决定转换链。
num_workers>0 时,Windows 要求训练入口受保护,否则子进程重新执行整份脚本,出现嵌套的 DataLoader 和重复下载。止血仍是先改回 0。Linux 上 workers 2 往往无感提升,因为 28×28 太小。不要从博客抄 num_workers=8, pin_memory=True 当人格设定。有 GPU 再考虑 pin_memory;没有 GPU 开它无益。把这两项当成性能旋钮,与学习率分开调,且只在形状已经正确之后。
最后一批不足 64 时,BN 会发出统计不稳定的抱怨,MLP 没有 BN 则无事。这是「入门不加 BN」的又一条理由:少一条与批次大小耦合的行为。等你加 CNN+BN,再回头看 drop_last。
__len__ 与 __getitem__;返回张量与整数标签num_workers 出问题先 0DataLoader 验收:getitem 返回单样本,loader 之后才出现 64 那一维;训练 shuffle 开,测试关;Windows 上 workers 先 0;设备迁移在循环里做。Normalize 是否使用必须写进卡片,因为它改 min max。类别名元组不进 Dataset 返回值。len 必须等于真实样本数。自定义 list 当层会在 4.2 翻车,自定义 Dataset 返回已带批次维会在这里翻车,两件事都是「把容器职责做进了元素」。记住元素永远是一条,容器才组批。与 tf.data 的元素对批次关系完全同构,只是对象拆成了两个类。
多进程下的随机增强会让每个 worker 若共用同一种子则做出相关翻转,验证波动变小却是假象。入门关闭增强就避开了。打开增强时,为每个 worker 设不同种子。Windows 上 besides workers=0,还要避免在 getitem 里再创建新的 DataLoader。套娃会递归爆炸。pin_memory 只在有 GPU 且批次确实要搬到显卡时有意义。persistent_workers 在 workers 为 0 时无效。把这些参数看成性能旋钮,与学习率分栏。形状未正确之前禁止拧性能旋钮。getitem 里不要写死 cuda,否则 CPU 机器无法跑同一 Dataset。转换在 CPU 完成,搬家在循环完成,和 tf.data 的 map 在 CPU、fit 内部再喂设备是同一分工。记住分工,3.3 的翻译表才不是单词表。
把 Windows 上 workers 先 0 写成默认,而不是写成失败后的退路。默认 0 能跑,再尝试 2,不行就回 0。不要默认 8。getitem 保持无副作用可重复,除非随机增强是故意的。len 用同一数据源计数。Subset 切验证时 transform 按实例传入,避免验证被训练增强污染。最后一批保留,评估更不要 drop_last,否则和 Keras evaluate 分母不同。这些细则看起来碎,对照时每一条都能单独制造「框架不稳定」的谣言。碎,但是闸门。闸门都关上,3.3 的验收单才打得下去。
把 Dataset 当成随机访问的抽屉柜,DataLoader 当成每次抽出若干抽屉并捆成一扎的工人。工人不会进抽屉里改衣服,搬家到 GPU 是循环的事。抽屉柜不知道批次。这个比喻用完即弃,但能挡住 getitem 里 unsqueeze 出批次维的手。挡住了,形状契约才能传到 5.1。传不到,Fashion 的 64 会变成 1 或 5 维,对照从第一行就作废。
审查第一批时同时看 labels.dtype。不是整数就停。停比继续训便宜。继续训会把错误编码学进权重,检查点也会脏。脏检查点不能当基线。基线必须从干净批次产生。干净的定义写在 3.3 验收单上,执行在本节的 next(iter)。执行了,抽屉柜比喻可以忘掉。没执行,比喻只是装饰。装饰挡不住 5 维张量。
本节毕业标准是:你能在不看稿的情况下说出 getitem 没有批次维、Loader 才有、Windows 先零 worker、设备在循环里迁。四句能说完,3.3 的翻译才有实物。说不全,回去打印第一批。打印是毕业考,比喻不是。抽屉柜只为挡住五维。挡住之后,批次维只允许出现在 Loader 返回值上,不允许出现在任何自定义 getitem 的返回值上。这是本节最后一条法律。法律优先于任何性能旋钮。旋钮可以明天再拧。今天只允许把形状拧对。
下一节把 tf.data 与 DataLoader 放进同一张取舍表:预取、打乱、预处理放哪、何时直接用内存数组。