最近在做一个图像分类项目,数据集是自己爬的,大概有10万张图片。我按网上教程写了个自定义Dataset类,里面用了torchvision的transforms做预处理(resize、归一化这些),但训练时每次load数据都特别慢,一个epoch要跑快两个小时。我看别人说用DataLoader的num_workers能加速,但我设了4个worker之后反而报错了,好像是内存炸了。想问下各位大佬,这种自定义数据集的情况,有没有什么常规的优化技巧?比如是不是应该先把图片转成Tensor存起来,还是说transforms里少用点随机操作?或者PyTorch有没有像TF那样可以直接从tfrecord读数据的方式?先谢过大家了。
用PyTorch写自定义Dataset时,数据加载太慢怎么办?
全部回复
共 152 条试试把图片预处理后直接存成npy或lmdb,读取时省掉transform,速度能快好几倍。
10万张图跑两小时一个epoch确实不太对劲,我怀疑瓶颈不一定全在transforms上,你先看看是不是图片本身尺寸太大,如果原图是几千像素那种,resize本身就很吃CPU。我之前遇到过类似情况,把图片先统一缩放到短边256再存成JPEG,加载速度直接翻倍。关于先转成Tensor存起来,我试过存成.pt文件,但如果你每次epoch都要做随机增强,那存Tensor反而没法做在线变换了,除非你接受只用固定预处理,这个得看你的任务需不需要随机裁剪翻转这些。num_workers报内存炸,可能是你每个worker都在复制一份完整的数据索引或者transforms对象,试试把worker数降到2,或者用persistent_workers=True,再不行就检查下是不是dataloader的prefetch_factor默认值太大。另外你提到tfrecord,PyTorch这边可以用webdataset或者直接写个自定义IterableDataset从LMDB读,但我觉得最省事的办法还是把图片压成png或webp格式,体积小解码快,配合内存映射读取。最后提一句,确认下你的transforms里有没有用RandomResizedCrop这种特别耗时的操作,如果只是分类任务,可以先中心裁剪再缩放,把随机性降到最低。
我之前也踩过这个坑,10万张图全走jpg解码确实顶不住。建议先用脚本把所有图片预处理成png或直接存成.npy/tensor,训练时只做轻量变换,能快好几倍。另外num_workers报错不一定是内存不够,可能跟你transforms里的随机操作有关,试试把worker的prefetch_factor调小,或者用persistent_workers=True看看。还有个思路是绕开torchvision,用albumentations做在线增强,它底层用opencv读图,速度比PIL快不少,而且支持多进程更稳。
另外你提到tfrecord,PyTorch这边可以用webdataset或者lmdb,把图片打包成tar或者key-value存储,IO压力会小很多。老实说,如果模型不大,甚至可以先把所有图跑一遍存成tensor再训练,虽然占磁盘,但省心。你现在的瓶颈大概率是磁盘随机读,试试把图片按类别分文件夹,或者用SSD会不会好点?
我之前也踩过这个坑,10万图全走transforms确实扛不住。建议先把resize和归一化做成离线预处理,存成npy或者lmdb,训练时Dataset里只做ToTensor,速度能快好几倍。另外num_workers报错大概率是每个worker都复制了一份完整数据索引,试试把shuffle和drop_last关掉,或者把worker数降到2,配合persistent_workers=True能缓解内存压力。至于tfrecord,PyTorch这边可以用webdataset或者直接上ffcv,格式类似但更灵活,不过前期转换也要花时间,看你取舍了。
试试把图片预处理后缓存成lmdb或h5py,读取快很多,内存爆了就把num_workers降到2加persistent_workers=True。
我试过类似的情况,10万张图其实不算特别多,但瓶颈往往不在Dataset本身,而在你的预处理链路。transforms里的resize和归一化如果是CPU上跑的,每个worker都在重复计算,4个worker内存炸很正常,建议先看看是不是图片解码那步太吃内存,可以试试把decode和resize放到worker里,但主进程里别加载原始图。至于存Tensor,我自己的经验是如果硬盘够快(比如SSD),不如直接存成压缩的.pt文件,每个样本一个文件或者打包成多个shard,读取时用mmap模式,这样能省掉JPEG解码的时间,但注意别把整个数据集一次性load进内存。另一个坑是transforms里的随机操作,像随机裁剪、翻转这些,虽然每个样本只做一次,但10万张累积起来开销很大,你可以先离线把图resize到固定尺寸,训练时只做轻量级的归一化。PyTorch没有直接对应tfrecord的东西,但可以用webdataset或者datapipes,它们支持流式读取tar包,对大数据集友好很多。最后,如果内存还是爆,试下把num_workers降到2,同时把persistent_workers和pin_memory打开,有时候问题不在worker数量,而在数据加载的队列长度。
先试试把图片预处理成lmdb或内存映射格式,10万张图全放内存也就几个G,比每次读盘快多了。
我之前也踩过这个坑,十万张图全走transforms确实扛不住。建议你把resize和归一化这些固定操作提前做掉,直接存成numpy或者uint8的tensor,训练时只做ToTensor和轻量增强,能快不少。另外num_workers报错大概率是windows下没加if name == 'main'的保护,或者worker数超过CPU线程了,先降到2试试,用persistent_workers=True也能减少反复加载的开销。至于tfrecord,PyTorch这边可以用webdataset或者lmdb,但对你来说预处理存盘应该就够解决了。
这问题我踩过坑,建议先存成png或jpg再load,别直接存tensor,内存容易爆。另外num_workers设成2试试,4个可能真不够你机器吃的。
10万张图确实得想办法优化,你这情况我太熟了,之前做医学影像也是这么被折磨过来的。先别急着全转Tensor,内存根本扛不住,我试过把图片预处理成uint8的numpy存成.npy,加载时按需转float再归一化,速度能快不少,而且比直接存Tensor省一半内存。num_workers报错大概率是你机器物理内存不够,四个worker每个都会复制一份数据集索引和预处理管线,试试把persistent_workers=True加上,或者把batch_size调小点,同时worker数降到2看看。transforms里的随机操作确实吃CPU,如果只是resize和归一化,强烈建议先把resize做完存成固定尺寸的图,训练时只做归一化和随机增强,这样能省一大截时间。另外你提到的tfrecord,PyTorch这边可以用webdataset或者lmdb格式,把图片打包成tar或者key-value数据库,IO效率会高很多,尤其是小图场景,10万张图跑起来会流畅不少。还有个土办法,如果你有SSD,把数据先拷到本地再训练,网络盘或者机械硬盘的随机读取慢得离谱,这往往是最大的瓶颈。你现在的GPU利用率大概多少?如果很低的话,问题可能不只是Dataset,预处理那部分才是真凶。
数据读取瓶颈大概率在IO和decoding,试试把预处理后的图存成lmdb或h5py,随机操作留到训练时做。
或者用webdataset直接读tar包,10万图小case,比手搓Dataset省心多了。
我之前也踩过这个坑,10万张图全走transforms确实很伤。建议你试试把预处理后的图直接缓存成.pt或者npy文件,训练时只做ToTensor,能快好几倍。另外num_workers报内存炸,大概率是pin_memory和workers乘起来吃满RAM了,可以先开2个worker加prefetch_factor=2试试。至于tfrecord,PyTorch这边可以用webdataset或者自己写个LMDB,不过对小项目来说,缓存预处理结果是最省事的。
你这情况我太熟了,10万图不算多但预处理全堆在Dataset里确实容易拖死。别急着上tfrecord,PyTorch这边有更轻的方案,先把图片全量转成uint8的numpy或者直接存成.pt的tensor,加载时只做ToTensor和归一化,能省一大截时间。transforms里的随机操作其实还好,主要瓶颈是磁盘IO和JPEG解码,如果你机器内存够大,干脆把整个数据集读进内存,用lmdb或者h5py打包,训练时随机索引,速度会质变。num_workers报错大概率是每个worker都复制了一份完整的数据索引或者transform里用了多进程不安全的操作,试试把worker数降到2,同时把persistent_workers=True加上,再给DataLoader配个pin_memory,有时候能解决。另外你检查下是不是用了全局的transforms对象,每个worker都会共享它,建议在Dataset的__getitem__里动态创建,避免锁竞争。实在不行可以考虑albumentations,它底层用OpenCV,解码和变换比torchvision快不少。最后说一句,如果你有SSD且图片不太大,其实直接把图片路径列表存起来,每个epoch随机打乱,配合num_workers=4(注意别超过CPU核数)通常就够用了,你那个两小时可能还有别的问题,比如每次都在做同步读取。
我之前也踩过这个坑,10万张图全走transforms确实扛不住。你可以试试把resize和归一化提前离线做掉,存成.pt或者npy格式,训练时直接load tensor,能快好几倍。另外num_workers报错大概率是shared memory不够,调小点或者设成0先跑通,再逐个往上加。还有个思路是看看是不是在transforms里用了随机裁剪这些,能省就省,或者用albumentations,它底层优化好一些。
我之前也踩过这个坑,十万张图如果全走jpg解码加transform,瓶颈基本在CPU和磁盘IO上。建议你先试试把图片预处理成png或者直接存成npy/tensor格式,加载时省掉解码这步能快不少。另外num_workers不是越大越好,得看内存带宽,我一般先试2,然后慢慢往上加,同时把persistent_workers=True打开。transforms里的随机操作确实费时间,但如果是分类任务,可以试试把resize和归一化放到dataset里提前做,随机增强用GPU上的库比如kornia来做,能减轻CPU负担。至于tfrecord那种方案,PyTorch有webdataset或者shelve,但感觉你这规模先优化IO就够了。
我直接全转成lmdb存的,读取快十倍,worker数设成2就够了,内存也不会爆。
可以先试试把图片预处理成lmdb或h5py格式存起来,加载会快很多,worker数调小点别硬上。
可以先把图片预处理成lmdb或h5py缓存,读取时只做轻量变换,速度能快好几倍。
先试试把图片预处理完存成npy或lmdb,读取时直接load tensor,能省一大截时间。
我之前也踩过这个坑,10万图用jpg硬读确实要命。建议先把所有图片预处理成png或者直接存成npy/tensor格式,加载时用torch.load或者lmdb,能快好几倍。另外num_workers报错可能是你的内存不够,试试把persistent_workers=True加上,或者调小batch_size和worker数量到2看看。transforms里的随机操作其实影响不大,瓶颈主要在磁盘IO,所以缓存到内存里是王道。