ARTICLE DETAIL

资讯详情

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

PyTorch自定义Dataset全指南:从数据加载到DataLoader调优的实践与避坑

PyTorch自定义Dataset全指南:从数据加载到DataLoader调优的实践与避坑 1. 为什么必须理解Dataset类从数据到模型的第一公里如果你刚接触PyTorch大概率最开始几周都在和张量、自动求导较劲。等到想跑一个像样的模型时突然发现自己卡在了一个看起来很简单、实际却很要命的问题上怎么把我自己的数据喂给模型这时候你就绕不开torch.utils.data.Dataset和它的搭档DataLoader。我见过不少初学者在网上找了一段加载MNIST或者CIFAR的代码跑通了就觉得自己会了。结果一换到自己整理的数据上——图片放在十几个文件夹里、文件名还带中文和空格、CSV里既有数值又有类别标签——马上就不会写了。原因很简单内置数据集帮你把脏活累活全干完了你根本没机会理解数据是怎么从磁盘变成Tensor的。而自定义数据集的加载恰恰是任何真实项目里最绕不开的一环。这篇文章要讲的就是围绕Dataset类做自定义数据集加载的完整套路。我会从设计原理讲起给出图像、表格两类最典型场景的完整代码再聊到DataLoader的参数调优和自定义collate_fn最后把我在实际项目中踩过的坑按问题清单的形式整理出来。适合刚学完PyTorch基础语法、正准备跑自己数据的读者也适合想把自己散落的数据整理成规范训练流程的从业者。2. 先别急着写代码Dataset和DataLoader各管哪一段2.1 不用Dataset类时你是怎么处理数据的很多人的第一版数据加载代码长这样把所有图片读进一个列表再用np.stack或者torch.stack拼成一个超大数组然后直接喂给模型。images [] labels [] for path in all_paths: img cv2.imread(path) img cv2.resize(img, (224, 224)) images.append(img) labels.append(get_label(path)) images np.stack(images) # 一次加载全部 labels np.array(labels)这段代码在小数据集上没问题但一旦数据量上升到几万张、几十万张机器内存直接爆掉。而且它有几个天然缺陷所有样本必须一次性进内存无法处理超出内存规模的数据集。训练时想打乱顺序、按batch取数据都得自己写索引逻辑。想对每个样本做在线数据增强得手动写在循环里代码很快就乱成一团。2.2 Dataset是数据说明书DataLoader是取货员PyTorch把数据从磁盘到模型这个流程拆成了两层Dataset只负责回答一个问题给定一个索引返回第i个样本的数据和标签。它是一个可索引的对象类似于一个有序的数据集合。具体数据是放在内存里还是从磁盘现读完全可以自己决定。DataLoader负责在此基础上做批处理按batch_size把样本攒成一捆、按shuffle决定是否需要乱序、用num_workers开多进程预取数据。这两层的分工你可以类比成菜单和服务员的关系。Dataset是菜单上面列好了每一道菜样本是什么DataLoader是服务员它会按照你的要求几桌一起上菜、要不要换顺序去后厨取菜。菜单不用关心后厨怎么做菜服务员也不用关心菜谱细节。这个拆分的核心价值在于你只需要把怎么根据索引拿到一个样本这件事写好剩下的batch、乱序、并行加载全部交给框架。更妙的是由于Dataset是按需取数的你完全可以做懒加载——每次__getitem__才去读文件几千张图片也不会占满内存。3. 自定义Dataset的三个核心方法把地基打牢3.1 先看骨架三个方法一个都不能少自定义Dataset需要继承torch.utils.data.Dataset然后实现三个方法from torch.utils.data import Dataset class MyDataset(Dataset): def __init__(self, ...): # 1. 初始化记录数据路径、标签、变换参数等元信息 pass def __len__(self): # 2. 返回数据集总样本数 pass def __getitem__(self, index): # 3. 根据索引返回一个样本数据 标签 pass我刚接触时最大的困惑是这三个方法为什么必须是这个名字尤其是__len__和__getitem__看着就很魔法方法。其实这就是Python协议的一部分——实现了这两个方法你的类就可以被len()调用、可以被下标索引行为类似一个内置列表。DataLoader在内部恰恰就是通过dataset[i]这种方式逐个取样本的所以这套协议必须完整。3.2__init__里到底该放什么一个常见的直觉错误是在__init__里就把所有图片读入内存。这样做违背了懒加载的设计初衷也让__init__变得又慢又占内存。正确的做法是__init__只负责构建样本清单——即每个样本对应的路径、标签、或者其他必要元数据。def __init__(self, img_dir, label_file, transformNone): self.img_dir img_dir self.transform transform self.samples [] # 每个元素是 (图片路径, 标签) # 解析标签文件构建样本清单 with open(label_file, r) as f: for line in f: filename, label line.strip().split(,) self.samples.append((os.path.join(img_dir, filename), int(label)))把文件清单在__init__里构建好有两点好处第一即使有十万条数据构建清单也只是读文本、拼路径速度极快第二DataLoader在训练前会先调用len(dataset)来确定总步数如果__init__太慢整个训练启动都会卡顿。3.3__getitem__才是真正的体力活每个样本的读取、解码、预处理逻辑都写在__getitem__里。这也是整个类中唯一涉及重活的地方def __getitem__(self, index): img_path, label self.samples[index] img Image.open(img_path).convert(RGB) if self.transform: img self.transform(img) return img, label注意这里index一定会是合法的整数范围在[0, len(dataset)-1]之内。因为DataLoader是先调用__len__知道边界再生成随机的索引序列来调用__getitem__。所以在__getitem__内部你通常不需要自己判断index是否越界——框架已经帮你限制了。一个更进阶的问题如果样本是变长的怎么办比如每个样本是一个不定长的文本序列或者一个不同形状的矩阵。你完全可以在这个方法里做padding或者截断保证返回的每个样本形状一致。这样一来DataLoader在堆叠batch时就不会因为形状不一致而报错。如果你希望在一个batch内部按最长样本做动态padding那就需要后面讲的自定义collate_fn这里先留个悬念。4. 实战两类最常见的自定义数据集写法4.1 场景一图像分类数据集带数据增强假设你的图片放在data/train/cat/和data/train/dog/两个文件夹里每个文件夹名就是类别名。这是最典型的图像分类场景完整的自定义数据集写法如下import os from PIL import Image from torch.utils.data import Dataset from torchvision import transforms class ImageFolderDataset(Dataset): def __init__(self, root_dir, transformNone): self.root_dir root_dir self.transform transform self.classes sorted(os.listdir(root_dir)) self.class_to_idx {cls: idx for idx, cls in enumerate(self.classes)} self.samples [] for cls in self.classes: cls_dir os.path.join(root_dir, cls) for fname in os.listdir(cls_dir): if fname.lower().endswith((.jpg, .jpeg, .png)): self.samples.append((os.path.join(cls_dir, fname), self.class_to_idx[cls])) def __len__(self): return len(self.samples) def __getitem__(self, index): img_path, label self.samples[index] img Image.open(img_path).convert(RGB) if self.transform: img self.transform(img) return img, label配合使用时的代码train_transform transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(), transforms.RandomRotation(15), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) train_dataset ImageFolderDataset(data/train, transformtrain_transform) train_loader DataLoader(train_dataset, batch_size32, shuffleTrue, num_workers4)这里有几个容易踩的细节。第一ToTensor()必须放在所有像素级变换之后、归一化之前。因为RandomHorizontalFlip这类操作作用在PIL图像上更高效而ToTensor会把PIL图像换成0到1之间的Tensor之后的Normalize则对Tensor操作。顺序反了轻则报错重则结果错误但程序不报错排查起来特别浪费时间。第二class_to_idx的构建要保证稳定。如果训练集和测试集分别构建数据集sorted(os.listdir(...))的排序结果可能因文件系统不同而不同导致同一个类别在两个数据集中映射到不同索引。稳妥的做法是单独维护一个classes.txt文件两边都按它来映射。第三__getitem__里Image.open之后别忘记convert(RGB)。灰度图、RGBA图如果不统一通道数后面张量拼接时会报维度不一致的错误。这是我被问过最多的问题之一。4.2 场景二CSV表格数据处理混合特征图像场景大家写得比较多表格数据的自定义Dataset反而很多人不会。我来写一个真实项目里很常见的场景CSV里有几列数值特征、一列类别特征、一列标签需要做归一化和类别编码。import pandas as pd import torch from torch.utils.data import Dataset class TabularDataset(Dataset): def __init__(self, csv_path, numeric_cols, category_cols, label_col): df pd.read_csv(csv_path) self.numeric df[numeric_cols].values.astype(float32) # 对类别列做简单的整数编码 self.categories [] for col in category_cols: codes, _ pd.factorize(df[col]) self.categories.append(codes) self.labels df[label_col].values.astype(float32) def __len__(self): return len(self.labels) def __getitem__(self, index): x_num torch.tensor(self.numeric[index]) x_cat [torch.tensor(self.categories[i][index], dtypetorch.long) for i in range(len(self.categories))] y torch.tensor(self.labels[index]) return x_num, x_cat, y注意这里我没有返回单个Tensor而是返回了一个包含多个Tensor的列表。DataLoader默认的collate_fn能够处理每个样本是一个元组/列表且元组内每个元素是Tensor的情况它会自动把每个位置上的多个样本堆叠成新的Tensor。所以你的__getitem__完全可以返回多个值不限于一个输入一个标签。表格数据场景有一个专属的坑归一化一定要在Dataset外部完成或者在__init__里基于全量统计量完成。如果放在__getitem__里对单个样本做(x - mean) / std而mean和std又来自整个训练集的统计量那没问题但你千万别在__getitem__里去算这个样本自己的mean和std——那就是per-sample归一化含义完全错了。实践中我建议在Dataset外部先算好训练集的统计量再传入__init__使用。4.3 一点补充不要重复造轮子如果你的数据就是文件夹名即类别的标准结构torchvision.datasets.ImageFolder已经内置实现了上面第一段代码的功能还带find_classes、make_dataset等辅助逻辑。是不是就没必要手写了不是。我自己仍然推荐手写一遍原因有二只有自己写过一遍你才能理解ImageFolder背后在做什么。遇到不标准的目录结构比如文件名里带有额外信息、不同子目录对应同一个类你才知道怎么改。真实项目的样本清单往往来自数据库导出、标注平台接口而不是纯文件目录。这时候内置类就不够用了你终究要回到自定义Dataset。手写一遍的过程就是给后面所有数据项目打地基的过程。5. DataLoader配合实战batch、乱序、多进程和自定义批处理5.1 四个参数决定训练体验DataLoader是Dataset的使用者。它的核心参数我用一张表总结参数作用我的推荐设置batch_size每次取多少样本按显存/内存来常见32、64、128shuffle每个epoch是否打乱顺序训练集True验证/测试集Falsenum_workers并行读取数据的进程数本机CPU核数的一半或4~8drop_last最后不足一个batch时是否丢弃训练集建议True验证集可以False很多新手对num_workers有误解以为越大越好。实际上这个参数是用更多CPU进程预取数据减少GPU等待。设置太大的后果是进程调度开销盖过了读取加速的收益反而拖慢训练。我常用的经验值是数据读取是IO密集型读图片、读文件时调到4~8如果数据本身就在内存里比如表格数据直接读数组num_workers2甚至0就够了。有一类特殊场景数据的读取依赖随机种子。比如__getitem__内部有随机数据增强那么shuffleTrue时每个epoch样本顺序不同此时你的随机增强最好依赖全局状态而非固定的epoch种子否则同一批样本在不同epoch看到同样的增强结果等于变相缩小了有效训练集。5.2 自定义collate_fn当默认堆叠不够用时默认的collate_fn做的事情很朴素把batch_size个样本按位置堆叠成Tensor。逻辑等价于def default_collate(batch): return torch.stack([item[0] for item in batch]), torch.stack([item[1] for item in batch])但有几类情况它做不了样本是变长序列你想要在一个batch内做padding而不是提前在getitem里pad到全局最大长度。样本里有非Tensor数据比如文本、变长的numpy数组。样本里包含字典结构你希望保留字典的键把每个键对应的值堆叠起来。最常见的是变长序列padding。假设__getitem__返回的是(序列Tensor, 长度, 标签)那么自定义collate_fn可以这样写from torch.nn.utils.rnn import pad_sequence def collate_fn(batch): seqs, lengths, labels zip(*batch) seqs_padded pad_sequence(seqs, batch_firstTrue) lengths torch.tensor(lengths) labels torch.tensor(labels) return seqs_padded, lengths, labels在这个函数里batch是一个长度为batch_size的列表列表里每个元素正是__getitem__返回的元组。你拿到这堆样本后可以随意做变换。这也是为什么很多NLP任务要自定义collate_fn的原因。如果每次都在DataLoader构造时写一遍collate_fn项目里数据集多了之后会显得重复。我自己倾向于把collate_fn写成类方法或独立函数放在数据集模块里和数据定义放在一起这样换数据集时一目了然。6. 我认为值得单独说的几个细节6.1 数据的生命周期从磁盘到GPU你可以把完整的数据流画在脑子里__getitem__从磁盘读文件 → 在CPU上做预处理 →collate_fn打包成batch →DataLoader的worker进程把batch送到主进程 → 在主进程里tensor.to(device)搬到GPU。这个流程决定了两个常见优化方向第一如果CPU预处理太慢GPU就会空转这时候优先检查__getitem__里是否有重复的、可以缓存的计算第二num_workers的本质是让多个CPU进程并行执行__getitem__所以如果你的__getitem__本身很轻比如只是从内存数组切一行开一堆worker纯属浪费。6.2 随机增强与样本绑定的陷阱有一个很容易被忽视的问题shuffleTrue时数据在每个epoch的读取顺序是变化的但__getitem__内部的随机增强每次调用都会重新随机。这本身没问题。问题是有些代码会这样写# 错误示范在Dataset外部准备了一份数据却没意识到每次取数都会重新增强 for epoch in range(epochs): for batch in train_loader: ...这里没问题。但如果你为了加速把预处理后的数据缓存成一个列表再从这个列表构建Dataset那么__getitem__每次返回的其实是同一份缓存数据随机增强就没多少意义了。类似的问题也出现在你想做确定性验证的场景——你需要把random.seed或者torch.manual_seed设置在合适的位置但通常不建议在__getitem__内部设置全局随机种子否则多进程场景下所有worker会生成完全相同的增强序列等于训练效果被严重削弱。6.3 分布式训练前的数据切分使用DataLoader做分布式训练时有一个容易被忽略的问题每个进程应该看到不同的数据分片。如果你只是简单地把同一个DataLoader分发到多个GPU而shuffleTrue的随机种子不同在某些实现中可能会出现数据重复或漏读。PyTorch 提供了DistributedSampler来解决这个问题用法是构造DataLoader时传入samplerDistributedSampler(dataset)并在每个epoch开始时调用sampler.set_epoch(epoch)来保证shuffle的差异性。这个内容虽然进阶但一旦你开始用多卡训练它几乎是必踩的坑。我的建议是从单卡到多卡切换时不要只改.to(device)和torch.distributed的部分数据侧也必须同步改造。你可以在自己的代码里加一个判断如果args.distributed为真就自动切换到DistributedSampler避免遗留bug。7. 常见问题与排查技巧实录7.1 问题速查表我在带项目时把这些年学员和同事问得最多的问题整理成了一张排查表分享给你症状常见原因解决办法IndexError: index out of range__len__返回值比实际样本数大检查__len__里返回的是不是len(self.samples)batch堆叠时报shape不一致__getitem__返回的样本形状不固定在__getitem__里做resize或padding训练进程卡死或极慢num_workers过大或__getitem__里有死锁调小workers检查是否有全局锁数据增强结果对不齐ToTensor和Normalize顺序错误按PIL变换→ToTensor→Normalize排列验证集结果比训练集好很多训练集看的是增强后的数据验证集没增强或评估逻辑不一致确认评估时用的也是未增强数据或同样的预处理DataLoader报RuntimeError: received 0 items of ancdata传输给worker的数据太大降低batch_size或减少单样本体积7.2 一个真实排查案例图片加载异常有一次我帮同事排查模型训练到一半突然报错的问题错误信息指向__getitem__里的Image.open打不开文件。查了半天发现是某些图片文件本身损坏或者文件名对不上。这个问题在数据量大时特别容易碰到因为少量坏数据不会在一开始就触发要等到训练到那个索引才爆炸。我的建议是在正式训练前先写一个数据集体检脚本把Dataset完整遍历一遍dataset MyDataset(...) for i in range(len(dataset)): try: _ dataset[i] except Exception as e: print(f样本索引 {i} 出错: {e}) break这个脚本看起来简单但能帮你在训练前发现绝大部分数据侧的异常。我基本上每个新数据集都会跑一遍花不了几分钟却能避免训练中途崩溃带来的时间浪费。7.3 内存泄漏与进程残留问题很多人在Windows上跑PyTorch时会发现训练停了但内存没有完全释放甚至Python进程结束了一部分子进程还残留。这个问题的根源往往是num_workers 0时DataLoader创建的worker进程没有正确退出。常规解法是确保使用DataLoader时代码位于if __name__ __main__:的保护之下。在Jupyter Notebook里则更容易踩坑因为笔记本的交互式环境会频繁创建worker。我的经验是在Notebook里调试时先用num_workers0逻辑确认无误再切到多进程同时尽量不要在同一个内核里反复创建和销毁多个DataLoader必要的话用del loader之后调用gc.collect()。7.4 数据预处理应该放在Dataset里还是外面这个问题的答案取决于你的场景。我的经验原则是应当放在Dataset外或__init__里提前算好的统计量计算均值、方差、缩放系数、词汇表构建、类别映射。应当放在__getitem__里的文件的读取、解码、单样本变换、数据增强。为什么因为__getitem__会被并行调用成千上万次任何重计算放这里都会被放大batch_size倍的代价。而统计量放到__init__里整个生命周期只算一次性价比高得多。比如图像归一化用的mean和std从全量数据集的统计来但你在__init__里只做一次全量遍历而不是每个样本都去算。8. 最后再分享一点我的实际体会写了这么多年PyTorch代码我对Dataset的理解经历过三个阶段。最开始觉得它就是必须继承的类照着模板抄后来开始把业务逻辑往里塞结果数据集代码越来越臃肿再后来才明白好的Dataset设计本质上是把数据如何呈现给模型这件事清晰地表达出来。__init__决定数据范围__getitem__决定样本形态collate_fn决定batch形态——各司其职边界清楚。一个验证你设计好不好的标准是如果换一个模型你的Dataset和DataLoader代码一行都不用改那说明数据层与模型层解耦得不错如果你发现每换一个任务就要大改数据脚本大概率是你在数据类里混入了太多不该有的业务逻辑。这套DatasetDataLoader的加载模式也是很多经典开源项目能快速在不同数据集间迁移的基础。把这块地基打好后面无论是跑resume、做多卡、上分布式还是接更复杂的多模态数据你都会比别人少花一半的时间去趟坑。
返回列表