最近在做一个图像分类项目,数据集是自己爬的,大概有10万张图片。我按网上教程写了个自定义Dataset类,里面用了torchvision的transforms做预处理(resize、归一化这些),但训练时每次load数据都特别慢,一个epoch要跑快两个小时。我看别人说用DataLoader的num_workers能加速,但我设了4个worker之后反而报错了,好像是内存炸了。想问下各位大佬,这种自定义数据集的情况,有没有什么常规的优化技巧?比如是不是应该先把图片转成Tensor存起来,还是说transforms里少用点随机操作?或者PyTorch有没有像TF那样可以直接从tfrecord读数据的方式?先谢过大家了。
用PyTorch写自定义Dataset时,数据加载太慢怎么办?
全部回复
共 152 条我之前也踩过这个坑,10万张图全放内存肯定炸,建议先试试把图片预处理成png或者npy存下来,加载时直接读tensor能快不少。另外num_workers不是越大越好,我一般设2-3个,配合pin_memory=True,如果还炸就把persistent_workers=False试试。transforms里那些随机操作其实不算太耗时间,瓶颈多半在磁盘IO和解码上,可以考虑用lmdb或者h5py打包数据,比散文件快很多。最后提一嘴,PyTorch现在也有datapipe,但感觉生态还不成熟,不如自己搞个缓存层实在。
我之前也踩过这个坑,10万图直接读确实要命。建议先别急着上num_workers,把transforms里的随机操作先去掉,用纯resize和归一化试试,瓶颈往往在解码和随机增强上。另外可以试试把图片预处理成uint8的numpy数组或者lmdb存起来,读的时候直接load tensor,速度能快好几倍。你内存炸可能是worker数开太高了,我一般先设2个,配合persistent_workers=True,然后观察内存占用再慢慢加。torchdata那个库也可以看看,但感觉现阶段不如自己手动缓存来得直接。
我之前也踩过这个坑,10万张图全放内存肯定炸,建议先把图片预处理成uint8的tensor存成pt文件,加载时再转float做归一化,能快不少。另外num_workers别一上来就调大,先2个试试,配合persistent_workers=True和pin_memory=True,内存会稳很多。要是还慢,可以试试把resize这类操作放到保存pt之前做掉,训练时只做轻量级增强,这样IO压力小很多。至于tfrecord,PyTorch有webdataset或者lmdb方案,但对你这个规模感觉有点杀鸡用牛刀了。
10万张图用transforms现算确实遭不住,我之前也是卡在这。建议把resize和归一化直接离线跑一遍,存成npy或者lmdb,训练时只做ToTensor,能快好几倍。num_workers报错八成是每个worker都复制了完整transforms链,试试把worker数降到2,或者用persistent_workers=True,内存能省不少。另外随机增强像RandomCrop这种,其实可以在缓存前固定生成几组版本,训练时随机选,效果差不多但快得多。
可以先把图片预处理成npy或者lmdb缓存,IO瓶颈比transforms大多了,worker数调成2试试。
我之前也踩过这个坑,十万张图全走transforms确实扛不住。建议先把resize和归一化提前做掉,存成npy或者lmdb格式,训练时只做ToTensor,能快好几倍。num_workers报错大概率是每个worker都复制了完整数据集索引,试试把shuffle和drop_last设好,再调小worker数到2,或者加persistent_workers=True看看。另外别用太多随机增强,比如随机裁剪这种CPU密集操作,真需要的话可以试试在GPU上用dali库。
先转成LMDB或内存映射试试,IO瓶颈比transforms大多了,我10万张图用这个直接快三倍。
可以试试把预处理后的图片缓存成lmdb或h5py,读取速度能快好几倍,内存不够就分批存。
我之前也踩过这个坑,10万张图全走transforms确实顶不住。你那个内存炸了大概率是num_workers设太高,每张图都走一遍CPU预处理再往GPU搬,内存自然爆,试试把worker降到2或者1,同时把persistent_workers=True加上,能省不少启动开销。
更根本的办法是把预处理结果缓存下来,比如先用一个脚本把所有图片resize并归一化后存成.pt或者.npy文件,训练时直接load tensor,这样transforms只做随机增强那部分轻量操作,速度能快好几倍。另外你提到tfrecord,PyTorch这边可以用webdataset或者lmdb,支持流式读取,但前期建库也要花时间,小项目直接存tensor最省事。
还有个容易忽略的点,检查一下你的transforms是不是有大量随机操作,比如random crop、flip这些,它们每次都要重新计算,如果只是做分类,可以先把固定部分(resize、归一化)在缓存阶段做完,训练时只加随机增强。最后建议你统计一下数据加载和GPU计算的时间占比,如果加载已经超过训练时间,那优化加载才有意义,不然白忙活。
试试把图片预处理后直接存成npy或lmdb,读的时候直接load tensor能快不少,10万张图内存够的话可以全缓存进去。
10万张图还带实时transform,瓶颈大概率在CPU解码上,尤其resize和随机裁剪很吃算力。你可以先试试把图片缩略图缓存成lmdb或h5py格式,训练时直接读数组,能快好几倍。num_workers报错八成是每个worker都复制了完整transforms,试试把worker数降到2,或者用persistent_workers=True,内存会稳一些。另外随机增强别全堆在Dataset里,有些操作比如随机翻转可以放到GPU上用tensor做,省CPU开销。
10万图全放内存不现实,先试试把图片resize成小尺寸再存成lmdb或h5py,读取会快很多。
我之前也踩过这个坑,10万图直接上transforms确实扛不住。建议先把图片离线处理成uint8的tensor存成.pt文件,或者干脆用lmdb/内存映射,训练时只做归一化,能快好几倍。另外num_workers报错大概率是共享内存不够,试试把persistent_workers=True加上,或者把batch_size调小点,别一上来就开4个。还有个思路是换个轻量级的jpg解码库,比如用pillow-simd替代原版pillow,随机裁剪翻转这些操作其实没那么耗时间,瓶颈主要卡在磁盘IO上。
我之前也踩过这个坑,10万图全放内存确实会炸,建议先用PIL读图存成jpg或png的路径,然后transforms里别用太多随机操作,像resize和归一化可以放dataset外面做。另外worker数不是越大越好,4个不够就试试2个,或者把num_workers设成0跑一下看是不是内存问题,还有可以试试把图片预处理成numpy缓存到磁盘,训练时直接读npy,速度能快不少。
我之前也踩过这个坑,十万张图raw读确实要命。建议先把所有图片预处理成uint8的numpy或者直接存成.pt的tensor,省得每次都要解码和resize,能快好几倍。num_workers炸内存大概率是每个worker都复制了完整的数据集引用,试试把prefetch_factor调小点,或者用persistent_workers=True看看。另外transforms里那些随机操作尽量放CPU上做,别全堆在GPU那边,不然IO和计算容易互相卡脖子。
大概率是IO瓶颈,先把图片缩成小图缓存成lmdb或h5py格式,读取速度能快好几倍。
10万图建议直接打包成lmdb或h5py,io瓶颈比transforms大多了,worker数得看着内存调。
我之前也踩过这个坑,10万张图全走transforms确实扛不住。建议先把resize和归一化这些固定操作提前做掉,存成numpy或者lmdb格式,训练时只读tensor,能快不少。另外num_workers报内存炸,可以试试把persistent_workers=True加上,或者把worker数量降到2,再配合prefetch_factor调小一点。至于随机操作,别全砍,留个随机翻转之类的就行,不然泛化会受影响。
10万张图确实不该这么慢,你八成是卡在IO上了。我建议先把图片预处理成uint8的numpy数组或者lmdb存起来,训练时候直接读内存映射,能快好几倍。transforms里的随机操作其实影响不大,真正吃时间的是每次从磁盘解码,你可以试试用torchdata的DataPipes做流的并行预处理。至于num_workers报错,多半是每个worker默认会复制一份完整dataset,内存直接翻倍,可以试试persistent_workers=True或者把worker数量降到2看看。
10万图一个epoch两小时确实不正常,瓶颈多半不在transforms而在磁盘IO。建议先把图片缩到合适尺寸再用torch.save缓存成uint8的tensor,加载时直接读文件省去解码;另外num_workers炸内存可以试试把prefetch_factor调小,或者用persistent_workers=True。tfrecord的话其实有webdataset这个库,思路类似但更轻量,不过你这规模先优化缓存应该就够了。