最近在做一个图像分类项目,数据集是自己爬的,大概有10万张图片。我按网上教程写了个自定义Dataset类,里面用了torchvision的transforms做预处理(resize、归一化这些),但训练时每次load数据都特别慢,一个epoch要跑快两个小时。我看别人说用DataLoader的num_workers能加速,但我设了4个worker之后反而报错了,好像是内存炸了。想问下各位大佬,这种自定义数据集的情况,有没有什么常规的优化技巧?比如是不是应该先把图片转成Tensor存起来,还是说transforms里少用点随机操作?或者PyTorch有没有像TF那样可以直接从tfrecord读数据的方式?先谢过大家了。
用PyTorch写自定义Dataset时,数据加载太慢怎么办?
全部回复
共 152 条10万张图用默认的ImageFolder加transforms确实容易卡,尤其是resize这种操作很吃CPU。你试试把图片预处理提前做好,比如用albumentations或者PIL先统一resize成固定尺寸,然后直接存成numpy或者pickle,DataLoader里只做to tensor和归一化,这样能省不少时间。num_workers报错可能是内存不够,可以降到2或者1试试,同时把pin_memory=True打开,另外检查下你的transform里有没有用RandomCrop这类随机操作,太频繁也会拖慢速度。torch没有tfrecord的等价物,但可以考虑用lmdb或者h5py把数据打包成二进制文件,随机读取会比文件系统快很多。
这问题我也踩过坑,10万张图用默认方式硬扛确实要命。你提到的转Tensor存起来其实是个好思路,预处理后的图片直接保存成.pt文件,训练时用torch.load加载,能省掉每次重复的transforms计算,内存占用也稳定很多。不过要注意,如果用了随机增强(比如随机裁剪、翻转),得在加载后在线做,否则会破坏数据多样性。另外num_workers报错大概率是内存爆了,可以试试把worker数量降到2或3,同时把prefetch_factor设小一点,比如2,别让worker一股脑塞满内存。还有一个冷门技巧:用lmdb或者h5py这类二进制格式存数据,读取速度比散列的小文件快一个数量级,我试过把10万张图打包成几个lmdb文件,训练时间直接砍半。至于tfrecord那种方案,PyTorch官方没有,但社区有像WebDataset这样的库,支持流式读取tar包,你可以看看合不合胃口。
你这情况我太熟了,10万张图用基础方案确实会卡得怀疑人生。num_workers报错大概率是内存不够,每个worker会复制一份数据,4个worker加上主进程大概要吃掉5倍的内存,建议你降到2个试试,或者把batch_size调小一点。另外我强烈建议你把预处理后的图片先存成.pt文件,就是用torchvision的transforms做完resize和归一化之后直接torch.save(),训练时load进来就只剩ToTensor那步了,能快好几倍。transforms里的随机增强像RandomCrop、ColorJitter这些可以保留,但建议你用GPU做在线增强,比如用albumentations库配合PyTorch,CPU只做最基本的解码。还有个骚操作是提前把图片转成LMDB或者HDF5格式,随机读取速度比文件系统快得多,尤其是小文件多的情况。至于tfrecord,PyTorch有WebDataset库,功能类似但生态不如TF成熟,不过如果你愿意折腾,用Apache Arrow或者NVIDIA的DALI也能达到类似效果。
我试过类似的场景,10万张图的话,推荐先把图片预处理成tensor或者h5py格式存起来,这样读取速度能快不少。num_workers报错大概率是内存不够,可以试试把batch_size调小点或者用prefetch_factor控制预取数量。另外torchvision的transforms里少用RandomCrop这类随机操作,能减轻不少加载压力。
预加载+lmdb存成二进制格式能解决IO瓶颈,或者试试用albumentations代替torchvision的transforms会快很多。
建议先离线把图片转成LMDB或HDF5格式存起来,训练时直接读序列化数据,能省去大量IO时间。
我之前也遇到过类似的问题,num_workers设太高确实容易爆内存,建议先降到2试试,同时把batch size调小一点。另外可以先把图片预处理成tensor存成.pt文件,训练时直接load,能省掉大部分IO和transforms的时间。还有就是检查下transforms里是不是有太多随机操作,像RandomResizedCrop这种在数据加载时做很费时,可以提前离线做好。
num_workers炸内存很常见,可以试试调低到2或者1,同时把batch_size减小一点。另外我建议你先把图片预处理成.pt文件存起来,训练时直接加载Tensor,能省掉resize这些操作的时间。随机操作别全删,但像RandomResizedCrop这种可以改成固定大小裁剪,实测能快不少。
学到了,感谢分享!
可以试试先把图片预处理后存成.pt文件,训练时直接加载Tensor能快不少。
10万张图确实不少,我遇到过类似问题,一个常用技巧是用lmdb或h5py先把图片批量转成内存映射格式,这样读取比从硬盘一张张IO快很多。另外num_workers报内存错误的话,试试把prefetch_factor调小到2,同时pin_memory=False也能省点显存。transforms里的随机操作倒不是主要瓶颈,但RandomResizedCrop这类如果配合多进程,可以放一部分到on_epoch_end里预计算。
10万张图确实不少,我之前也踩过这个坑。num_workers报错很可能是你机器内存不够,建议先降到2试试,同时把batch_size也调小点。另外可以试试先把图片预处理成.pt或者.npy格式存下来,训练时直接加载tensor,能省掉反复resize和归一化的时间。至于transforms里的随机操作,像随机翻转这些留几个关键的就行,别全堆上,不然IO和CPU都扛不住。
可以把图片预处理后存成numpy或lmdb格式,加载时直接读tensor能快不少。num_workers设4就爆内存的话,试试把prefetch_factor调小点。
我之前也碰到过类似问题,num_workers设太高确实容易内存炸,建议先从1或2慢慢试,同时检查下transforms里是不是有像RandomCrop这种耗资源的操作,能少用就少用。另外把图片预处理成npy或pt文件存起来是个好办法,load的时候直接读tensor比每秒解压图片快很多。至于tfrecord那种格式,PyTorch有官方的WebDataset或者你可以用lmdb自己做内存映射,读写效率会高不少。
试试用lmdb或h5py把图片打包成二进制文件,内存加载会比IO快很多,worker数设成cpu核心数别太高。
说实话10万张图一个epoch两小时确实有点离谱了,我怀疑瓶颈可能不在transforms本身。你设num_workers=4炸内存,大概率是因为每个worker都会复制一份完整的Dataset对象,如果图片原始分辨率太高或者没做缓存,四个worker同时加载会把内存撑爆。我自己的做法是先在预处理里把图片统一resize到256x256这种小尺寸,然后保存在内存里或者用lmdb/arrow这种格式存成二进制,读取的时候直接load tensor而不是从文件读再resize。另外transforms里的随机操作像RandomCrop、ColorJitter这些确实会增加计算量,但通常不会慢到这种程度,主要还是IO和格式转换拖后腿。你可以试试把图片先转成npy或者pt文件,读的时候用torch.load直接读取tensor,配合num_workers=2或者3,应该能快很多。至于像tfrecord那样的方案,PyTorch官方有个WebDataset库,或者你可以用h5py把数据打成h5文件,读取速度比散图快一个量级。
可以先试试把图片缓存成LMDB或HDF5格式,读起来快很多,worker数建议从2慢慢往上调。
说实话你这个情况我前段时间刚踩过坑,10万张图一个epoch两小时确实太折磨了。num_workers报内存炸大概率是因为每个worker都会复制一份transforms里的随机操作状态,加上图片解码后的原始数据本来就大,4个worker同时干很容易把内存撑爆。我建议你先别急着转Tensor存,那个反而会让单张图片体积变大——除非你存成.pt格式然后配合内存映射,但小数据集还行,10万张操作起来挺麻烦的。
更实际的解法是:把transforms里那些CPU密集型的操作(比如resize、归一化)先离线跑一遍,把预处理好的图片存成.pt或者.npy文件,训练时Dataset里直接加载这些预处理后的张量,这样DataLoader几乎只做IO操作,内存压力小很多,而且你甚至可以设到8个worker。另一个技巧是检查下图片读取库,PIL默认的Image.open其实挺慢的,换成cv2.imread或者用imageio之类的库能肉眼可见提速。至于tfrecord那种方案,PyTorch这边可以用WebDataset或者FFCV,但学习成本略高,你先把离线预处理加多worker调通,应该就能把epoch时间压到20分钟以内。
我之前也遇到过类似的问题,10万张图确实不少。建议你把图片预处理后的结果先存成.pt或者.npy文件,这样训练时直接加载tensor,能省下每次都要做transforms的时间。另外num_workers报错大概率是内存不够,可以试试把batch_size调小一点,或者用prefetch_factor控制一下预取的数量。PyTorch没有tfrecord那种格式,但可以用webdataset或者lmdb来做加速,效果还不错。
我之前也踩过这个坑,num_workers设太高确实容易爆内存,建议先降到2试试,同时把batch size调小一点。另外你那个transforms里如果有RandomCrop这类随机操作,可以先改成固定的Resize,等训练稳定了再加回去。还有一个骚操作是把预处理完的图片直接缓存成.pt文件,读的时候用torch.load比反复用PIL快很多,10万张图大概能省一半时间。