☰
PyTorch Dataset类实战指南:核心方法、DataLoader协作与常见坑解析
2026/10/2 2:10:58 网站建设 项目流程

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的配合调参顺序

我建议按下面的顺序排查数据加载性能瓶颈:

  1. 先在__getitem__里打印时间,看单次取样的耗时。如果超过50毫秒,说明数据读取本身太重了
  2. 如果单次取样很快,但整体训练还是慢,调num_workers
  3. 如果num_workers调高后内存暴涨,说明worker进程数比CPU核数还多,降到等于物理核数
  4. 如果某个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__返回的图片通道数不对灰度图没转成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__里返回了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这些参数,数据这块基本就拿捏住了。我在多个项目里沿用了这套模式,无论图像分类、表格回归还是分割任务,都是改改写写就能用,没有翻过车。你上手之后,大概率会发现,自定义数据集加载其实比想象中简单得多。

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

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

立即咨询