看到标题里的PyTorch3我先愣了一下。PyTorch目前真没有3.x这个版本我猜这大概是PyTorch和教程第3期这类说法被搅在了一起。名字不重要重要的是Dataset与DataLoader这套数据装载机制——所有PyTorch训练脚本都绕不开它们同时它们也是新手从跑通demo走向用自己数据集训练的第一道坎。我见过太多这样的场景从GitHub上clone了一个项目模型、损失函数都调通了结果一换到自己的数据训练脚本直接原地爆炸。报错五花八门但翻来覆去根子大多在数据装载上——要么Dataset写得不对要么DataLoader参数设置得离谱。这篇教程不打算照着官方文档念一遍而是从数据到底是怎么从硬盘流进GPU的讲起把自定义数据集的完整写法、DataLoader每个参数的取舍逻辑都说透再把我真实踩过的几个坑和排查思路分享出来。适合谁看刚配好环境、第一次准备跑自定义数据的小白也包括已经写过不少训练脚本、但一直对collate_fn、num_workers这些地方知其然不知其所以然的朋友。如果你是后者可以直接跳到第3节。1. DataLoader到底在替你扛什么活先理解三件事1.1 没有DataLoader的日子为什么朴素循环走不远很多朋友第一次接触数据加载是从这样的代码开始的for epoch in range(num_epochs): for i in range(len(x_train)): sample x_train[i] label y_train[i] # forward数据量小的时候这代码完全能跑。一个几千样本的MNIST几个epoch下来也没觉得哪不对。但一旦换成上万张图片、每个样本还要做随机裁剪、翻转、归一化问题就来了循环体里塞了太多杂事数据预处理、batch组装、迭代顺序控制全部混在一起代码会越来越乱改一处崩三处。更关键的是一件事这个循环根本没法利用多核CPU。PyTorch训练时GPU在算CPU大部分时间是闲着的但你没法让读下一批数据和训练当前这批数据同时进行。于是GPU利用率上不去显存占用像过山车训练时间翻倍都是轻的。DataLoader的出现就是为了解决这几件事帮你把数据按需求打乱顺序、按batch_size凑成一批、用多进程提前读好下一批、并且通过collate_fn把零散样本拼成模型能直接吃的张量。你只需要告诉它数据怎么取和怎么组合剩下的脏活累活它全包了。1.2 DataLoader与Dataset的分工后厨备菜与前台传菜理解这两个类最简单的方式是把它想成一家餐厅。Dataset是后厨备菜的厨师。它不关心菜怎么端上桌、一桌几道菜它只认一个东西编号。你给它一个编号索引idx它给你一份原材料单个样本。这份原材料可以是一张图一个标签可以是一句话一个分类也可以是一段视频一组边界框。只要你能用编号定位到它厨师就能把它取出来。DataLoader是前台的传菜员。它负责按一定规则向后厨喊号每批要多少份batch_size、按什么顺序喊shuffle决定是否随机、要不要一次多喊几份提前备着prefetch、num_workers决定并行度、喊出来的菜要怎么摆盘collate_fn决定怎么把单个样本拼成一个batch。这个分工特别重要。很多人写不好自定义数据集就是因为在Dataset里想干DataLoader的活或者在DataLoader里想干Dataset的活。记住一条准则**Dataset只负责按索引返回一个样本DataLoader负责把多个样本变成一个小批量并送到训练循环手里。**这条边界划清楚了后面所有问题都好解决。1.3 和TensorFlow的tf.data对比一个直观的映射如果你是从TensorFlow转过来的会发现在数据加载这件事上两家思路差别不小。TensorFlow的tf.data.Dataset更像一条完整的流水线定义数据来源、shuffle、batch、map、prefetch一气呵成有点像PyTorch里Dataset DataLoader transform三样东西叠在一起的效果。PyTorch则把这些步骤拆得比较散Dataset管单样本、transform管预处理、DataLoader管批次和并发。相比之下PyTorch的拆法更Pythonic也更贴近普通人的思考习惯先想清楚单个样本长什么样再想怎么批量化。而且因为Dataset本质是一个普通的Python类里面的逻辑可以写得非常灵活——比如动态从数据库拉数据、根据当天日期切换数据源、在网络请求失败时换备用数据。这些场景在tf.data里实现起来要别扭得多。所以如果你是从TF转来的不用在这套设计哲学上纠结太久上手写一个Dataset类跑一遍训练循环自然就顺了。2. 手写Dataset最核心的三个方法调用时机一次讲透2.1init只登记不干活按我自己的经验写自定义Dataset最容易犯的错误就是在__init__里把所有数据都读进内存。看起来省事但实际上很危险。一个10万张图片的数据集每张图如果平均几MB__init__一跑就是几十个GB的内存占用更麻烦的是后面如果设置了多个num_workers每个子进程都会对Dataset做一次处理内存占用还会成倍翻。正确的做法是__init__里只记录从哪拿数据的信息而不是数据本身。比如图片分类任务__init__只需要把文件夹下所有图片的路径和对应标签整理成两个list文本任务__init__只需要读一个CSV的路径和类别就可以不用把全部文本内容都塞进内存。我一般习惯把数据结构组织成一个list每个元素是一个小字典或者元组比如class MyDataset(Dataset): def __init__(self, root_dir): self.samples [] # 每个元素为 (img_path, label) for cls_name in sorted(os.listdir(root_dir)): cls_path os.path.join(root_dir, cls_name) for fname in os.listdir(cls_path): if fname.lower().endswith((.jpg, .jpeg, .png)): self.samples.append((os.path.join(cls_path, fname), cls_name))记住这个原则**__init__里做尽量少的耗时操作只做索引和路径的建立。**真正读图片、读文本、做解码的活全部留到__getitem__里。2.2len__与__getitem按编号备菜的规则__len__很简单返回样本总数。这个数值会直接影响DataLoader对数据量的判断包括drop_last要不要丢最后一个batch、进度条总共显示多少步等。写错它不会立刻报错但会出现训练了几个epoch但实际数据没用全或最后一个batch永远凑不满这类隐性问题。__getitem__是真正干活的地方。DataLoader内部会生成一堆索引然后逐个调用dataset[idx]这实际上就是触发了你的__getitem__(idx)方法。一个容易被忽略的细节__getitem__里应该处理异常。比如PIL打开一张损坏的图片时会直接抛异常整个训练进程就崩了。我遇到这种情况会在读取时加一个防御逻辑比如跳过损坏样本返回相邻样本def __getitem__(self, idx): while True: img_path, label self.samples[idx] try: image Image.open(img_path).convert(RGB) break except Exception: idx (idx 1) % len(self.samples) # 取下一个样本当然更稳妥的方案是在训练前置阶段跑一遍数据完整性检查把坏图提前踢掉。但防御代码的存在至少能让你在野外数据上不至于一个epoch都撑不过去。2.3 把图片文件夹变成Dataset一个可以直接抄的例子光讲原理不落地总是虚的我直接给一个完整可用的图片二分类/多分类Dataset代码import os import torch from torch.utils.data import Dataset from PIL import Image 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: i for i, cls in enumerate(self.classes)} self.samples [] for cls_name in self.classes: cls_folder os.path.join(root_dir, cls_name) for fname in os.listdir(cls_folder): if fname.lower().endswith((.jpg, .jpeg, .png)): self.samples.append(( os.path.join(cls_folder, fname), self.class_to_idx[cls_name] )) def __len__(self): return len(self.samples) def __getitem__(self, idx): img_path, label self.samples[idx] image Image.open(img_path).convert(RGB) if self.transform is not None: image self.transform(image) return image, label几个我特别想强调的点第一路径尽量用绝对路径。因为训练脚本的工作目录很可能随时变化用了相对路径今天能跑通明天在另一个目录下运行就莫名其妙FileNotFoundError。这一点在团队协作里几乎天天遇到。第二transform参数建议从外部传入不要写死在Dataset内部。这样同一个Dataset类可以轻松组合不同预处理流程训练集用随机增广验证集只用ResizeToTensor代码复用率高出很多。第三如果你用的是图片分类之外的任务比如目标检测那么__getitem__返回的就不能只是(image, label)。它可能需要返回一个字典里面包含image、boxes目标是哪些框、labels每个框的类别。这种数据结构后面第3节会详细讲因为DataLoader要正确处理它必须配一个专门的collate_fn。2.4 常见的标注格式VOC/YOLO/JSON怎么接进来很多朋友拿到的数据集不是干净的文件夹结构而是带标注文件的格式比如VOC的xml标注、YOLO的txt标注、还有各种json标注。这里有一个很多人会问的问题Dataset对不同标注格式是不是要写不同模板答案是**Dataset本身不关心标注格式是什么它只关心你能不能在__getitem__里返回一个规范的样本。**你需要做的是在__init__阶段把这些五花八门的标注文件解析成统一的内存结构这样__getitem__就只需要按索引取出这份结构然后读图片即可。举个例子VOC格式的xml标注我一般会在__init__里用ElementTree解析一遍import xml.etree.ElementTree as ET for xml_file in self.xml_files: tree ET.parse(xml_file) root tree.getroot() boxes [] for obj in root.findall(object): name obj.find(name).text bbox obj.find(bndbox) boxes.append({ label: name, xmin: int(bbox.find(xmin).text), ymin: int(bbox.find(ymin).text), xmax: int(bbox.find(xmax).text), ymax: int(bbox.find(ymax).text), }) self.annotations.append({image_path: img_path, boxes: boxes})这样解析一次之后__getitem__里的工作就只剩下读图片和把boxes转成你想要的tensor格式比如YOLO需要的归一化坐标或者detectron2的format。如果你每次__getitem__都去重新解析xml不仅慢而且容易出现句柄泄漏之类的隐性毛病。顺便提一句网上看到有人问douyin comment dataset怎么处理这类从接口或爬虫抓下来的评论数据格式多半是JSON或CSV。处理思路一样__init__读JSON或CSV把文本内容和标签组织成list然后__getitem__按索引取出来做tokenizer编码。千万不要一上来就把几十万条评论全在__init__里编码成token序列你会发现内存直接爆炸。先存原文等访问到再编码。3. DataLoader参数逐个拆每一个都决定了训练能不能跑起来3.1 batch_size、shuffle、drop_last最容易被忽略的组合拳这三个参数是最常见的但组合起来有很多门道。batch_size决定了每个batch多少样本。很多人是拍脑袋定的我建议先看显存和内存如果是常规分类模型、输入224x224batch_size从32或64起步显存不足就往下减梯度更新频率不够平滑就往上加。另外batch_size和drop_last配合很重要如果数据总量不能被batch_size整除最后一个batch样本数会很少。如果你的网络里有BatchNorm最后一个batch样本太少会导致统计量极其不准甚至个别版本会直接报错或输出NaN。训练集上我一般倾向drop_lastTrue宁可丢掉几十个样本也不让一个异常的batch干扰训练但验证集上为了评估全部样本通常drop_lastFalse。shuffleTrue是一个容易理解但很多人忽略的逻辑训练集打乱顺序是为了避免模型学到数据集的固定排列规律验证集和测试集不能打乱是为了评估指标的可复现性。还有一个细节用了shuffle就不要显式传sampler参数两者是互斥的同时传会直接抛ValueError。3.2 num_workers、pin_memory、prefetch_factor并行的性价比陷阱num_workers是很多人特别容易冲动设置的一个参数总觉得越大越好。实际上这个想法有问题。每个worker是一个独立的Python进程它们负责不断调用dataset[idx]来准备数据然后把数据通过队列传给主进程。建议是num_workers从CPU逻辑核心数的一半起步最多不要超过CPU核心数。可以在训练时用nvidia-smi观察GPU利用率和CPU占用如果GPU利用率长期低于70%并且CPU没跑满那说明worker数量不够如果CPU已经快满了而GPU仍在等待那就不是加worker能解决的可能瓶颈在磁盘IO或者预处理本身。pin_memoryTrue这个参数很多人忽略其实它对GPU训练有明显收益。原理是把数据放在锁页内存pinned memory里GPU可以直接通过DMA方式从这块内存拷贝到显存省去中间一次交换。代价是锁页内存不能被系统换出到磁盘会占用物理内存。所以机器内存够、显存够的情况下大胆开pin_memoryTrue如果内存非常紧张就权衡一下宁可不开也别导致系统内存耗尽。prefetch_factor是每个worker预先准备几个batch的数据默认是2。样本比较大时prefetch_factor调低到1可以少占点内存样本很小、数据源很慢时调高一点能改善流水线吞吐。注意这个参数在PyTorch 1.10版本之后才支持非默认值老版本别乱设。3.3 collate_fn默认行为与你必须重写的时机collate_fn是DataLoader里最容易被忽略、但实战中价值最高的参数因为它决定了多个独立样本怎么变成一个batch。默认情况下DataLoader会把所有样本放进一个list然后调用default_collate把一组tensor用torch.stack摞起来组成一个新的维度batch维。这对所有样本shape一致的情况完美适用。但有很多场景默认collate会直接报错文本数据经过tokenizer后每个句子的长度不一样目标检测任务里每张图的边界框数量不一样样本返回的是字典里面键相同但值的shape不一致。这些情况都需要你自己写collate_fn。我一般这么写def collate_fn(batch): images torch.stack([item[image] for item in batch]) lengths torch.tensor([item[length] for item in batch]) labels torch.tensor([item[label] for item in batch]) return { images: images, lengths: lengths, labels: labels, }有一点要提醒DataLoader拿到一批样本后只调用一次collate_fn处理整个batch。所以你不用担心它在单样本层面做什么你只需要拿到一个Python listbatch里的每个元素就是__getitem__的返回值然后在这个函数里把它们处理成模型需要的结构。还有一个容易被忽略的坑如果你写了自己的collate_fn默认的tensor类型转换、跨设备操作都不会自动发生。比如你的Dataset返回的是PIL图片DataLoader不会自动帮你堆叠必须自己在collate_fn里先ToTensor()再torch.stack。所以我在设计Dataset的时候会尽量让__getitem__的返回已经是tensor或纯数值类型避免在collate里再做转换。3.4 sampler家族样本不平衡时遇到的WeightedRandomSampler如果你遇到的是类别极度不平衡的数据集比如正样本只有1万负样本却有9万光靠shuffle是不够的因为每个epoch里负样本依然占大多数。这时候WeightedRandomSampler就是正确答案。思路很简单给每个样本一个权重权重越高的样本被抽中的概率越大。连写都不用写循环这样构造from torch.utils.data import WeightedRandomSampler # 假设dataset.samples里存着 (path, label) class_counts {} for _, label in dataset.samples: class_counts[label] class_counts.get(label, 0) 1 weights [] for _, label in dataset.samples: weights.append(1.0 / class_counts[label]) sampler WeightedRandomSampler(weights, num_sampleslen(weights), replacementTrue) dataloader DataLoader(dataset, batch_size32, samplersampler)注意用了sampler就不要同时设shuffleTrue了因为采样器已经控制了取样本顺序。replacementTrue表示同一个epoch里允许重复抽到同一个样本这是处理极端不平衡时常用的手段作用等效于对少数类样本做重采样。4. 跑通一次完整训练从数据集到batch的流动全记录4.1 组合一个可复用的数据流水线有了前两节的积累现在可以组合一个完整可跑的流程了。我用第2节的ImageFolderDataset加一个常见的图像增广组合然后划分训练集和验证集。import torch from torch.utils.data import DataLoader, random_split from torchvision import transforms train_transform transforms.Compose([ transforms.Resize((224, 224)), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) val_transform transforms.Compose([ transforms.Resize((224, 224)), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]), ]) full_dataset ImageFolderDataset(data/classification, transformtrain_transform) train_size int(0.8 * len(full_dataset)) val_size len(full_dataset) - train_size train_dataset, val_dataset random_split(full_dataset, [train_size, val_size]) # 注意验证集要用不同的transform所以这里不要用同一个dataset这里想提醒一个细节验证集不应该做随机增广因为增广带来的随机性会影响验证指标的稳定性。所以验证集我会单独换一个transform传入或者直接单独构造一个ImageFolderDataset实例。有些人图省事在random_split之后用同一个Dataset结果验证时也做RandomHorizontalFlip评估指标忽高忽低很难定位问题。划分之后分别构造DataLoadertrain_loader DataLoader( train_dataset, batch_size32, shuffleTrue, num_workers4, pin_memoryTrue, drop_lastTrue ) val_loader DataLoader( val_dataset, batch_size32, shuffleFalse, num_workers4, pin_memoryTrue )训练集和验证集的三个区别值得记住shuffle一个True一个Falsedrop_last一个True一个Falsetransform一个带增广一个不带。其他参数保持一致这样比较结果才有意义。4.2 训练循环里batch是如何一步步变成loss的一个标准的训练循环长这样device torch.device(cuda if torch.cuda.is_available() else cpu) model MyModel().to(device) criterion nn.CrossEntropyLoss() optimizer torch.optim.Adam(model.parameters(), lr1e-3) for epoch in range(num_epochs): model.train() for batch_idx, (images, labels) in enumerate(train_loader): images images.to(device) labels labels.to(device) outputs model(images) loss criterion(outputs, labels) optimizer.zero_grad() loss.backward() optimizer.step() if batch_idx % 50 0: print(fepoch {epoch} step {batch_idx} loss {loss.item():.4f})注意这里的images已经是一个四维tensor了[B, 3, 224, 224]。它就是从一堆独立样本经由DataLoader的默认collate合并出来的。这就是DataLoader替你干的活——如果你绕过它手动做batch得自己循环、自己stack、自己保证张量shape对齐麻烦得多。关于.to(device)有个细节DataLoader产出的batch还在CPU内存上不会自动搬去GPU。不少人第一次写训练循环忘了把images和labels挪到GPU然后遇到一个莫名其妙的input type (CPUFloatTensor) and weight type (CUDAFloatTensor) should be the same报错。这个报错本身就说得很清楚了但很多新手会被吓到。我建议把.to(device)养成肌肉记忆训练循环里写入这几行时脑子里默念数据在哪、模型在哪、目标值在哪。4.3 调试数据流的三板斧next(iter)、打印shape、检查dtype数据流最容易出问题的地方就是你猜不透__getitem__到底返回了什么。我调试时最常用的三个手段第一招用next(iter(train_loader))手动取一个batch。这个方法极其好用它不会启动一个完整epoch只需几秒钟就能拿到一个组装好的batchsample_batch next(iter(train_loader)) print(sample_batch[0].shape, sample_batch[0].dtype) print(sample_batch[1].shape, sample_batch[1].dtype)如果这里就报错那问题一定出在Dataset或collate_fn跟模型无关别傻傻地等训练循环报错后再回头查。第二招每个epoch的第一次迭代打印shape。很多人把打印放在循环外最后看半天也没看出名堂。放在第一个batch的forward前后能同时确认输入shape是否匹配第一层全连接或卷积层的期望输入。第三招检查dtype和数值范围。图像数据经过ToTensor后应该是torch.float32且范围在0到1之间如果某个地方粗心把图片读成了0到255的整型tensor模型虽然能算但loss可能一开始就非常大甚至梯度爆炸。这个我用sample_batch[0].min(), sample_batch[0].max()一眼就能看出来。另外如果__getitem__返回的类型不是你预期的那样可以用type(sample)或者sample.__class__.__name__打印出实际类型。这会比你逐行读代码快得多毕竟Python的鸭子类型在Dataset里经常玩出花来。5. 自定义数据集踩坑实录四条完整排查链路5.1 closed dataset报错文件句柄的生命周期先于worker结束先说一个我在真实项目中遇到的报错大意是cannot perform this operation on a closed dataset。当时我用zipfile打包了一批图片数据想在Dataset里直接读取。第一次写的时候我在__init__里用了with zipfile.ZipFile(...)想着这样最简洁、能自动关闭文件。但跑起来就发现__getitem__执行到一半抛了上面那个报错。原因其实很清晰with语句块在__init__返回前就执行完了zip文件句柄被关闭。等DataLoader在Worker进程里调用__getitem__访问这个文件时句柄早就没了。这个问题的本质是文件句柄的持有时间必须覆盖所有你会访问它的方法。我当时排查的顺序是这样的先单开一个dataset[0]直接跑发现没报错——说明__init__本身没问题。再用DataLoader设num_workers0跑也没报错。这就把问题缩小到了多进程相关。设num_workers2必现。继续看堆栈报错指向__getitem__里访问zip文件的那一行。回看代码确认是with提前关闭了句柄。修复方案有两种要么把__init__里打开的zip文件句柄保存为实例属性不在with里关闭要么干脆在每个__getitem__里临时打开、读完就关闭。后者虽然每次多一次open开销但最稳妥也不会在不同worker之间出现并发访问同一文件句柄的隐患。我的建议是如果文件数量不多数据预处理又不重就现场open现场close安全第一。5.2 Windows下num_workers0就疯狂重启spawn与__main__保护这个大坑基本是Windows用户的噩梦。代码在Linux上跑得好好的num_workers4一路顺畅换到Windows训练脚本要么卡死要么看起来像是不断重新执行整个脚本。根因是进程创建方式不同。Linux上Python默认用forkWindows上只能用spawn。spawn模式下每个新进程都会重新导入主模块。如果你的DataLoader创建代码和训练循环写在了模块顶层没有放进if __name__ __main__:保护块里那么每个worker进程一启动就会再次创建DataLoader而创建DataLoader又触发新的worker进程形成死循环。排查链路不复杂先看代码主体是否都在if __name__ __main__:里如果没写立刻补上。如果补上了还卡检查num_workers是否大于0暂时改为0看是否恢复正常。如果num_workers0正常说明是spawn和顶层代码交互的问题检查是否有全局变量、模块级代码在重复执行。标准写法是把所有训练逻辑包进main()函数然后在文件末尾只留一句if __name__ __main__: main()这不仅仅是为了规范在Windows上这是multiprocessing能正常工作的前提。另外补充一句被反复导入的模块里如果存在耗时操作比如读配置、建日志、初始化数据库连接也会被每个worker重复执行一遍这也是Windows上运行慢的隐形原因。可以在模块顶层用if __name__ __main__或者把初始化收敛进main()来规避。5.3 增广始终一模一样随机种子被全局锁死的真相有段时间我发现训练时看到的图片增广结果几乎是一样的横竖翻转都是同一个方向裁剪位置也差不多。起初以为是错觉后来在__getitem__里打印了每次的随机值全是同一个数才意识到问题出在随机种子的传播机制上。原因是这样的DataLoader的多进程worker在启动时会继承主进程的随机数生成器状态。如果你在训练脚本里设置了torch.manual_seed(0)那么每个worker拿到的全局随机状态可能是相同的于是它们在不同worker里做随机增广生成的一系列随机数完全一致。排查路径先确认有没有全局设置随机种子。如果设置了就别指望多进程Worker各自独立随机了。在DataLoader的worker_init_fn里为每个worker重新设置独立的随机种子def worker_init_fn(worker_id): torch.manual_seed(42 worker_id) np.random.seed(42 worker_id)如果还在__getitem__里用了Python内置的random模块也要记得在worker_init_fn里给random.seed(42 worker_id)。做完这些每个worker就有了独立的随机数流增广效果才会真正多样。注意这里的target是保证复现和多样性之间的平衡主进程的seed控制整体训练流程可复现worker级seed控制数据增广不重复。5.4 内存被8个worker干爆子进程对Dataset实例的复制机制之前接手过一个项目数据集不大但预处理很重__init__里把每张图片读成tensor后统一存在Dataset实例里。单进程跑没问题一改成num_workers8内存直接突破警戒线。原因在于Linux下DataLoader默认用fork创建worker进程子进程会以写时复制copy-on-write的方式继承父进程内存。理论上不会立刻真实复制可一旦子进程需要修改某部分内存比如做数据增广时写数据复制就会发生。而Windows下spawn直接会pickle整个Dataset实例传给子进程如果Dataset里存了上万个图片tensor每个子进程都会真实拥有完整数据集的一份拷贝。8个worker就是8份内存不炸才怪。这个问题的排查思路比较清晰打开nvidia-smi或系统监控看内存涨幅与worker数量的关系。将num_workers调为0内存回归正常大概率就是Dataset实例太大导致的复制。看Dataset的__init__是否做了重量级的数据预载。最终修复方向有两个一是把重量级数据移出Dataset实例改成从磁盘、内存映射或共享存储里按需读取二是减小每个worker的复制代价比如先预处理成单独的npy文件__getitem__里用np.load按需读取。同步提一句如果你的数据是纯tensor且内存总量能放下可以尝试把tensor放进torch.utils.data的TensorDataset它内部对内存的管理会高效很多但也不是万能的数据量特别大时依然会受限。6. 数据加载性能优化小文件、大文件与分布式缓存6.1 先别堆num_workers从IO类型判断瓶颈很多朋友遇到训练慢第一反应是加num_workers。但我要先泼一盆冷水如果瓶颈在磁盘IO上加worker作用很有限。判断瓶颈的方法很朴素在训练脚本里单独把DataLoader的数据读取跑一遍不训练只看每秒能产出多少个batch。如果这个吞吐远高于训练所需说明DataLoader不是瓶颈如果这个吞吐本来就很低再多的worker也救不了你。具体来说需要区分两种情况随机IO密集和顺序IO密集。前者是从几万个分散文件夹里随机读小文件瓶颈通常在随机读性能后者是从一个大文件里按顺序连续读取瓶颈通常可以是带宽。想要压榨性能就得对症下药。随机IO密集场景核心思路是把小文件变成大文件减少随机寻址次数顺序IO场景重点看预读取深度和系统页面缓存。6.2 小文件灾难的解法打包与内存映射我自己在一台普通工作站上遇到过血泪教训一个约2万张图片的数据集散落在几千个文件夹里机械硬盘上跑一个epoch光读文件就花掉40多分钟。后来做了两个改动时间直接降到5分钟以内。第一个改动是把散碎文件打包成单一的文件比如HDF5、LMDB、tar归档。打包之后文件系统寻址次数大幅减少随机IO变成了近似顺序IO。我后来用LMDB比较多因为它在小数据集上API简单读写都快。示意代码大概是import lmdb import pickle # 写入 env lmdb.open(images.lmdb, map_size1e10) with env.begin(writeTrue) as txn: for idx, (img_bytes, label) in enumerate(all_samples): txn.put(str(idx).encode(), pickle.dumps((img_bytes, label))) # Dataset里读取 with env.begin() as txn: img_bytes, label pickle.loads(txn.get(str(idx).encode()))第二个改动是利用操作系统的页面缓存和内存映射。如果你的数据原本就是numpy数组可以直接np.load(path, mmap_moder)这样不会一次性加载全部数据到内存而是按页惰性读盘。对于特征向量这类数据用mmap_mode几乎是无痛的性能提升。顺带说如果你看到网上有人讨论loading redis is loading the dataset in memory那说的其实是把数据集提前加载到内存/缓存型存储里思路和我上面说的减少每次IO是一脉相承的。别惦记着用Redis去解决所有IO问题——单机场景用文件和内存映射才是主流。6.3 Redis/共享缓存适合哪些场景多机训练才值得有一种场景我用过Redis多机多卡分布式训练时想让所有节点共享同一份预处理后的样本。比如你把原始数据统一做了一次增广缓存之后每个epoch都从缓存里取能省掉大量重复的图像解码。这种情况下Redis做共享缓存是有价值的因为它天然跨进程、跨机器且读取速度比磁盘快很多。但如果你只是单机单卡我的建议是不要让Redis参与数据加载。它引入的额外网络/进程开销往往比本地文件读取更慢而且还得维护一个额外服务调试起来非常痛苦。网络上有些性能优化文章把Redis说得神乎其神实际上属于吃饱了撑的优化。先用好上一条的内存映射和打包方案比什么都强。6.4 预处理位置的选择在getitem里做还是提前做最后一个优化点也是新手最容易困惑的数据预处理到底放在__init__、__getitem__还是训练前一次性完成我的判断标准很简单**如果预处理不受随机种子影响且结果固定尽量提前做如果预处理依赖随机性比如增广就放在__getitem__里做。**比如归一化、去均值、缩放这类固定操作完全可以在离线阶段把所有图片统一处理并保存成tensor训练时直接读tensor。而随机裁剪、随机翻转这类操作必须留在__getitem__里因为每个epoch都需要不同的结果。具体操作时我会给数据集做一次预热for idx in range(len(full_dataset)): sample full_dataset[idx]如果预热阶段就慢得受不了那多半是预处理太重考虑把固定操作离线化如果预热快但训练时GPU很闲那问题大概率在磁盘IO或多进程配置回到本节前面几小节去排查。另外提一个细节如果使用torchvision的v2版本它提供了不少向量化、编译优化的接口比如torchvision.io.decode_image这类底层解码函数在我实测中比PIL的Image.open()快不少。对图像IO有硬需求的朋友可以优先用torchvision.io替代PIL但需要注意的是它返回的是tensor和PIL的转换逻辑要理顺。换掉之后同样的预处理流程整体吞吐经常能快10%到20%。最后说一点个人体会我在不同项目里折腾Dataset和DataLoader也有几年了最大的感受是这一层抽象虽然简单但它是整个训练系统的地基。地基不稳后面模型结构、损失函数再花哨也跑不出稳定的结果。真心建议所有刚入门的朋友不要跳过数据加载这一课花个半天把它彻底吃透后面能帮你省下大量排查训练异常的时间。再分享一个小技巧开发阶段调模型时把num_workers直接设0或者用一个专门的小采样器只抽几百个样本。这样能最快暴露模型和loss的问题等逻辑调通了再把num_workers打开。不要一上来就开8个worker到时候模型报错、数据加载报错混在一起排查难度直接翻倍。最后等你把数据流理顺、模型训完再考虑转ONNX部署那些事——那是另一个故事。先把数据喂好训练这件事才算真正开始。
