ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

PyTorch Dataset类实战指南:核心方法、DataLoader协作与常见坑解析

PyTorch Dataset类实战指南:核心方法、DataLoader协作与常见坑解析 PyTorch里最容易被新手玩坏的就是Dataset类。很多人写了两三行就跑起来结果碰到点奇怪的数据就卡壳或者数据集一大就慢得像蜗牛。我自己刚开始学的时候也被它坑过几次所以这篇就打算把Dataset类彻底讲清楚它到底是什么、为什么非它不可、三个核心方法怎么用、不同场景的实例怎么写、跟DataLoader怎么配合以及一路踩过的坑。环境默认你已经装好了PyTorch版本2.x和1.x都适用不需要额外装任何东西。1. 为什么偏偏要用Dataset直接读数据不行吗很多人一开始都会想我直接把图片读进列表、把表格塞进numpy数组不也能训练吗确实能但那只是数据量小、场景简单的时候。一旦数据多了或者要做随机增强、多线程加载、乱序采样这些正经训练流程直接塞内存的做法就崩了。1.1 Dataset在训练链路里的真实定位先看一个标准训练循环里数据是怎么流动的for epoch in range(num_epochs): for batch_data, batch_label in dataloader: # 模型前向、反向、更新参数这里的dataloader是DataLoader实例它负责把Dataset按batch切好、打乱顺序、可能开多进程去加载。Dataset则是一个“提供单个样本”的东西。DataLoader不关心你的数据在硬盘上是什么结构它只负责调用dataset[i]拿到第i个样本然后帮你打包成batch。这个设计最大的好处是解耦。Dataset管“怎么拿到一个样本”DataLoader管“怎么把这些样本高效地喂给模型”。两件事拆开各自只需要做好自己那一摊。比如今天你数据是文件夹里一堆jpg明天换成CSV表格后天换成h5文件只需要换Dataset的实现训练代码一行都不用动。1.2 什么时候你才需要自己写Dataset不是说任何情况都要自定义Dataset。官方torchvision.datasets已经覆盖了ImageFolder、CIFAR、MNIST这些常见玩意直接用就行。但下面几种情况你就逃不掉了数据不是标准的目录结构比如所有图片在一个文件夹里标签在另一个CSV文件里一个样本对应多个输入比如同时要读图片和对应的json标注信息样本本身是序列化格式比如h5、npz、pkl需要自定义解析逻辑数据量太大没法全部塞进内存只能在__getitem__里按需读取要做特殊的样本级别预处理或者返回多任务学习里的多个标签判断标准其实很简单如果dataset[i]这种取一个样本的操作你用现成的类搞不定那就自己写。大多数比赛和工业场景都得自己来。2. 三个核心方法照着抄就行自己写Dataset类本质上就是继承torch.utils.data.Dataset然后实现三个方法。一个都不能少少了直接报错。from torch.utils.data import Dataset class MyDataset(Dataset): def __init__(self, ...): # 初始化比如读文件列表、读标签、设置transform def __len__(self): # 返回样本总数 def __getitem__(self, idx): # 根据索引idx返回一个样本2.1__init__把所有准备工作做在这里很多新手容易犯一个错就是把__getitem__当成了整个数据加载逻辑的入口所有东西都往里面塞。其实__init__才应该承担大部分重活。一般__init__里做的事情包括读取所有样本的文件路径列表读取标签文件整理成list或dict初始化transform做数据划分比如训练集/验证集分离打印一下样本总量确认数据加载正确我自己的习惯是在__init__里把文件路径和标签全部对齐形成一个self.data列表每一个元素是一个(image_path, label)元组。这样__getitem__的逻辑就极简单按idx从self.data里拿一条记录读文件做transform返回。def __init__(self, img_dir, label_csv, transformNone): self.img_dir img_dir self.transform transform self.data [] # 假设label_csv是两列filename, label import pandas as pd df pd.read_csv(label_csv) for _, row in df.iterrows(): img_path os.path.join(img_dir, row[filename]) self.data.append((img_path, row[label])) print(f共加载 {len(self.data)} 个样本)这里要强调一点文件路径核对放在__init__里做不是等到训练时才发现路径错了。我见过有人偷懒不检查路径结果训练了10分钟才报FileNotFoundError白白浪费时间。最好在__init__里抽查几个路径是否存在。2.2__len__骗谁也别骗这个__len__就一行代码返回len(self.data)就行了。千万别在这里做什么复杂计算DataLoader会频繁调用它来判断一个epoch有多少个batch。有个误区是有人觉得__len__返回的是batch数而不是样本数。不对就是返回单个样本的总数。至于一个epoch有多少个batch是DataLoader根据batch_size自己算的。2.3__getitem__核心中的核心__getitem__接收一个整数idx返回第idx个样本。这里的重点在于你想返回什么类型的数据。可以返回一个(image_tensor, label)元组——最常见一个(image, mask)元组——分割任务一个(image, label, extra_info)元组——需要额外信息时一个dict——样本本身是多种数据时比如{image: ..., label: ..., name: ...}看一个实际的图像分类例子def __getitem__(self, idx): img_path, label self.data[idx] from PIL import Image image Image.open(img_path).convert(RGB) if self.transform: image self.transform(image) return image, label关键字眼是convert(RGB)。PIL打开图片时灰度图不会自动变成三通道不转的话同一个模型输入维度不稳定训练直接炸。另外打开图片后如果要做什么尺寸调整建议放在transform里用torchvision自带的Resize而不是在这里自己用PIL去resize能省很多事。__getitem__里还有个大忌讳别在这里做太耗时的操作比如每次读取都重新解析一个大型JSON文件或者做非常复杂的预处理。__getitem__会被DataLoader以极高频调用一个epoch跑几万次。耗时操作放这里训练速度会肉眼可见地变慢。正确做法是那些跟具体样本无关的固定操作能提前做就提前做放__init__里。3. 从最简单到最实战三种Dataset写法直接抄纸上谈兵没用直接上菜。下面三个实例覆盖了大部分使用场景图像分类、表格数据、分割/多输出任务。每个我都会给完整代码和解释。3.1 图像分类文件路径加CSV标签这是竞赛和业务里最常碰到的场景图片都堆在一个文件夹标签放在CSV里。有的数据源是图片文件名就是标签更简单但大多数时候还是CSV稳妥。import os from PIL import Image from torch.utils.data import Dataset import pandas as pd class ImageClassificationDataset(Dataset): def __init__(self, img_dir, label_csv, transformNone): self.img_dir img_dir self.transform transform self.data [] df pd.read_csv(label_csv) for _, row in df.iterrows(): self.data.append((row[filename], int(row[label]))) # 抽查前3个路径是否存在 for fname, _ in self.data[:3]: fpath os.path.join(img_dir, fname) assert os.path.exists(fpath), f文件不存在: {fpath} def __len__(self): return len(self.data) def __getitem__(self, idx): fname, label self.data[idx] fpath os.path.join(self.img_dir, fname) image Image.open(fpath).convert(RGB) if self.transform: image self.transform(image) return image, label重点说两个细节。第一label转成int这个很重要。很多CSV里标签读出来是字符串比如0、1如果忘了转训练时loss计算直接类型错误。但你的标签是字符串类别名比如cat、dog那就在这里或者外面做一个类别到整数的映射。第二assert检查只做前几个就够了没必要全部检查全部检查在数据量大的时候本身也有开销。使用方式from torchvision import transforms transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize([0.485, 0.456, 0.406], [0.229, 0.224, 0.225]) ]) dataset ImageClassificationDataset( img_dirtrain_images/, label_csvtrain_labels.csv, transformtransform )如果你没有现成的CSV只是图片文件名里带标签比如cat_001.jpg、dog_002.jpg那__init__里解析文件名就够了def __init__(self, img_dir, transformNone): self.img_dir img_dir self.transform transform self.data [] for fname in os.listdir(img_dir): if not fname.endswith((.jpg, .jpeg, .png)): continue # 假设文件名格式类别_编号.jpg label_str fname.split(_)[0] label 0 if label_str cat else 1 self.data.append((fname, label))3.2 表格数据回归或分类都一样表格数据很多人习惯用pandas一把梭全部读进内存然后转numpy再转torch。这种方法在小数据集上没问题但有几个尴尬场景一是数据量大到内存吃紧二是需要在线做特征工程或数据增强三是训练集验证集需要在相同逻辑下做样本级变换。Dataset写法如下import numpy as np import torch from torch.utils.data import Dataset class TableDataset(Dataset): def __init__(self, features, targetsNone): # features: numpy数组或DataFrame # targets: numpy数组或Series可以为None预测场景 if isinstance(features, pd.DataFrame): features features.values if targets is not None and isinstance(targets, pd.Series): targets targets.values self.features torch.from_numpy(features).float() if targets is not None: self.targets torch.from_numpy(targets).float() else: self.targets None def __len__(self): return len(self.features) def __getitem__(self, idx): x self.features[idx] if self.targets is not None: y self.targets[idx] return x, y return x这里把特征直接全部转成torch.Tensor放内存好处是__getitem__几乎零开销训练时能把CPU瓶颈降到最低。有些同学还会在这里做标准化我建议标准化在外面用sklearn做不要在Dataset里做因为你还要保证验证集用训练集的均值方差来做标准化放在Dataset里容易混。3.3 分割任务图片和像素级掩码一起返回分割、检测这类任务一个样本不只是图片本身还有对应的像素级掩码或边界框标注。以分割为例import numpy as np from PIL import Image class SegmentationDataset(Dataset): def __init__(self, img_dir, mask_dir, transformNone, mask_transformNone): self.img_dir img_dir self.mask_dir mask_dir self.transform transform self.mask_transform mask_transform self.filenames sorted(os.listdir(img_dir)) # 过滤掉非图片文件 self.filenames [f for f in self.filenames if f.endswith(.png) or f.endswith(.jpg)] def __len__(self): return len(self.filenames) def __getitem__(self, idx): fname self.filenames[idx] image Image.open(os.path.join(self.img_dir, fname)).convert(RGB) mask Image.open(os.path.join(self.mask_dir, fname)) if self.transform: # 注意transform和mask_transform必须是同步的 seed torch.initial_seed() torch.manual_seed(seed) image self.transform(image) torch.manual_seed(seed) mask self.mask_transform(mask) # 掩码转成long tensor因为分割标签是类别索引 mask torch.as_tensor(np.array(mask), dtypetorch.long) return image, mask这里有个极其隐蔽但极其关键的坑图片和掩码的随机增强必须同步。如果你给图片做了随机翻转掩码也必须做同样角度的随机翻转否则数据就错位了。上面代码的技巧是先固定随机种子然后分别对image和mask做transform。但要注意如果你用的是torchvision的v2版本或者transform里用了RandomResizedCrop这种会改变尺寸的增强简单固定种子可能不够。更稳妥的做法是自己在__getitem__里实现一个同步逻辑或者用一些专门处理分割增强的库比如albumentations。用albumentations的话它天生支持image和mask同时变换一张代码一张图方便得多。4. 和DataLoader的配合几个参数直接决定训练速度Dataset写好了喂给DataLoader就完事了吗没那么简单。DataLoader有一堆参数每一个都可能让你训练变慢或者跑崩。4.1 batch_size、shuffle、drop_last怎么配batch_size大多数时候是2的幂次16、32、64。不是越大越好取决于你的显存。训练时如果OOM先把batch_size减小一半再说。shuffle训练集必须开验证集和测试集一般不开。有人为了省事全程不开shuffle结果模型在epoch之间看到的数据顺序完全一样可能会学到顺序上的伪特征。drop_last当样本数刚好不能被batch_size整除时最后一个batch可能只有几个样本。有些人希望每个epoch的batch数固定就会设drop_lastTrue。如果不设默认False意味着最后一个batch会被保留但batch大小和其他不一样可能导致模型训练的稳定性稍微受影响。我的习惯是设True省心。4.2 num_workers这个参数很多人一辈子就设为0num_workers决定用几个子进程去预取数据。默认是0意思是数据在主进程里同步加载模型在GPU上算完一批CPU才去加载下一批二者串行GPU经常闲着等数据。设成4或8之后加载下一批数据的操作会提前在另一个进程里做好模型一算完数据已经等在那边了训练速度能快好几倍。但是num_workers不是越大越好。开太高了会有两个问题一是每个worker都要复制一份Dataset的内存副本实际是按需复制内存占用暴涨二是进程切换的开销反而拖慢速度。我自己的经验法则是普通办公CPU先设216核左右的机器设6到8内存紧张时一律往小了调还有一个Windows上的坑num_workers大于0时Dataset相关代码要放在if __name__ __main__:里面否则会无限递归炸内存。Linux上没这个问题但Windows上必须注意。4.3 collate_fn当你返回的东西不是规整张量时默认的collate_fn做的事情是把__getitem__返回的各个样本在第一个维度上堆起来变成batch。比如返回的是(3, 224, 224)的图片张量那堆完就是(32, 3, 224, 224)。但如果你返回的东西包含变长序列、或者本身是dict、或者图片尺寸不统一比如没做Resize默认的collate_fn就会报错。这时候你需要自定义collate_fn。比如处理变长文本序列时通常要做paddef collate_batch(batch): images, labels zip(*batch) # 假设images已经是torch.Tensor尺寸一致只是堆成batch images torch.stack(images, dim0) labels torch.tensor(labels) return images, labels如果你__getitem__返回的是dict那collate_fn可以这样写def collate_dict(batch): return { image: torch.stack([item[image] for item in batch], dim0), label: torch.tensor([item[label] for item in batch]), name: [item[name] for item in batch] }这里有个判断标准如果你的样本里所有元素都可以直接堆叠成张量就用默认collate省事。只要有非张量元素或者元素形状不是完全一致的就老老实实自己写。5. 实操中一定会踩的坑我帮你提前填平这部分是我自己在无数轮训练里踩出来的血泪教训每一条都值得记下来。5.1 图片增强和标签必须同步尤其分割任务前面提过分割任务的同步问题这里再单独强调一遍。如果你用torchvision的Compose同时处理image和mask直接写两个transform对象会发现分别作用后图片翻转了但掩码没翻转——augmentation不同步。我自己之前做分割时就因为这个模型训练了半天指标死活上不去最后逐样本检查才发现掩码和图片错位了。一个最简单的解决方案是使用albumentations库它天然支持同步变换多张图import albumentations as A from albumentations.pytorch import ToTensorV2 transform A.Compose([ A.HorizontalFlip(p0.5), A.RandomResizedCrop(224, 224, scale(0.8, 1.0)), A.Normalize(), ToTensorV2() ]) def __getitem__(self, idx): # 读原图 image cv2.imread(img_path) # BGR image cv2.cvtColor(image, cv2.COLOR_BGR2RGB) mask cv2.imread(mask_path, cv2.IMREAD_GRAYSCALE) augmented transform(imageimage, maskmask) image augmented[image] mask augmented[mask] return image, mask.long()这个方案在语义分割、目标检测任务里非常省心。图像分类任务只用transform处理一张图就无所谓同不同步torchvision的Compose就行。5.2__getitem__里做transform还是外面做这是个问题有个经典设计问题torchvision.DataLoader的官方示例里transform都是放在Dataset里的。但有些高性能框架会把基础图像解码、尺寸调整这些常见操作放到另一个阶段去流水线化。对普通玩家来说transform放Dataset里的__getitem__做没毛病。但有一种例外情况如果你做的是多模态任务比如video加audio或者数据量大到解码成了瓶颈建议把耗时操作解码视频帧、读取并解析JSON放在__init__里预计算好或者第一次访问后缓存起来别每次训练都重新解码一遍。我当时做过一个视频分类项目每个样本要从视频文件里抽取10帧每次__getitem__都要重新用OpenCV打开视频、逐帧读取。一个epoch要重复读几万次视频文件速度慢到离谱。后来改成在__init__里提前抽好所有帧存成图片文件__getitem__变成简单的读图操作训练速度快了将近5倍。5.3 Dataset和验证集的纠缠一种很容易犯的错是把transform用在训练集和验证集上时逻辑不一致。有些人图省事训练集验证集用同一个Dataset实例结果验证时也做了随机增强评估指标忽高忽低完全不可信。正确做法是训练集和验证集要用不同的transformtrain_transform transforms.Compose([ transforms.RandomHorizontalFlip(), transforms.RandomRotation(15), transforms.Resize((224, 224)), transforms.ToTensor(), ]) val_transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), ]) train_dataset CustomDataset(..., transformtrain_transform) val_dataset CustomDataset(..., transformval_transform)随机增强是给训练集用的验证集只需要做尺寸调整和归一化保持确定性。5.4 Worker进程里跑Debug疑似内存泄漏常见现象是训练到一半内存占用不断上涨最后直接卡死或者OOM。除了模型本身和优化器的正常内存占用一个常见原因是num_workers设置过大而且每个worker都会复制一份Dataset引用。如果你在__getitem__里有大量矩阵操作或者频繁的文件读取多个worker同时跑内存翻倍增长。另外如果你在Windows上Debug时发现程序反复启动十有八九是忘记把训练代码包在if __name__ __main__里子进程递归加载了主模块。这个坑极其隐蔽报错信息又不直观网上搜出来的多半是让你把num_workers设回0治标不治本。5.5 索引错位shuffle之后索引还对不对有一个非常经典的错误在__getitem__里如果你用的idx是来自某个数据集的原始下标而你又手动实现了shuffle逻辑而不是交给DataLoader的shuffle参数很容易出现一个样本被重复读取或漏读的情况。我一直的主张是shuffle就老老实实交给DataLoader不要在Dataset内部自己做打乱。DataLoader的shuffle参数会先生成一个打乱过的索引序列然后逐个传给__getitem__这是最标准的机制。如果你在__init__里手动shuffle了self.data列表然后又开了DataLoader的shuffle结果就是样本顺序被连续打乱两次虽然不算致命错误但会让调试时想“复现同一批数据”变得很难。6. 性能调优细节从数据侧把训练速度拉满很多人以为训练慢是模型的问题其实数据加载往往是最大的瓶颈。GPU算的再快数据喂不上来也是白搭。这节分享几个我从实际项目里总结出来的数据加载调优技巧。6.1 把数据先打包成内存友好格式如果你反复读取几千张小图片文件磁盘IO和文件系统开销会非常可观。一个很实用的加速技巧是把图像数据打包成WebDataset或LMDB格式或者干脆把所有样本缓存进内存。对于内存足够的情况直接把整个数据集load进内存是最暴力的解决方案class InMemoryDataset(Dataset): def __init__(self, img_dir, label_csv, transformNone): self.transform transform self.images [] self.labels [] df pd.read_csv(label_csv) for _, row in df.iterrows(): img Image.open(os.path.join(img_dir, row[filename])).convert(RGB) # 可选先resize到固定尺寸节省内存 img img.resize((256, 256)) self.images.append(np.array(img)) self.labels.append(row[label]) print(f已加载 {len(self.images)} 张图片到内存) def __getitem__(self, idx): image self.images[idx] label self.labels[idx] image Image.fromarray(image) if self.transform: image self.transform(image) return image, label这种写法本质上是牺牲内存换速度对于几千到几万张图片的数据集完全可行。但要注意如果图片很大且没有resize几万张可能直接占掉几十G内存反而害了自己。6.2 Dataset与DataLoader的配合调参顺序我建议按下面的顺序排查数据加载性能瓶颈先在__getitem__里打印时间看单次取样的耗时。如果超过50毫秒说明数据读取本身太重了如果单次取样很快但整体训练还是慢调num_workers如果num_workers调高后内存暴涨说明worker进程数比CPU核数还多降到等于物理核数如果某个epoch结束时总要卡顿一下很可能是最后一个batch数据不足导致的等待把drop_last设为True这个排查顺序帮我解决了很多莫名其妙的性能问题一步步来不用瞎猜。6.3 结合PyTorch的pin_memory参数还有一个很少人提到但实际很有效的参数pin_memoryTrue。当你的数据在CPU上模型在GPU上时数据从CPU内存拷贝到GPU显存前如果先把CPU内存锁页拷贝速度会快很多。只要你的机器有CUDA且内存不紧张打开这个参数基本上是无脑收益。train_loader DataLoader( dataset, batch_size32, shuffleTrue, num_workers4, pin_memoryTrue )如果是CPU训练这个参数没意义设不设都行。7. 常见问题速查表把我在社区答疑和实际项目中遇到的高频问题整理成一张速查表方便你对照自查。问题现象可能原因快速解决方案__getitem__返回的图片通道数不对灰度图没转成RGBImage.open(...).convert(RGB)RuntimeError: stack expects each tensor to be equal size同一个batch里的图片尺寸不一致在transform里加Resize到固定尺寸TypeError: default_collate: batch must contain tensors...__getitem__返回了非张量类型检查是否忘了转torch.Tensor或者自己写collate_fn训练时内存不断上涨num_workers太高worker复制Dataset开销大降低num_workers或在__init__里把数据load成轻量格式Windows上程序反复重启/崩溃没把代码包在if __name__ __main__里训练脚本主逻辑包进main函数并加判断验证集指标忽高忽低验证集用错了带随机增强的transform给验证集单独建transform不开RandomFlip、RandomCrop自定义Dataset在__getitem__里返回了dictDataloader报错默认collate_fn不支持dict或dict里的key不统一写自定义collate_fn按key分别整理数据加载慢GPU利用率上不去num_workers0数据加载和训练串行调大num_workers开pin_memoryTrue索引越界IndexError__len__和__getitem__返回的数据长度不一致检查self.data在__init__里是否被截断或重复添加这表里最后一条我特别想多说一句因为我自己犯过。有一次做数据采样我在__init__里按比例截取了self.data的一部分比如只取前80%但忘了更新self.data的长度索引结果__len__返回的还是截取后的数量__getitem__却跑到截取范围外取数据训练到中途直接IndexError。排查了半天才发现是截断之后忘了重新赋值列表。8. 一个完整的实战从文件夹到可训练数据管道前面知识比较碎这里给一个完整集成的例子从文件夹里的图片和CSV标签开始构建一个能直接送入训练循环的DataLoader。假设你的目录结构是data/ train/ cat_001.jpg cat_002.jpg dog_001.jpg train_labels.csv完整流程代码import os import torch import pandas as pd from PIL import Image from torch.utils.data import Dataset, DataLoader from torchvision import transforms class MyImageDataset(Dataset): def __init__(self, img_dir, label_csv, transformNone): self.img_dir img_dir self.transform transform self.data [] df pd.read_csv(label_csv) for _, row in df.iterrows(): self.data.append((row[filename], int(row[label]))) # 自查路径 missings [f for f, _ in self.data if not os.path.exists(os.path.join(img_dir, f))] if missings: raise FileNotFoundError(f缺失 {len(missings)} 个文件例如: {missings[:3]}) def __len__(self): return len(self.data) def __getitem__(self, idx): fname, label self.data[idx] image Image.open(os.path.join(self.img_dir, fname)).convert(RGB) if self.transform: image self.transform(image) return image, label if __name__ __main__: train_transform transforms.Compose([ transforms.RandomHorizontalFlip(), transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) train_dataset MyImageDataset( img_dirdata/train, label_csvdata/train_labels.csv, transformtrain_transform ) train_loader DataLoader( train_dataset, batch_size64, shuffleTrue, num_workers4, pin_memoryTrue, drop_lastTrue ) # 测试一个batch是否正确 images, labels next(iter(train_loader)) print(images.shape) # torch.Size([64, 3, 224, 224]) print(labels.shape) # torch.Size([64])这段代码里有几个地方值得解释一下。第一if __name__ __main__:配合num_workers4才能保证Windows不崩。第二我在训练前先跑了一次next(iter(train_loader))这是我最喜欢的调试手段用最小的代价检查数据管道是否通了。很多人在训练跑起来之后才发现数据有问题白白浪费几十分钟。现在把所有东西整合起来你会发现PyTorch里Dataset类本身没有多少玄机核心就是三个方法。把__init__里的准备工作做扎实在__getitem__里保证返回数据的类型和形状稳定和DataLoader配合时想清楚batch、shuffle、worker这些参数数据这块基本就拿捏住了。我在多个项目里沿用了这套模式无论图像分类、表格回归还是分割任务都是改改写写就能用没有翻过车。你上手之后大概率会发现自定义数据集加载其实比想象中简单得多。
返回列表