最近在做一个图像分类项目,数据集是自己爬的,大概有10万张图片。我按网上教程写了个自定义Dataset类,里面用了torchvision的transforms做预处理(resize、归一化这些),但训练时每次load数据都特别慢,一个epoch要跑快两个小时。我看别人说用DataLoader的num_workers能加速,但我设了4个worker之后反而报错了,好像是内存炸了。想问下各位大佬,这种自定义数据集的情况,有没有什么常规的优化技巧?比如是不是应该先把图片转成Tensor存起来,还是说transforms里少用点随机操作?或者PyTorch有没有像TF那样可以直接从tfrecord读数据的方式?先谢过大家了。
用PyTorch写自定义Dataset时,数据加载太慢怎么办?
全部回复
共 152 条我之前也踩过这个坑,十万张图全走transforms确实慢,尤其是随机resize那类操作特别吃CPU。你可以试试先把预处理后的图片存成npy或者lmdb格式,训练时直接读数组,能快好几倍。另外num_workers报错不一定是内存不够,可能是你用了Windows,得把代码放到if name == 'main'里保护一下,或者把worker数降到2试试看。至于tfrecord,PyTorch有webdataset库或者直接写个简单的LMDBReader,但感觉对你这个规模来说,存成npy是最省事的。
我之前也踩过这个坑,10万张图全走transforms确实扛不住。建议先把预处理好的图像用torch.save存成.pt文件,或者直接转成numpy的uint8数组落盘,训练时只做ToTensor和归一化,能快好几倍。
另外num_workers报错大概率是每个worker都复制了一份完整的数据集索引,内存直接翻倍。可以试试把worker数降到2,或者用persistent_workers=True,再不行就上lmdb或者h5py这种内存映射格式。
还有个小技巧,如果硬盘是机械的,强烈建议把图片打包成zip或者用WebDataset,减少随机IO的寻道开销。transforms里的随机操作其实影响不大,主要瓶颈在解码和IO,别全甩锅给它。
我之前也踩过这个坑,10万张图其实不算特别大,但瓶颈多半不在transforms本身,而是每次都在做磁盘IO加解码。你试试把图片提前处理成uint8的numpy数组或者直接存成.pt文件,加载的时候一次性mmap进来,能快非常多。另外num_workers=4报内存炸,大概率是每个worker都会复制一份dataset的引用,如果你的transforms里有随机操作,每个worker还会额外维护自己的随机状态,内存就翻倍了。可以试试把workers降到2,同时把prefetch_factor调成2,或者用persistent_workers=True,这样能减少重复初始化的开销。还有个偏方是,如果机器内存够大,干脆把所有图片解码后塞进RAM disk或者直接用torch.load缓存成内存张量,训练时走内存读取,基本能快到飞起。至于tfrecord那种格式,PyTorch这边有webdataset或者ffcv,但学习成本高,不如先试试imagefolder+datapipe的组合。对了,transforms里的随机操作比如RandomCrop,建议改成先统一resize再随机裁剪,别在每次load时都做两次插值,这个也吃CPU。最后检查一下你的图片是不是都特别大,如果原图是4K分辨率,resize到224的消耗远大于你想象的,可以提前用PIL批量缩到512×512再存,训练时二次resize就快了。
试试把图片预处理完存成lmdb或h5py,读取时直接load tensor,能快好几倍。
把图片预处理完缓存成lmdb或h5py格式,读取速度能快好几倍,内存不够就换memmap。
我之前也踩过这个坑,十万张图直接读确实慢,建议先把图片预处理成uint8的numpy数组存成.pt文件,训练时直接load tensor,能快一大截。transforms里的随机操作可以放在训练循环里做,别全塞在Dataset里。另外num_workers报错大概率是内存爆了,试试把persistent_workers和prefetch_factor调小点,或者用2个worker先跑通。至于tfrecord,PyTorch这边有WebDataset或者简单的LMDB,但对你这个场景可能有点过度设计了。
我之前也踩过这个坑,十万张图确实不能直接硬扛。你提到的先转成Tensor存起来是个思路,但别存成单个文件,最好用lmdb或者h5py分块存,读的时候按索引取,IO压力会小很多。另外transforms里那些随机操作确实费CPU,尤其是随机裁剪和翻转,你可以试试先把resize和归一化这种固定操作在预处理阶段做掉,只把随机部分留在训练时,能省不少时间。num_workers报错大概率是内存爆了,因为每个worker都会复制一份数据集索引,你可以把persistent_workers=True加上,同时把prefetch_factor调小点,比如2,这样能缓解。还有个偏方是先把图片解码成numpy数组缓存到内存里,如果内存够大的话,10万张224x224的RGB图大概也就15GB左右,一次性load进去后面就快了。至于tfrecord,PyTorch这边可以用webdataset或者ffcv,格式上跟tfrecord类似,但配置起来有点麻烦,我建议你先试试优化现有流程,别急着换框架。最后检查一下你的磁盘是不是机械硬盘,如果是的话,换SSD提升比调参还明显。
我之前也踩过这个坑,10万张图直接读确实扛不住。建议先把所有图片预处理成tensor存成.pt或者h5py格式,训练时直接load tensor,能快很多。另外transforms里的随机操作别放IO线程里,能省则省,随机性留在训练循环里做也行。num_workers报错大概率是内存溢出,试试把prefetch_factor调小,或者用persistent_workers=True,别一次性开太多。
我之前也踩过这个坑,10万张图全走jpg解码加随机变换确实扛不住。建议先把图片预处理成png或npy存起来,顺便把resize和归一化提前做了,训练时只读tensor省掉transforms那部分开销。num_workers报错大概率是每个worker都复制了一份完整dataset,内存峰值翻倍,试试把persistent_workers=True加上,或者把batch_size调小点。另外你如果不太吃随机增强,可以把crop和flip这类操作放到GPU上用张量实现,CPU只做IO,速度能快不少。
你要是懒得改代码,还有个偏方是把图片直接打包成lmdb或h5py格式,读取速度比散文件快很多,PyTorch虽然没tfrecord那么现成,但这两个格式用起来也不难。我项目里最后是npy+多进程预加载解决的,epoch时间直接从两小时压到二十分钟,你可以试试。
我之前也踩过这个坑,10万张图其实不算多,但瓶颈多半在IO和解码上。建议先试试把图片直接存成npy或者lmdb格式,加载的时候用np.load或者内存映射,能快不少。另外transforms里少用随机操作确实有用,特别是随机裁剪这种,可以放batch之后再做。num_workers别贪多,先试2个,配合persistent_workers=True,有时候能缓解内存问题。还有个野路子,如果机器内存够大,干脆把整个数据集预加载到内存里,epoch之间直接读Tensor,速度起飞。
10万张图一个epoch跑两小时确实有点离谱了,我怀疑瓶颈不在transforms本身,而在每次都在磁盘上做随机IO。你试试把num_workers调回0,然后先在内存里用lmdb或者h5py把图片打包好,读取时直接索引,这样能快很多。至于transforms里的随机操作,其实CPU开销没那么大,除非你用了大量的随机擦除或者cutout,不然不是主要矛盾。
我自己的经验是,先做一个小的缓存层,把resize后的图片存成uint8的numpy数组,训练时只做ToTensor和归一化,这样能把IO压力彻底降下来。另外你设4个worker内存炸,大概率是每个worker都在复制完整的dataset引用,如果数据本身已经加载进RAM了,那会直接撑爆,建议用persistent_workers=True配合共享内存,或者干脆用iterable dataset分片读取。
PyTorch没有像tfrecord那样一站式的东西,但可以用webdataset或者ffcv,后者据说能压到接近纯GPU上限。不过最省事的办法还是先转成512x512的压缩JPEG,然后全部塞进一个tar包,用torchdata的DataPipes去读,我试过效果不错。你现在的transforms如果每个epoch都在做随机裁剪,那最好把裁剪前的resize提前做掉,不然等于每次重复算。最后检查一下你的数据存储格式,如果是几十万个小文件散在目录里,那linux的inode缓存也会拖慢速度,可以先合并成几个大文件。
我之前也踩过这个坑,10万张图全走jpg解码确实扛不住。建议你先试试把图片预处理成uint8的numpy数组或者直接存成png的tensor格式,读写快很多,内存也稳。另外num_workers不是越大越好,可以先从2开始调,配合persistent_workers=True和pin_memory=True,能缓解内存抖动。transforms里的随机操作尽量放轻量级的,像resize这种可以提前做一次缓存,别每次epoch都重复算。
我之前也踩过这个坑,10万图全走jpg实时解码确实扛不住。建议先把图片预处理成uint8的tensor或numpy存成.pt或.h5,训练时直接load内存映射,能快好几倍。另外num_workers报错大概率是workers数×batchsize×图片尺寸超了共享内存,可以试试把persistent_workers=True和prefetch_factor调小点,或者干脆用2个worker加pin_memory。transforms里的随机操作其实还好,主要瓶颈在IO和解码,别全堆在dataset里,能离线做的都离线做掉。tfrecord的话,pytorch也有webdataset或者直接写个多进程prefetch,但你这场景先缓存成张量最省事。
我之前也踩过这个坑,十万张图其实不算多,但问题往往出在transforms里的随机操作上,它们会强制在CPU端同步执行,拖慢整个流水线。建议把resize和归一化这类固定的预处理直接离线做一遍,存成npy或者lmdb格式,训练时只加载和做随机增强,能快好几倍。另外num_workers报错不一定是内存不够,试试把persistent_workers=True加上,或者把worker数量降到2,同时把batch_size调小点,很多时候是数据队列积压导致的。至于tfrecord,PyTorch这边可以用WebDataset或者简单的h5py,但我觉得最省事的还是先预处理成内存映射格式。
10万张图的话瓶颈大概率不在transforms那点随机操作,先把图片直接缓存成jpg缩略图或者npy格式试试,读盘速度能快好几倍。num_workers报内存炸了可能是你设的4个worker每个都复制了一份完整数据集索引,试试把persistent_workers=True加上,或者把batch_size调小点配合worker数量。另外你如果内存够大,干脆用lmdb或者h5py把整个数据集打包成二进制,训练时直接随机读,比从文件夹一张张找快得多。我之前遇到过类似问题,最后是先把所有图片resize到256存成uint8的tensor,加载时只做normalize,epoch时间直接砍半。
先试试把图片预处理结果缓存成本地文件,训练时直接读,能省一大截时间。
我之前也踩过这个坑,十万张图全走transforms确实慢,尤其随机操作会拖垮CPU。你试试把预处理结果直接缓存成.pt文件,训练时用map-style dataset读,能快不少。另外num_workers报错大概率是pin_memory设太高了,改成0或者2先试试,内存炸了就先别开persistent_workers。还有个土办法,把图片缩到固定尺寸再存成npy,不光省空间,加载时也省了resize的开销。
我之前也踩过这个坑,10万张图全走JPG解码确实顶不住。建议先把图片预处理成png或直接存成npy/tensor格式,读的时候省掉resize和归一化,能快一半以上。另外num_workers不是越大越好,我一般设成CPU核心数的一半,然后配合persistent_workers=True,内存会稳一点。至于transforms里的随机操作,训练时留个RandomCrop和Flip就够,别堆太多,不然CPU pipeline容易成瓶颈。还有个笨办法,如果你显存够大,干脆把整个数据集一次性load进内存,跑起来跟飞一样。
我之前也踩过这个坑,10万张图直接硬扛肯定不行。建议先把图片预处理成uint8的tensor存成.pt文件,加载时省掉大部分transforms开销,随机操作放到训练循环里做。另外num_workers报错大概率是每个worker都在复制一份数据集引用,试试把persistent_workers=True加上,或者用prefetch_factor调小一点。还有检查下是不是磁盘IO瓶颈,SSD和机械硬盘速度差好几倍。
我之前也踩过这个坑,10万张图全放内存肯定爆,建议先把图片缩到合适尺寸存成png或者npy,读取时再load,能快不少。另外transforms里那些随机操作确实拖速度,尤其resize和归一化,可以试试先离线处理好存成tensor,训练时只做简单增强。num_workers报错大概率是内存不够,可以试试把persistent_workers和prefetch_factor调小,或者用2个worker先跑通。还有个思路是直接上WebDataset或者用lmdb打包,读取效率比散文件高很多。