ARTICLE DETAIL

资讯详情

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

PyTorch数据加载完全指南:从Dataset到DataLoader的高效实现

PyTorch数据加载完全指南:从Dataset到DataLoader的高效实现 1. 先想清楚为什么要单独写一篇文章讲数据加载先把话说在前头任何一个正经的深度学习项目超过一半的坑都出在数据加载这一层。很多人模型结构写得飞起损失函数背得滚瓜烂熟结果一到训练就开始各种花式报错显存不够、训练贼慢、loss震荡、迭代到一半崩溃……根子往往不在模型而在数据没喂好。我为什么专门挑Dataset和DataLoader出来写一篇因为这俩东西在PyTorch里是被严重低估的两个类。很多人学深度学习上来就怼卷积、怼Transformer把官方教程里的数据加载代码直接当模板复制粘贴从来没想过这里面的设计逻辑是什么。你问他为什么要用Dataset他说大家都是这么写的你问他num_workers设成几他说随便设了个8。这种状态干活十个项目九个要出问题。这篇东西适合两类人第一类刚入门深度学习、想搞清楚PyTorch数据管线到底怎么回事的初学者第二类已经写了一阵子训练代码、但总感觉数据这块使不上力、想系统梳理一遍细节的进阶玩家。我会从底层逻辑讲到实战代码再把踩过的坑全部抖出来争取让你读完就能自己写出一套干净、高效、不崩溃的数据加载模块。先明确一个基本认知PyTorch的数据加载体系本质上是三层分工。第一层是原始数据在硬盘上的存储形态可能是图片文件夹、CSV表格、JSON标注或者别的什么第二层是Dataset负责把硬盘上的文件映射成内存里的样本第三层是DataLoader负责把单个样本打包成成批的张量并且处理并发读取、随机打乱这些脏活累活。记住这个分层后面所有细节都是在这三层上做文章。2. Dataset把数据从硬盘变成样本的核心封装2.1 三个必须实现的方法PyTorch的Dataset类本质是一个抽象接口。你自定义的数据集类要继承torch.utils.data.Dataset然后实现三个方法__init__、__len__、__getitem__。__init__负责初始化通常在这里解析文件列表、读取标注信息、做一些全局性的预处理。__len__返回数据集的总样本数DataLoader要靠它知道一个epoch要迭代多少步。__getitem__接收一个索引值返回对应位置的样本通常是一个(输入, 标签)的元组。有一个最常见的理解误区我必须先说清楚__getitem__才是真正干活的方法它会在训练过程中被反复调用每一次调用都会从硬盘读数据、做变换、转张量、返回结果。而__init__只在创建数据集对象时执行一次。所以永远不要把耗时操作塞进__init__里挨个跑一遍否则你会在数据准备阶段就等到怀疑人生。我给你们看一个我实际用过的图像分类数据集写法这个例子非常经典几乎覆盖了所有基础场景import torch from torch.utils.data import Dataset import os from PIL import Image from torchvision import transforms class ImageClassificationDataset(Dataset): def __init__(self, root_dir, transformNone): root_dir: 数据根目录结构为 root_dir/ class_0/ img_001.jpg img_002.jpg class_1/ img_001.jpg ... self.root_dir root_dir self.transform transform # 扫描所有类别文件夹 self.classes sorted(os.listdir(root_dir)) self.class_to_idx {cls_name: idx for idx, cls_name in enumerate(self.classes)} # 收集所有(图片路径, 标签)对 self.samples [] for cls_name in self.classes: cls_dir os.path.join(root_dir, cls_name) for file_name in os.listdir(cls_dir): if file_name.lower().endswith((.jpg, .jpeg, .png)): file_path os.path.join(cls_dir, file_name) self.samples.append((file_path, self.class_to_idx[cls_name])) def __len__(self): return len(self.samples) def __getitem__(self, idx): file_path, label self.samples[idx] image Image.open(file_path).convert(RGB) if self.transform: image self.transform(image) # 注意这里标签没做transform但如果你做数据增强 # 涉及目标检测或分割任务时标签也要同步变换 label_tensor torch.tensor(label, dtypetorch.long) return image, label_tensor这个代码有几个细节点值得讲。os.listdir出来的文件顺序是不稳定的不同操作系统、不同文件系统下顺序可能不一样所以我在代码里加了sorted保证类别的顺序稳定。千万别小看这个我有一个朋友就在这儿吃过亏他在Windows上训练好的模型换到Linux服务器上推理发现结果全乱了排查半天最后发现是os.listdir的返回顺序变了类别索引对不上了。图片读取用Image.open的时候记得.convert(RGB)。很多灰度图或PNG图是四通道或者单通道的不统一转成RGB后面送给预训练模型的时候就会报通道数不匹配。我见过太多新人在这上面浪费一下午。标签转成torch.long类型是给分类任务的交叉熵损失函数准备的nn.CrossEntropyLoss要求目标张量是长整型你传个float32进去它会直接给你报错。2.2 transform的底层逻辑和常见组合transform参数是PyTorch里一个非常优雅的设计它遵循装饰器模式的思想把数据预处理步骤拆成一个个可独立组装的小模块。你可以在__init__和__getitem__之间灵活切换也可以在外部定义。对于图像任务最常写的一段transform组合是from torchvision import transforms # 训练集transform数据增强 归一化 train_transform transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomCrop(224), transforms.RandomHorizontalFlip(), transforms.ToTensor(), transforms.Normalize(mean[0.485, 0.456, 0.406], std[0.229, 0.224, 0.225]) ]) # 验证集transform只做必要处理不搞数据增强 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]) ])训练集和验证集的transform一定要区分开。训练集可以加随机裁剪、随机翻转、颜色抖动这些增强手段让模型见过更多样的样本起到正则化的作用。验证集不要加任何随机操作保证每次评估结果可复现、可比对。这个是行业里约定俗成的规范你别图省事一套transform打天下。很多人不理解为什么随机裁剪和随机翻转能提升模型泛化能力。你换个角度想数据增强本质上是在制造一个虚拟样本池每张原图通过随机变换可以产生无数种变体模型在训练时永远看不到两张完全一样的图它就很难背答案只能去学习真正稳定的特征。这个逻辑跟正则化是一模一样的。那Normalize里的mean和std又是怎么来的这套数值是ImageNet数据集的统计值几乎所有预训练模型都是基于ImageNet训练的所以你用torchvision.models里的预训练权重时输入数据就要用这套数值做标准化否则输入分布和模型训练时不一致效果直接崩。如果你从头训练自己的模型这个数值可以用你训练集的全局均值和标准差来算或者直接用默认值也能收敛。2.3 大规模数据时的内存与IO优化如果你的数据集特别大比如几万张图片、几十万条文本或者干脆是上百GB级别的视频__init__阶段把所有数据一次性load进内存就是不现实的。这时候有几条经验法则第一__init__里只存路径和索引不读实际数据。这样即使有一百万个样本内存占用也就是几百万个路径字符串的量完全可以接受。真正的数据读取延迟到__getitem__时才做。上面的例子就是按照这个原则设计的。第二如果你的数据本身不大几千张图而且内存充足可以在__init__里把所有图片都读进来缓存到内存里训练时直接从内存取省去每次IO的开销。这个做法通常会带来非常明显的加速。但要注意数据增强还是要放在__getitem__里做因为你不能让同一个epoch里每张图只出现一次变换版本不然增强的意义就没了。第三对于超大文件的数据集比如毫米波雷达数据、基因序列建议使用内存映射或h5py这类持久化格式按需读取指定索引的数据块。这块属于进阶话题知道有这个方向就行真遇到了再深入研究。3. DataLoader真正驱动训练的引擎3.1 常用参数全解析DataLoader本身不负责数据从哪里来它只负责把Dataset喂给它的样本打包、排队、批量分发。官方签名里的参数很多但真正天天用的就那几个我一个个给你讲透。from torch.utils.data import DataLoader dataloader DataLoader( datasettrain_dataset, # 上面自定义的Dataset实例 batch_size32, # 每个batch多少个样本 shuffleTrue, # 每个epoch是否打乱顺序 num_workers4, # 用几个子进程去加载数据 drop_lastFalse, # 最后一个batch不够batch_size时是否丢弃 pin_memoryTrue # 是否锁页内存加速CPU→GPU传输 )batch_size是模型一次前向传播看到的样本数它的选择直接影响训练速度、显存占用和梯度质量。太大容易显存溢出太小训练不稳定且慢。一个经验值是先从num_samples // 100左右起步比如一万个样本就用batch_size64或128然后根据显存占用上下调整。8的倍数比较友好因为很多硬件和库的底层优化都对齐这个数值。shuffleTrue非常关键。如果不打乱模型在每个epoch看到的样本顺序完全相同它会学到这个顺序上的伪规律导致收敛变慢、泛化变差。验证集和测试集的DataLoader则应该设置shuffleFalse因为你评估时不需要随机性且保持顺序有助于复现结果。num_workers这个参数我放到后面的避坑章节详细讲这里先给结论不是越大越好。drop_lastTrue的常规场景是当样本总数不能被batch_size整除时丢弃最后一个不完整的batch。做多卡分布式训练时建议开启因为各卡需要对齐batch数量。单卡训练时丢不丢影响不大。pin_memoryTrue会在主机内存里分配页锁定内存让GPU能更快地读取数据。配合to(device, non_blockingTrue)使用效果更佳。大多数人写代码用的是data.to(device)这是一种同步操作GPU会等数据拷贝完才继续执行改成data.to(device, non_blockingTrue)变成异步拷贝GPU可以边拷贝边计算吞吐量能提升不少。3.2 迭代过程内部七步DataLoader到底是怎么把Dataset和batch联系起来的我用大白话拆解一遍第一步DataLoader创建一个迭代器。第二步如果shuffleTrue它内部会调用torch.utils.data.RandomSampler重新打乱索引。第三步BatchSampler指向RandomSampler或SequentialSampler按照batch_size把打乱后的索引切分成一个个batch的索引列表。第四步多个num_workers子进程分别拿到索引调用Dataset.__getitem__取出原始样本。第五步拿到一批原始样本后DataLoader调用collate_fn把单个样本整理合并成一个batch张量。第六步把整理好的batch从CPU内存搬到GPU显存如果设置了pin_memory。第七步把batch返回给训练循环。正常情况下你训练循环写的是for batch_idx, (images, labels) in enumerate(dataloader): images images.to(device, non_blockingTrue) labels labels.to(device, non_blockingTrue) # 前向传播、反向传播...这里有个挺隐蔽的点DataLoader的迭代过程是惰性的。它并不会一次性把所有batch都准备好而是每次for循环进行到下一次时才去取下一个batch的数据。这样做的好处是内存占用不会随数据总量增长坏处是如果你在训练循环里加了很耗时的操作复杂的数据后处理、长时间阻塞数据加载进程可能会提前把好几个batch的样本准备好并排队等待这时候你会看到内存占用量慢慢往上涨——这是正常的不是内存泄漏。3.3 collate_fn到底在干什么collate_fn是很多人完全没接触过的概念因为默认行为已经能满足90%的需求了。它的本质是Dataloader手里有一堆样本怎么把它们拼成一个张量默认行为是如果__getitem__返回的是元组它就分别对每个位置做stack或者cat。先说默认行为适配的场景。__getitem__返回(image, label)其中image是一个torch.Tensorlabel是另一个torch.Tensor。DataLoader会把这个batch里所有样本的image堆叠成一个四维张量(batch_size, C, H, W)所有label堆叠成一个一维张量(batch_size,)。前提条件是一个batch里所有图片的尺寸必须相同因为张量Stack要求形状严格一致。如果你的数据形状不固定比如变长的文本序列、不同尺寸的图片默认的collate_fn就会报错。这时候必须自定义collate_fndef custom_collate_fn(batch): images torch.stack([item[0] for item in batch], dim0) # 文本序列长度不一致手动padding sequences [item[1] for item in batch] max_len max(seq.size(0) for seq in sequences) padded_sequences torch.zeros(len(sequences), max_len, dtypetorch.long) for i, seq in enumerate(sequences): padded_sequences[i, :seq.size(0)] seq labels torch.tensor([item[2] for item in batch], dtypetorch.long) return images, padded_sequences, labels另一种常见场景是物体检测一张图里有数量不定的目标框框的数量也不一样。这种数据天生没法堆叠成一个固定形状的张量所以通常的做法是用一个列表去接一个batch的标注信息。这种时候你也要在collate_fn里做特殊处理。说一句大实话数据加载系统里最绕的地方就是这个collate_fn各种形状不匹配的报错80%都能追溯到这一步。你只要记住DataLoader某个batch返回的维度不对第一时间去看collate_fn处理逻辑而不是改模型输入层。4. 从零到一一个可运行的完整训练数据管线4.1 完整代码猫狗分类项目实战光讲概念太空我直接给一个完整的示例从Dataset定义到DataLoader配置再到训练循环接入全流程走一遍。这个例子基于一个简单的猫狗分类数据集文件结构是train/cat/xxx.jpg这种形式。代码可以直接复制运行我尽可能把细节都写进注释里。import os import torch import torch.nn as nn import torch.optim as optim from torch.utils.data import Dataset, DataLoader from torchvision import transforms, models from PIL import Image # ---------- 第1步定义Dataset ---------- class CatDogDataset(Dataset): def __init__(self, root_dir, transformNone): self.root_dir root_dir self.transform transform self.classes [cat, dog] self.class_to_idx {cat: 0, dog: 1} self.samples [] for cls_name in self.classes: cls_dir os.path.join(root_dir, cls_name) for file_name in os.listdir(cls_dir): file_path os.path.join(cls_dir, file_name) self.samples.append((file_path, self.class_to_idx[cls_name])) def __len__(self): return len(self.samples) def __getitem__(self, idx): file_path, label self.samples[idx] image Image.open(file_path).convert(RGB) if self.transform: image self.transform(image) return image, torch.tensor(label, dtypetorch.long) # ---------- 第2步定义数据增强 ---------- train_transform transforms.Compose([ transforms.Resize((256, 256)), transforms.RandomResizedCrop(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]) ]) # ---------- 第3步创建数据集和加载器 ---------- train_dataset CatDogDataset(data/train, transformtrain_transform) val_dataset CatDogDataset(data/val, transformval_transform) train_loader DataLoader( train_dataset, batch_size32, shuffleTrue, num_workers4, pin_memoryTrue, drop_lastTrue ) val_loader DataLoader( val_dataset, batch_size64, shuffleFalse, num_workers4, pin_memoryTrue ) # ---------- 第4步定义模型 ---------- model models.resnet18(weightsmodels.ResNet18_Weights.DEFAULT) num_features model.fc.in_features model.fc nn.Linear(num_features, 2) model model.cuda() # ---------- 第5步定义损失和优化器 ---------- criterion nn.CrossEntropyLoss() optimizer optim.Adam(model.parameters(), lr1e-4) # ---------- 第6步训练循环 ---------- num_epochs 10 for epoch in range(num_epochs): model.train() running_loss 0.0 for batch_idx, (images, labels) in enumerate(train_loader): images images.cuda(non_blockingTrue) labels labels.cuda(non_blockingTrue) optimizer.zero_grad() outputs model(images) loss criterion(outputs, labels) loss.backward() optimizer.step() running_loss loss.item() if (batch_idx 1) % 50 0: print(fEpoch [{epoch1}/{num_epochs}], Batch [{batch_idx1}/{len(train_loader)}], Loss: {loss.item():.4f}) epoch_loss running_loss / len(train_loader) print(fEpoch [{epoch1}/{num_epochs}], Average Loss: {epoch_loss:.4f}) # 验证 model.eval() correct 0 total 0 with torch.no_grad(): for images, labels in val_loader: images images.cuda(non_blockingTrue) labels labels.cuda(non_blockingTrue) outputs model(images) _, predicted torch.max(outputs, 1) total labels.size(0) correct (predicted labels).sum().item() val_acc 100 * correct / total print(fValidation Accuracy: {val_acc:.2f}%)4.2 参数选择的计算思路这个示例里几个关键参数的选择我解释一下背后的考量。batch_size32的选择依据是单张ResNet18在224x224分辨率下前向加反向的显存占用大约在1到2GB之间32张图同时算配合2GB到4GB显存的中端显卡可以跑得比较舒服。如果你的显卡显存只有2GB建议降到8或16如果是RTX 3090或4090这种大显存显卡可以开到64甚至128。num_workers4的来由是这个例子中的数据集是普通硬盘上的小图片单张图片解码耗时大约几十毫秒。当num_workers0时所有数据准备工作都在主进程里排队GPU经常空转等数据训练吞吐量大概能差出30%到50%。开4个子进程后数据准备好了在队列里等GPU训练速度就有了显著提升。我实际测试过再往上加到8速度提升就微乎其微了反而可能因为进程切换和内存带宽出现负优化。shuffleTrue配合drop_lastTrue的组合是训练集的标配保证了每个epoch数据顺序完全随机且所有batch大小一致没有最后一个batch样本数不足带来的统计偏差。验证集用shuffleFalse因为评估阶段的顺序不影响结果而且还可以把预测结果按原始顺序保存下来做进一步分析。4.3 模型训练时学习率的联动调整很多人只调batch_size不调学习率这是不对的。根据经验公式学习率应当随batch_size近似线性缩放batch_size翻倍学习率也大致翻倍。你从32换成64之后如果把学习率保持原值收敛速度会变慢如果batch_size从32降到16学习率过高则会导致损失函数震荡。在这个示例里我用Adam优化器加lr1e-4初始值。作为对比如果用SGD加Momentum一般建议从0.01到0.1之间起步配合学习率衰减策略使用。初学者我建议直接用Adam超参数少、鲁棒性强不挑任务。5. 高频踩坑实录这些问题我全遇到过的5.1 num_workers为什么会导致程序卡死这是我在社区答疑时被问得最多的问题之一。现象是DataLoader设置num_workers大于0后训练循环第一轮正常第二轮或者中途程序突然卡住CPU占用率飙升但训练不再迭代。原因通常出在Windows平台的spawn多进程模型上。Windows不像Linux那样用fork来创建子进程它要重新导入主模块如果你的数据集类或者包含数据类定义的代码不在if __name__ __main__:的保护块里子进程就会反复递归导入主脚本导致死锁或者无限重启。解决办法有两个。第一个是把训练代码放进main()函数并用if __name__ __main__: main()包裹这是Windows下PyTorch的标准做法。第二个是Windows上num_workers的设置要保守一般2到4就够了设太大反而容易触发一些底层库的兼容问题。另外我发现很多人以为num_workers等于开ThreadPool线程数这是一个理解性误区。num_workers开的是独立的子进程每个子进程内部才可能有多个线程做IO。进程间通信依赖队列和共享内存当数据量特别大、单样本特别重时进程间的序列化和传输开销反而可能抵消掉并行加载的好处。5.2 数据张量形状不匹配的四大根源训练中遇到RuntimeError: size mismatch的报错时90%的情况是数据管线出了问题而不是模型定义错误。我总结过四种最常见的根源第一种图片通道数不同。有些图片是灰度图单通道有些是RGB图三通道直接stack必然报错。解决办法是在__getitem__里统一做.convert(RGB)。第二种图片尺寸不一致。默认的collate_fn要求一个batch里所有图片形状相同如果你的数据集不经过resize就直接进DataLoader尺寸稍微差一个像素都会炸。解决办法是transform里务必加上Resize或RandomResizedCrop。第三种标签类型不对。CrossEntropyLoss要求标签是torch.long如果你写的是torch.tensor(label)保持默认的int64大概率没问题但如果你代码里其他地方把它转成了float损失函数就会报类型错误。所以最佳实践是在__getitem__里就严格设定标签dtype。第四种batch维度的意外扩张。比如你在transform里用了unsqueeze(0)每个样本就多了一个维度导致最终stack出来的张量维度全部多一维。这种问题比较隐蔽因为它不报错只是训练的每个batch维度不对模型的前向传播会报错。排查方法是打一条print(images.shape)日志看清楚每个阶段的维度变化。5.3 数据增强的过度与不足数据增强不是越狠越好。我把两种极端情况都见过一种是什么增强都不做训练集和验证集长得一模一样模型学到了识别具体图片的作弊特征测试集上泛化效果极差另一种是增强过度比如对医学影像加高斯噪声、随机擦除把病灶区域都擦没了模型根本学不到有效特征。一个折中的、稳健的数据增强组合是随机水平翻转加轻微随机旋转加随机裁剪加标准化。这种组合在大多数视觉任务上都能稳定提升泛化能力不会引入太多噪声。高级增强策略如CutMix、MixUp、RandAugment确实能在分类任务上带来额外收益但需要更多调参经验新手不建议一上来就套用。另外记住**数据增强只在训练集用验证集测试集永远只做必要的resize、归一化和类型转换。**这条原则我反复强调因为真的有很多人在验证集上做了随机裁剪导致验证准确率忽高忽低不稳定。5.4 采样器的高级用法torch.utils.data里除了默认的RandomSampler和SequentialSampler还有几个非常实用但很多人不知道的采样器。WeightedRandomSampler专门处理样本不均衡的问题。当你的数据集中有几十个类别的样本数量差异巨大时普通随机采样会让模型严重偏向多数类。用WeightedRandomSampler时给少数类样本分配更高的采样权重让它们在每个epoch里被抽中的概率更大。SubsetRandomSampler可以用来手动划分训练集和验证集而不需要创建两个Dataset对象。这在你想快速做一个随机划分实验时非常方便from torch.utils.data import SubsetRandomSampler dataset CatDogDataset(data/train, transformtrain_transform) num_samples len(dataset) train_indices list(range(int(num_samples * 0.8))) val_indices list(range(int(num_samples * 0.8), num_samples)) train_loader DataLoader( dataset, batch_size32, samplerSubsetRandomSampler(train_indices), num_workers4 ) val_loader DataLoader( dataset, batch_size64, samplerSubsetRandomSampler(val_indices), num_workers4 )这条路线还有一个好处就是如果你要做K折交叉验证只需要循环不同的索引划分然后替换SubsetRandomSampler传入的列表就行不需要重复创建Dataset。别忘了设置固定的随机种子保证每次实验可复现。5.5 GPU显存碎片化和数据加载速度的权衡训练中大batch配大图很常见一块显存常常被一张图占掉大半。如果DataLoader在pin_memory时还同时加载了多个batch的数据预取到内存中显存里存储的batch张量在多次迭代后会留下不少碎片。虽然PyTorch的缓存分配器会自动整理但频繁申请大块显存确实会导致额外开销。我的做法是大分辨率任务比如遥感影像用较小的batch_size加梯度累积策略不要盲目追求大batch小图任务比如CIFAR、ImageNet的224尺寸可以开到64到128的batch同时配合pin_memoryTrue让数据搬运速度成为瓶颈之前先被GPU计算拖住。你需要做的实验很简单——把num_workers从小到大挨个试一遍观察每个epoch耗时曲线找到吞吐的拐点。6. 聊聊PyTorch新版的数据加载变化如果你用的是PyTorch 2.x版本有个变化值得知道DataLoader内部集成了很多编译优化pin_memory和num_workers的行为在某些场景下会更激进。另外PyTorch 2.0引入了torch.compile可以让模型的执行图被静态编译显存占用更小、训练速度更快。但要注意torch.compile目前对动态输入形状的支持还不完善如果你的数据没法保证固定的长宽比编译优化效果可能打折。现在很多大模型训练框架也开始用IterableDataset替代Dataset。IterableDataset适合处理数据流不固定、无法随机访问的场景比如实时采集的流式数据、在线数据增强。它跟Dataset最大的区别是没有__getitem__取而代之的是实现__iter__方法DataLoader对它的采样逻辑也不同。如果你的数据是无限流式产生的就需要研究这个方向。不过对于绝大多数常规任务掌握上面讲的Dataset DataLoader这套组合拳就完全够用了。先把基础吃得透透的别急着追新概念。7. 关于数据加载的几条经验铁律最后分享几条我实际工作中反复验证过的经验总结你可以直接收进自己的工具箱。数据加载永远不要全部塞进训练循环里。数据预处理该在Dataset的__init__做的就放__init__该在__getitem__做的就放__getitem__该在collate_fn里做的就放collate_fn。图层职责清晰代码才不会越改越乱。先跑通一个小规模数据子集。我真的建议新手在做任何训练前先取几百张图构建一个迷你数据集打印几个batch出来看看形状对不对、数值范围是否合理再上全量数据。你会在调试阶段省下大量时间。print大法永远不过时。在__getitem__里放一行print(idx, file_path)在DataLoader输出的地方打一行print(images.shape, labels.shape)你就能把整个数据管线的每一环都盯住。用完了再删掉成本极低收益极高。数据管线和模型训练建议分开调。先保证数据管线的输出完全正确再开始调模型结构。很多人把两块混在一起排查结果到最后一头雾水。一次只改一个变量这是最朴素的调试哲学。由我个人的习惯来说我特别喜欢在写训练脚本的第一天就把num_workers、batch_size、pin_memory这些参数抽成配置文件后面调参就改配置不改代码。这个习惯看起来不起眼但项目复杂到一定程度你就知道多香了。每次实验的记录都会跟超参数绑定哪个配置跑出来的结果好一目了然。
返回列表