☰
PyTorch Dataset与DataLoader完全指南:手写实现与性能优化实战
2026/10/6 8:28:57 网站建设 项目流程

动手写 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, label

2.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_len=128): 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, truncation=True, padding='max_length', max_length=self.max_len, return_tensors='pt' ) input_ids = encoded['input_ids'].squeeze(0) attention_mask = encoded['attention_mask'].squeeze(0) return input_ids, attention_mask, torch.tensor(label, dtype=torch.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_workers=0 表示数据加载在主进程完成,简单但慢。num_workers=n 表示开 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_memory=True,DataLoader 会先把数据放到锁页内存中,之后从 CPU 到 GPU 的拷贝速度会明显更快,因为锁页内存能被显卡驱动直接 DMA 访问,而不需要先复制到中间缓冲。

shuffle 的作用大家知道,是每个 epoch 开始前打乱顺序,但注意它和 sampler 是互斥的,两者不能同时设置。我的习惯是:训练集 shuffle=True,验证集 shuffle=False,这样验证时每次看到的数据顺序完全一致,指标可比性更强,也更方便保存预测结果。

有一点容易被忽略:shuffle=True 时打乱的是索引列表,而不是数据本身。也就是说,每个 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_samples=len(weights), replacement=True) loader = DataLoader(dataset, batch_size=32, sampler=sampler)

这里 replacement=True 表示允许同一个样本在同一个 epoch 中被抽到多次,这样才真正打破了原始数据的分布。还有一个细节:WeightedRandomSampler 的 num_samples 你可以设置为一个自定义值,比如想要每个 epoch 恰好采样 2000 个样本,就传 2000,loader 的 epoch 长度就不再由 len(dataset)//batch_size 决定了。

drop_last=True 会在最后一个 batch 不完整时直接丢弃。如果你做的是 batch normalization,并且 batch_size 比较小,这个参数很关键,因为 BatchNorm 在小 batch 上的统计量非常不稳定。我一般习惯:训练集 drop_last=True,保证所有 batch 大小一致,省去很多模型内部维度不匹配的奇怪 bug;验证集 drop_last=False,最大化样本覆盖。

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_size=32, num_workers=8) start = time.time() for batch in loader: pass print(f'Data load time: {time.time() - start:.2f}s')

这个测试很有用。如果数据加载本身就要 30 秒,而你的模型一个 epoch 只算 10 秒,那瓶颈肯定是数据端。此时优先优化getitem、增加 num_workers、使用 pin_memory,或者考虑在后面加缓存层。

还有一种常见情况是数据增强太重。我在一次图像分类任务里用了随机裁剪+旋转+颜色抖动,num_workers=8 都扛不住。最后的解法是把数据增强改成 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_first=True) attention_masks = torch.nn.utils.rnn.pad_sequence(attention_masks, batch_first=True) labels = torch.stack(labels) return input_ids, attention_masks, labels

第四个是MemoryError或CUDA out of memory。如果发生在数据加载阶段,很可能是你的init把所有数据塞进了内存。如果是训练到一半才爆显存,要看 batch_size 是否过大、模型是否多层保留梯度。此时试试在 loader 中加入pin_memory=True并配合non_blocking=True的.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_size=64, shuffle=True, num_workers=8, pin_memory=True, drop_last=True, persistent_workers=True )

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 关掉或者固定随机种子。很多时候所谓“模型不收敛”,其实只是数据顺序和归一的统计量在捣乱。数据和模型是同一张桌上的两个角色,只有数据传输带顺畅了,模型才能真正跑起来。

需要专业的网站建设服务?

联系我们获取免费的网站建设咨询和方案报价,让我们帮助您实现业务目标

立即咨询