
动手写 PyTorch 的数据管道之前我一直觉得 Dataset 和 DataLoader 是两个文档里绕不开、但很少有人讲透的东西。直到我自己被内存炸掉过几次、被训练速度卡到崩溃之后才真正体会到它们就是训练流程里的“数据传输带”——一端连接着你硬盘上的原始数据另一端连接着 GPU 的显存和模型的反向传播。这篇文章我会从这两者的分工逻辑出发手写一个自定义 Dataset再把 DataLoader 各个参数背后的坑一个个踩给你看最终让你不用再靠网上零散的代码片段拼凑数据流程。1. Dataset 与 DataLoader 的分工逻辑1.1 为什么不能把所有数据一次性塞进内存任何一个做过深度学习的人最开始的直觉都是把所有图片、文本加载成一个 numpy 数组然后直接 for 循环丢给模型。小数据集没问题但一旦数据量到了几十 GB你立刻会发现三个问题。第一个是内存暴涨。假设一张 224x224 的 RGB 图片转成 float32 后就是约 150KB一万张就是 1.5GB。如果做数据增强还要生成多份副本内存很快就不够用。第二个是训练和预处理互相干扰。如果你在主进程里先做归一化、再做裁剪、再转 tensor这些操作会让 GPU 在等待 CPU 处理数据利用率直线下降。第三个问题更隐蔽你无法做随机打乱和分批次采样。全部数据都放在内存里想每轮重新打乱顺序必然要复制一份数组时间和内存都浪费。PyTorch 给出的解法是Dataset 负责“定义数据的组织方式”DataLoader 负责“高效地取数据”。这个分工让两者各自做好自己的事——Dataset 不需要关心 batch、不需要关心进程、不需要关心设备它只回答两个问题这个数据集有多长给定第 i 条数据它的训练样本和标签是什么1.2 DataLoader 补上了哪三块关键短板真正让数据能“流”起来的是 DataLoader。它在 Dataset 之上封装了三个核心能力这也是我后来手动实现数据管道时才明白的。第一是批量组装。每次给你拼好一个 batch让 GPU 可以批量计算而不是一条条喂。第二是多进程预取。用子进程提前把未来几个 batch 的数据从硬盘读出来放进内存队列GPU 算完当前 batch 时下一个 batch 已经在内存里等着了。第三是随机化与流式控制。通过 shuffle 控制洗牌时机通过 sampler 控制每条样本的出现频率通过 drop_last 控制最后一个不完整 batch 的处理方式。用一个生活化的类比Dataset 是你的书架——它知道自己有多少本书也告诉你第 n 本是什么。DataLoader 则是你的助手——他按你的要求一次拿 32 本出来、顺序打乱、并提前把后面几批书从仓库运到桌边。你只关心每次从助手手里接下 32 本书就可以。2. 手写自定义 Dataset三步搞定你的专属数据格式2.1init、len、getitem三个方法缺一不可大多数教程只让你照着模板抄但不解释这三个方法各自的位置。我建议你把init只用来记录“元信息”——比如文件路径列表、标签表、一个 CSV 的引用千万别在这里把所有图片读进内存。因为init只在创建 Dataset 时执行一次而getitem会在每个 epoch 被调用 N 次。如果你在init里做了太重的预处理后续想改一个参数就要全部重新加载非常被动。len要返回总样本数。它决定了 DataLoader 的 epoch 长度也决定了 len(loader) 返回多少个 batch。很多人忽略了这个方法结果数据库明明有一万条数据却只跑了一个 epoch 就提前结束——因为 len() 默认返回 0。getitem是真正干活的地方。它接收一个整数索引返回 (样本, 标签) 元组。你要在这里写清楚“给定第 i 条如何把原始数据变成张量”。这一步不要只 return 一个 numpy 数组而要转成 torch.tensor并且把维度、类型都确认好。通常的做法是def __getitem__(self, idx): img_path self.paths[idx] image Image.open(img_path).convert(RGB) image self.transform(image) label self.labels[idx] return image, label2.2 一个可直接运行的文本分类 Dataset 案例光说理论容易飘我直接写一个真实可跑的文本分类 Dataset处理的是 CSV 格式的新闻标题与情感标签。import pandas as pd import torch from torch.utils.data import Dataset from transformers import BertTokenizer class TextClassificationDataset(Dataset): def __init__(self, csv_path, max_len128): self.data pd.read_csv(csv_path) self.tokenizer BertTokenizer.from_pretrained(bert-base-uncased) self.max_len max_len self.label_map {negative: 0, positive: 1} def __len__(self): return len(self.data) def __getitem__(self, idx): row self.data.iloc[idx] text str(row[text]) label self.label_map[row[label]] encoded self.tokenizer( text, truncationTrue, paddingmax_length, max_lengthself.max_len, return_tensorspt ) input_ids encoded[input_ids].squeeze(0) attention_mask encoded[attention_mask].squeeze(0) return input_ids, attention_mask, torch.tensor(label, dtypetorch.long)划几个重点我在这里每次调用 tokenizer 都会重新分词这个成本其实是偏高的。更好的做法是在init里预先 tokenize 一遍把 input_ids 存成列表。但这会让内存开销变大属于用空间换时间。实际项目中如果你的机器内存够大强烈建议预先把 token 化结果缓存到内存训练速度能提升两倍以上。2.3 map-style 与 iterable-style选错内存就白费了PyTorch 的 Dataset 其实分两种。上面写的是 map-style它通过 idx 随机访问任意第 i 条样本。只要你的数据可以按索引定位永远优先选它因为它天然支持 shuffle、sampler 和多进程 worker。iterable-style 需要实现iter而不是getitem像流水线一样依次产出数据。它适合数据无法随机访问的场景比如实时流式日志、数据库游标、网络爬虫抓取。但注意iterable-style 的 shuffle 支持很弱多进程时每个 worker 会拿到同一个迭代器的复制品你需要自己写 worker 间数据切分逻辑。我的经验是除非你的数据源真的无法落盘否则不要轻易用 iterable-style。它带来的采样控制麻烦远超它省下的那点内存。3. DataLoader 参数深挖默认值背后的真实含义3.1 num_workers不是越大越快这是一个被误解最深的参数。num_workers0 表示数据加载在主进程完成简单但慢。num_workersn 表示开 n 个子进程并行加载数据通过队列传给主进程。理论上 workers 越多读取越快但实际有两个硬约束第一个是 CPU 核数第二个是 I/O 瓶颈。如果你做的是硬盘读取密集型任务比如图片解码加 workers 明显有效。如果你的数据已经在内存里比如纯随机生成的 numpy 数组workers 再多也只是增加进程切换开销。我试过在 32 核服务器上把 num_workers 从 8 调到 32结果训练速度反而变慢——因为每个 worker 都要从内存队列取数据进程间通信成了瓶颈。一个相对靠谱的调参起点是num_workers 设为 CPU 物理核心数的一半或四分之一然后观察 GPU 利用率。如果 GPU 利用率一直在 80% 以下并且 CPU 没有跑满再逐步加。另外在 Windows 系统下多进程 worker 需要把数据加载代码放到if __name__ __main__:保护块里否则会无限递归报错Linux 下则没有这个问题。3.2 batch_size、shuffle、pin_memory 是怎么协同工作的batch_size 决定每次送入 GPU 的样本量它直接影响显存占用和梯度估计的稳定性。很多人一味调大 batch_size却发现显存爆了。这里要理解一个链条数据先被 loader 拼成 batch tensor然后通过.to(device)传到 GPU。如果你用了 pin_memoryTrueDataLoader 会先把数据放到锁页内存中之后从 CPU 到 GPU 的拷贝速度会明显更快因为锁页内存能被显卡驱动直接 DMA 访问而不需要先复制到中间缓冲。shuffle 的作用大家知道是每个 epoch 开始前打乱顺序但注意它和 sampler 是互斥的两者不能同时设置。我的习惯是训练集 shuffleTrue验证集 shuffleFalse这样验证时每次看到的数据顺序完全一致指标可比性更强也更方便保存预测结果。有一点容易被忽略shuffleTrue 时打乱的是索引列表而不是数据本身。也就是说每个 epoch 生成一个新的随机排列然后按这个排列逐个访问getitem。这保证了同一个 batch 内部的样本不会总是来自数据集某个固定区间对训练稳定性很重要。3.3 sampler 与 drop_last处理不均衡数据的关键如果你处理的是类别极度不平衡的分类问题光靠 shuffle 是不够的。这时候要用 WeightedRandomSampler。它的核心是给每个样本一个采样权重比如少数类样本权重更高、多数类样本权重更低然后按权重做带放回采样让每个 batch 里各类别的期望比例更均衡。实现方式如下from torch.utils.data import DataLoader, WeightedRandomSampler labels dataset.get_all_labels() class_counts torch.bincount(labels) weights 1.0 / class_counts[labels] sampler WeightedRandomSampler(weights, num_sampleslen(weights), replacementTrue) loader DataLoader(dataset, batch_size32, samplersampler)这里 replacementTrue 表示允许同一个样本在同一个 epoch 中被抽到多次这样才真正打破了原始数据的分布。还有一个细节WeightedRandomSampler 的 num_samples 你可以设置为一个自定义值比如想要每个 epoch 恰好采样 2000 个样本就传 2000loader 的 epoch 长度就不再由 len(dataset)//batch_size 决定了。drop_lastTrue 会在最后一个 batch 不完整时直接丢弃。如果你做的是 batch normalization并且 batch_size 比较小这个参数很关键因为 BatchNorm 在小 batch 上的统计量非常不稳定。我一般习惯训练集 drop_lastTrue保证所有 batch 大小一致省去很多模型内部维度不匹配的奇怪 bug验证集 drop_lastFalse最大化样本覆盖。4. 性能优化实战与问题排查4.1 数据加载慢如何判断瓶颈在哪训练时 GPU 空转是最常见的问题。我排查性能瓶颈有一套自己的顺序按这个顺序走基本能快速定位是哪一层的锅。首先用 nvidia-smi 看 GPU 利用率。如果利用率经常在 0%-30%说明 GPU 在等数据。此时看 CPU 占用率。如果 CPU 跑满了说明是数据预处理太慢优先检查getitem里的耗时操作比如是否有重复的解码、缩放、类型转换。如果 CPU 没跑满但 GPU 仍然空闲那可能是主进程和 worker 之间的数据拷贝太慢或者 batch_size 太大导致单个 batch 传输时间远大于计算时间。第二步是看 DataLoader 的耗时。你可以单独测一下一个完整 epoch 的数据加载时间而不让模型参与from torch.utils.data import DataLoader loader DataLoader(dataset, batch_size32, num_workers8) start time.time() for batch in loader: pass print(fData load time: {time.time() - start:.2f}s)这个测试很有用。如果数据加载本身就要 30 秒而你的模型一个 epoch 只算 10 秒那瓶颈肯定是数据端。此时优先优化getitem、增加 num_workers、使用 pin_memory或者考虑在后面加缓存层。还有一种常见情况是数据增强太重。我在一次图像分类任务里用了随机裁剪旋转颜色抖动num_workers8 都扛不住。最后的解法是把数据增强改成 CUDA 上运行——用 torchvision.transforms 中的 GPU 版本或者直接用 DALI。但如果你是新手先用最简单的策略做一个缓存版本的 Dataset把增强后的结果按索引存到内存字典里第二次访问就直接返回缓存。4.2 常见报错与解决方案速查我整理了一下自己实战中经常遇到的几个 DataLoader 相关报错每一条都是踩过的坑。第一个是RuntimeError: DataLoader worker (pid(s) X) exited unexpectedly。这多半是getitem里抛了异常但异常发生在子进程里无法正常回传。排查方法是把 num_workers 设为 0让代码在主进程跑一遍看到完整 traceback 再针对性修改。常见原因包括文件路径不存在、图片损坏无法解码、数据中出现 NaN。第二个是IndexError: index out of range。这通常是len返回的数大于getitem能接受的最大索引。比如你按文件列表长度返回了 len但列表里有空行或者某些文件被过滤掉了导致实际可取的数据比 len 少。解决方法是在init里做完过滤后用self.paths [p for p in self.paths if os.path.exists(p)]然后确保__len__ len(self.paths)。第三个是ValueError: Expected input batch_size to match target batch_size。原因是某个 batch 里最后一条数据维度不一致通常出现在 collate_fn 的默认行为搞不定变长输入的情况。解决方案是自定义 collate_fn把不同长度的文本补齐到当前 batch 最大长度或者对图像做 pad 操作。def collate_fn(batch): input_ids [item[0] for item in batch] attention_masks [item[1] for item in batch] labels [item[2] for item in batch] input_ids torch.nn.utils.rnn.pad_sequence(input_ids, batch_firstTrue) attention_masks torch.nn.utils.rnn.pad_sequence(attention_masks, batch_firstTrue) labels torch.stack(labels) return input_ids, attention_masks, labels第四个是MemoryError或CUDA out of memory。如果发生在数据加载阶段很可能是你的init把所有数据塞进了内存。如果是训练到一半才爆显存要看 batch_size 是否过大、模型是否多层保留梯度。此时试试在 loader 中加入pin_memoryTrue并配合non_blockingTrue的.to(device)这通常能减少 CPU 侧的临时占用。4.3 进阶自定义 collate_fn 才真正决定数据形态很多教程只守在getitem层面但实际工作中 collate_fn 才是决定数据最终形态的地方。默认的 collate_fn 会做三件事把列表中的 tensor 堆叠、把数字转成 tensor、把无法堆叠的数据做成列表。当你的样本是变长文本、多模态特征、或者要输出多个目标时默认行为根本不够用。我写过一次目标检测的数据管道每张图的标注数量不同默认 collate_fn 会直接报错。自定义 collate_fn 的职责就是把它拼成模型需要的结构——把图像 stack 成 [B, C, H, W]把标注做成一个带 batch 维度的列表或 padded tensor。这对模型前向传播的输入适配至关重要。还有一些时候我故意在 collate_fn 里做数据增强或 Mixup而不是放在getitem里。因为我希望增强操作能同时看到整个 batch 的样本比如用 batch 内其他样本的标签做插值。这就是“batch-level augmentation”的思路比逐样本增强更灵活也能减少重复计算。5. 亲测有效的数据加载性能优化清单实践下来下面这组配置是我的通用起点适合大部分 CV / NLP 任务。当然你还是要根据实际数据源微调但至少不会出大错。loader DataLoader( dataset, batch_size64, shuffleTrue, num_workers8, pin_memoryTrue, drop_lastTrue, persistent_workersTrue )persistent_workers 是 PyTorch 1.8 之后引入的参数我强烈建议开启。它的作用是让 worker 进程在跑完一个 epoch 后不退出而是继续存在于内存中等待下一个 epoch。如果不开启每个 epoch 结束都会销毁 workers 再重新创建这个开销在小数据集上非常明显。我实测过一个四万样本的数据集开启 persistent_workers 后每个 epoch 的切换时间从 8 秒降到了 1 秒左右。如果你用的是显存很大的 GPU还想进一步压缩 CPU 侧的预处理时间可以考虑两个方向的扩展。第一个是把预处理好的 tensor 直接存成硬盘上的二进制格式训练时只需要做反序列化省去图片解码开销。第二个是引入内存映射比如用 lmdb 或 h5py 存储数据配合 num_workers 读取可以显著降低小文件读写的性能损耗。这两个方案我都在项目中试过lmdb 版本比原版读写快了大约 2.5 倍。在配置 transformer 模型的数据流时还有一点值得提醒HuggingFace 的 Dataset 对象本身带.set_format(torch)方法你可以直接把它传给 DataLoader。但内部实现是会把 Batch 转成 dict如果你的模型需要多个输入字段这种形式反而更方便。不过要注意不要再用默认 collate_fn 去处理已经是 dict 的数据我建议显式传一个能把各字段分别 stack 的 collate_fn避免 PyTorch 版本升级后的行为差异。说实话Dataset 和 DataLoader 这两个类看过文档的人多真正用好的人少。我踩过最大的坑就是在init里做重活、或者盲目堆 num_workers结果内存动不动就上 30GB训练反而更慢。后来我慢慢形成了一套自己的习惯init只存路径getitem只做单样本变换collate_fn 负责批量整合sampler 控制采样分布pin_memory 和 persistent_workers 无脑开。这套组合拳打下来我的训练速度基本都能稳定压满 GPU 利用率也再没遇到数据加载导致的莫名其妙的崩溃。最后再分享一个调试小技巧当你的训练结果出现奇怪的随机性时先别急着改模型试着把 DataLoader 的 shuffle 关掉或者固定随机种子。很多时候所谓“模型不收敛”其实只是数据顺序和归一的统计量在捣乱。数据和模型是同一张桌上的两个角色只有数据传输带顺畅了模型才能真正跑起来。