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, transform=None): 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, transform=None): 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_dir="train_images/", label_csv="train_labels.csv", transform=transform )如果你没有现成的CSV,只是图片文件名里带标签,比如cat_001.jpg、dog_002.jpg,那__init__里解析文件名就够了:
def __init__(self, img_dir, transform=None): 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, targets=None): # 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, transform=None, mask_transform=None): 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), dtype=torch.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_last=True。如果不设默认False,意味着最后一个batch会被保留,但batch大小和其他不一样,可能导致模型训练的稳定性稍微受影响。我的习惯是设True,省心。
4.2 num_workers:这个参数很多人一辈子就设为0
num_workers决定用几个子进程去预取数据。默认是0,意思是数据在主进程里同步加载,模型在GPU上算完一批,CPU才去加载下一批,二者串行,GPU经常闲着等数据。设成4或8之后,加载下一批数据的操作会提前在另一个进程里做好,模型一算完,数据已经等在那边了,训练速度能快好几倍。
但是num_workers不是越大越好。开太高了会有两个问题:一是每个worker都要复制一份Dataset的内存副本(实际是按需复制),内存占用暴涨;二是进程切换的开销反而拖慢速度。我自己的经验法则是:
- 普通办公CPU,先设2
- 16核左右的机器,设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。比如处理变长文本序列时,通常要做pad:
def collate_batch(batch): images, labels = zip(*batch) # 假设images已经是torch.Tensor,尺寸一致,只是堆成batch images = torch.stack(images, dim=0) 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], dim=0), "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(p=0.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(image=image, mask=mask) 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实例,结果验证时也做了随机增强,评估指标忽高忽低,完全不可信。
正确做法是训练集和验证集要用不同的transform:
train_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(..., transform=train_transform) val_dataset = CustomDataset(..., transform=val_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, transform=None): 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_memory=True。当你的数据在CPU上,模型在GPU上时,数据从CPU内存拷贝到GPU显存前,如果先把CPU内存锁页,拷贝速度会快很多。只要你的机器有CUDA且内存不紧张,打开这个参数基本上是无脑收益。
train_loader = DataLoader( dataset, batch_size=32, shuffle=True, num_workers=4, pin_memory=True )如果是CPU训练,这个参数没意义,设不设都行。
7. 常见问题速查表
把我在社区答疑和实际项目中遇到的高频问题整理成一张速查表,方便你对照自查。
| 问题现象 | 可能原因 | 快速解决方案 |
|---|---|---|
__getitem__返回的图片通道数不对 | 灰度图没转成RGB | Image.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__里返回了dict,Dataloader报错 | 默认collate_fn不支持dict,或dict里的key不统一 | 写自定义collate_fn,按key分别整理 |
| 数据加载慢,GPU利用率上不去 | num_workers=0,数据加载和训练串行 | 调大num_workers,开pin_memory=True |
| 索引越界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, transform=None): 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_dir="data/train", label_csv="data/train_labels.csv", transform=train_transform ) train_loader = DataLoader( train_dataset, batch_size=64, shuffle=True, num_workers=4, pin_memory=True, drop_last=True ) # 测试一个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_workers=4才能保证Windows不崩。第二,我在训练前先跑了一次next(iter(train_loader)),这是我最喜欢的调试手段,用最小的代价检查数据管道是否通了。很多人在训练跑起来之后才发现数据有问题,白白浪费几十分钟。
现在把所有东西整合起来,你会发现,PyTorch里Dataset类本身没有多少玄机,核心就是三个方法。把__init__里的准备工作做扎实,在__getitem__里保证返回数据的类型和形状稳定,和DataLoader配合时想清楚batch、shuffle、worker这些参数,数据这块基本就拿捏住了。我在多个项目里沿用了这套模式,无论图像分类、表格回归还是分割任务,都是改改写写就能用,没有翻过车。你上手之后,大概率会发现,自定义数据集加载其实比想象中简单得多。