最近在做一个图像分类项目,数据集是自己爬的,大概有10万张图片。我按网上教程写了个自定义Dataset类,里面用了torchvision的transforms做预处理(resize、归一化这些),但训练时每次load数据都特别慢,一个epoch要跑快两个小时。我看别人说用DataLoader的num_workers能加速,但我设了4个worker之后反而报错了,好像是内存炸了。想问下各位大佬,这种自定义数据集的情况,有没有什么常规的优化技巧?比如是不是应该先把图片转成Tensor存起来,还是说transforms里少用点随机操作?或者PyTorch有没有像TF那样可以直接从tfrecord读数据的方式?先谢过大家了。
用PyTorch写自定义Dataset时,数据加载太慢怎么办?
全部回复
共 152 条我之前也踩过这个坑,10万张图全走transforms确实顶不住。你可以试试先把resize和归一化后的结果直接缓存成.pt或者npy文件,训练时只做ToTensor,速度能快好几倍。另外num_workers报内存炸了不一定是你RAM不够,可能是每个worker都复制了一份完整数据集引用,试试把persistent_workers=True加上,或者把batch_size调小点看会不会好。还有个思路是如果你不太依赖随机增强,干脆预处理全离线做完,训练时纯读张量,省得每次都要算。
我之前也踩过这个坑,十万张图直接读确实要命。建议你别用随机裁剪和翻转了,先离线把resize和归一化做完,存成numpy或者lmdb格式,训练时只做张量化和简单的tensor变换,速度能快好几倍。另外num_workers不是越大越好,内存炸的话试试把persistent_workers和prefetch_factor调小,或者干脆用2个worker加pin_memory。
你这情况我也踩过坑,十万张图全走transforms确实扛不住。建议先把resize和归一化这些固定操作提前跑一遍,存成npy或者lmdb格式,训练时直接读预处理好的数据,能快一大截。num_workers报内存炸多半是每个worker都复制了完整数据集引用,试试把worker数降到2,或者用persistent_workers=True看下。另外别用太多随机增强,像随机裁剪这种可以留到训练时做,但别全堆在Dataset里。
我之前也踩过这个坑,十万张图全走transforms确实扛不住。你那个内存炸了大概率不是num_workers本身的问题,而是每个worker都会复制一份dataset的副本,如果transforms里有大量随机操作或者缓存了什么东西,内存就翻倍涨了。我建议先把图片统一预处理成固定尺寸的tensor存成.pt文件,或者干脆转成npy,训练的时候直接load进内存,IO瓶颈一下就没了。另外transforms里的随机增强别放太多,尤其像RandomResizedCrop这种每次都要解码再裁剪的,特别吃CPU,能挪到GPU上用torchvision的GPU算子就挪过去。还有个偏方是试试把图片转成LMDB或者HDF5格式,读取速度比散落的小文件快很多,毕竟十万张图每次epoch都要几千次小文件IO,系统调用开销不是开玩笑的。最后,如果实在想用DataLoader加速,先试试worker=2,配合pin_memory=True和persistent_workers=True,有时候比盲目调大num_workers更稳定。PyTorch没有像TFRecord那种官方格式,但社区有webdataset,也挺好用。
我之前也踩过这个坑,10万图用jpg硬读确实顶不住。建议先做一步离线预处理,把图片全部resize成统一尺寸再存成png或者npy,虽然占点磁盘但训练时省掉大部分解码时间。另外transforms里的随机操作尽量放lightning或者GPU上做,别全堆在CPU端。num_workers报错大概率是共享内存不够,试试把persistent_workers=True和prefetch_factor调小点,或者直接降到2个worker看看。还有个野路子是把小图打包成lmdb或者h5py,读取比散文件快很多,你可以试试。
我之前也踩过这个坑,10万张图raw读的话确实能把你等哭。你那个num_workers=4内存炸了,大概率是每个worker都会复制一份完整的Dataset引用,加上transforms里的随机操作在CPU上跑,内存直接翻好几倍,建议先看看是不是transforms里放了太多需要额外缓存的东西。我后来是把图片全压成256x256的jpg或者png存成lmdb,读取速度能快一个数量级,而且内存占用稳定很多,不过lmdb那个库有点老,你得先把键值对设计好。至于直接把Tensor存下来,如果是uint8的话可以试试,但要注意别一次性全load进内存,10万张图就算压缩了也可能好几个G,最好做成memmap或者分片存。transforms里那些随机翻转、颜色抖动确实会拖慢速度,但瓶颈主要在IO和decode上,你可以开个计时器测一下,大概率发现decode占了80%时间。PyTorch没有直接等效tfrecord的东西,但有个webdataset库能用tar包流式读,效果还行,但也要配合num_workers用。另外你设num_workers报错,也可能是你用了Windows,多进程在Windows下容易出问题,可以试试把if name=='main'包好,或者改用persistent_workers=True加prefetch_factor调小点。还有个骚操作是先把所有图片缩成小图放内存当缓存,训练时先读缓存,不够清晰再补原图,但10万张全内存估计吃不消,可以只缓存高频类别。总之IO这块得profile,别一上来就堆worker,先试试单worker用pillow的线程模式,或者换成opencv的imdecode,往往比torchvision自带的loader快不少。
我之前也踩过这坑,10万张图全走transforms确实扛不住。建议先把resize和归一化提前做掉,存成.pt或者npy格式,训练时直接load tensor,能快好几倍。num_workers报错多半是内存爆了,试试把persistent_workers=True加上,或者干脆降到2个,配合pin_memory=True。另外别用太多随机增强,像RandomCrop这种每次都要读图算,特别拖速度,真要随机就放轻量级的。
我之前也踩过这个坑,十万张图如果每张都现读现做resize和归一化,瓶颈其实在磁盘IO和CPU预处理上,GPU反而在空转。你设4个worker内存炸,大概率是每个worker都会复制一份完整的Dataset引用,再加上transforms里的随机操作会额外生成中间变量,内存自然扛不住。一个比较直接的办法是先把图片批量预处理成.pt或者.npy文件,比如把resize和归一化提前做完存下来,训练时Dataset里只做读取和ToTensor,这样速度能快好几倍。另外你问的tfrecord类似物,PyTorch生态里可以用WebDataset或者简单的LMDB,但前期转换也有成本,我倒觉得先试试把num_workers降到2,配合persistent_workers=True和pin_memory=True,再给DataLoader加个prefetch_factor,说不定就能缓解。还有个小技巧,如果图片尺寸统一,可以试试把transforms里的RandomResizedCrop这类随机操作放到训练循环外面,用albumentations库替代,它内部优化过,比torchvision的纯Python实现快不少。最后提醒下,别把所有数据一次性load进内存,十个G的Tensor也够呛,分片缓存或者用内存映射mmap方式读取会更稳。
之前也踩过这个坑,10万张图每次现读现处理确实扛不住。你提到把图片转成Tensor存起来,这个方向是对的,但更推荐直接存成memmap或者LMDB格式,因为单张npy文件小文件多了随机读取也会卡IO。我自己是把预处理后的图像直接拼成一个大数组,用np.memmap做映射,配合DataLoader的num_workers=8(注意别超过物理核数),内存占用反而比读原图小,因为省了解码和resize的中间缓存。
另外transforms里的随机操作(比如随机裁剪、翻转)确实会增加不少CPU开销,但如果你需要数据增强,建议把随机操作放在worker进程里做,主进程只做轻量归一化。报内存炸了很可能是每个worker都复制了一份完整数据集索引,试试把dataset里存图片路径而不是加载到内存,让每个worker按需读取。PyTorch没有直接对应tfrecord的东西,但可以用WebDataset,它把数据打包成tar格式读取效率很高,社区里挺多人用这个。
还有个冷门技巧:如果你的图片尺寸统一,可以提前把所有图片解码成uint8数组存成hdf5,h5py支持并行读取,速度比单张读jpg快一个量级。最后检查下DataLoader的pin_memory和persistent_workers选项,这两个对内存管理也有影响,有时候报错是它们没配好,不是worker数量本身的问题。
我之前也踩过这个坑,十万张图全走transforms确实扛不住。建议把resize和归一化提前做掉,存成预处理好的npy或者jpg,训练时Dataset里只做ToTensor,能快一大截。num_workers报错大概率是每个worker都复制了一份完整数据索引,试试把persistent_workers=True加上,或者把batch_size调小点,内存压力会小很多。另外随机操作别全放transforms里,像RandomCrop这种可以改成在加载时只做一次,或者用albumentations库,速度比torchvision快不少。
10万张图一个epoch两小时,瓶颈大概率在硬盘IO和PIL解码上,不是transforms本身。可以试试把图片提前解码成小尺寸的numpy数组或者LMDB打包,训练时直接读内存,速度能快好几倍。num_workers开到4就爆内存,估计是每个worker都复制了一份数据索引,加上PIL缓存没释放,适当调小batch或者用persistent_workers试试。另外NVIDIA的DALI对图像解码加速挺明显,值得折腾一下。
试试把图片预处理后存成numpy或lmdb,比每次读原图解码快很多,num_workers内存炸可能是batch太大。