最近在做一个图像分类项目,数据集是自己爬的,大概有10万张图片。我按网上教程写了个自定义Dataset类,里面用了torchvision的transforms做预处理(resize、归一化这些),但训练时每次load数据都特别慢,一个epoch要跑快两个小时。我看别人说用DataLoader的num_workers能加速,但我设了4个worker之后反而报错了,好像是内存炸了。想问下各位大佬,这种自定义数据集的情况,有没有什么常规的优化技巧?比如是不是应该先把图片转成Tensor存起来,还是说transforms里少用点随机操作?或者PyTorch有没有像TF那样可以直接从tfrecord读数据的方式?先谢过大家了。
用PyTorch写自定义Dataset时,数据加载太慢怎么办?
全部回复
共 152 条我之前也踩过这个坑,10万张图全走JPEG解码加transforms确实扛不住。你那个内存炸了很可能是num_workers设太高,加上每个worker都会复制一份数据集索引和transform状态,4个worker如果图片还特别大,直接就把RAM吃满了。建议先试试把num_workers降到2,同时把persistent_workers=True加上,这样能省掉反复启动worker的开销。
至于把图片转成Tensor存起来,我个人觉得分两步走比较稳:第一次先把所有图resize成统一尺寸(比如256x256),用uint8格式存成npy或者LMDB,这样磁盘占用比原始JPEG还小,加载时只需要做ToTensor和归一化,省掉了大部分解码时间。如果内存允许,甚至可以全部load进RAM再用DataLoader的pin_memory,那个速度提升是质变的。
transforms里的随机操作确实会拖慢速度,但如果是训练集,随机裁剪和翻转这种增强还是得保留,不然容易过拟合。你可以把随机操作放到GPU上做,比如用torchvision.transforms.v2,或者干脆在Dataset里只做基础预处理,把增强挪到训练循环里用tensor操作实现,这样能避开CPU瓶颈。
另外PyTorch确实没有像tfrecord那样的官方格式,但可以用WebDataset或者FFCV库,它们能把数据打包成tar或者二进制分片,配合多进程读取效率很高,尤其适合这种大规模自定义数据。不过要是懒得改架构,最直接的办法是先把全部图片预处理成.pt文件,每个文件存个几百张的tensor批量,用torch.load读,速度能快好几倍。你试试看哪个方案更顺手,有问题再交流。
我之前也踩过这个坑,10万张图直接读确实难受。建议先把图片预处理成uint8的numpy或者torch tensor存成.pt文件,加载时省去decode和resize,能快不少。另外num_workers不是越大越好,得看内存带宽,你4个worker爆了可以试试2个,或者把persistent_workers和pin_memory打开。transforms里的随机操作别省,但可以移到加载后在GPU上做,或者用albumentations这类库,比torchvision的快。
我之前也踩过这个坑,10万张图raw读确实顶不住。建议先把图片resize成小尺寸(比如256)再存成jpg或者npy,transforms里只留归一化,随机增强放到训练时用另一个轻量版,能快一大截。另外num_workers报错大概率是共享内存不够,试试把persistent_workers=True加上,或者把batch_size调小点。你要是内存充裕,直接全量load进RAM用lmdb映射也行,比Tensor快不少。
试试把图片预处理后存成lmdb或h5py,读取时直接load tensor,能省一大截解码时间。
先别转Tensor硬存,试试把图片预处理改成多进程cache到lmdb或者h5py,能快不少。
先转成pin_memory再配合num_workers调小点试试,或者干脆把预处理好的图直接存成npy格式,读起来快很多。
我之前也踩过这个坑,10万张图直接读确实要命。建议你先试试把图片预处理完缓存成lmdb或者h5py格式,读取速度能快好几倍,而且内存占用也稳。另外num_workers报错不一定单是内存问题,看看是不是DataLoader的pin_memory和worker配合出了问题,可以先把pin_memory关掉再调worker数。transforms里的随机操作其实影响不大,瓶颈基本都在磁盘IO上,所以优先解决存储格式才是正经事。
先转成内存映射的npy或lmdb再训,另外num_workers别瞎调,得看内存带宽和队列超时。
把图先预处理成npy或lmdb存起来,读取时直接load tensor,10万张能快好几倍。另外num_workers别硬怼,先降到2试试,内存不够就换persistent_workers=True。
我之前也踩过这个坑,10万张图如果每次实时从磁盘读原图再resize,瓶颈基本卡在IO和JPEG解码上,CPU根本忙不过来。你设num_workers报错大概率是默认的prefetch_factor太大,加上每个worker都会复制一份Dataset,内存直接翻好几倍,可以先试试把num_workers降到2,同时把persistent_workers=True加上,能省掉反复创建进程的开销。
至于把图片预处理成Tensor存起来,这个方向是对的,但别直接存成单个大文件,内存和加载都不友好,建议用LMDB或者HDF5,把resize和归一化提前做掉,训练时只做ToTensor和标准化,能快好几倍。transforms里的随机操作(比如随机裁剪、翻转)不会拖慢加载,但如果你用了RandomResizedCrop这种带插值的,确实比纯Resize费时间,可以只在训练集用,验证集直接Resize。
另外你问的tfrecord类似物,PyTorch有WebDataset或者用torchdata的DataPipes,不过上手成本比LMDB高一点,我个人觉得对于图像分类,先把图片压缩成WebP或者JPEG质量调低一点,再配合LMDB就够了,10万张图其实不算多。还有个野路子,如果你显存够大,干脆把整个数据集预处理后一次性塞进内存里(比如用numpy memmap映射),训练时直接切片,速度起飞,但得看你机器配置了。最后建议用torch.profiler跑一下看看具体卡在哪个环节,别盲目调参。
我之前也踩过这个坑,十万张图直接读确实扛不住。建议你先别急着转Tensor,试试把图片预处理成LMDB或者h5py格式,读取速度能快好几倍,内存也稳。另外num_workers报错大概率是pin_memory设成True加上worker数太多导致的,可以先关掉pin_memory,worker设成2看看。transforms里的随机操作其实不太影响速度,瓶颈主要在图IO,真要提速可以考虑把resize和归一化提前到生成数据集时做掉,训练时只做tensor转换。
我之前也踩过这坑,把图片预处理后直接存成.pt文件能快不少,内存炸就调低num_workers试试。
大概率是transforms里随机操作卡CPU,先把图预处理成npy或lmdb存着,训练时只读不改会快很多。
我之前也踩过num_workers的坑,4个不够就试8个,但记得把内存换大点或者用shared memory。
我之前也踩过这个坑,10万张图全放内存里确实容易爆。你试试把transforms里的随机操作去掉,或者用albumentations这种更快的库,能省不少时间。另外,可以先转成LMDB或HDF5格式存起来,读取时用内存映射,比一张张从磁盘读快多了。num_workers报错的话,先降到2看看,或者把prefetch_factor调小一点,内存压力会小很多。
我试过类似的情况,10万张图全走transforms确实扛不住。建议先离线把resize和归一化做完,存成.pt或者npy,训练时只做ToTensor,能快好几倍。另外num_workers报错大概率是shared memory不够,试试把worker数降到2,或者加一下prefetch_factor。还有个土办法,不用transforms里那些随机增强,先跑通基线再说,后面再慢慢加。
我之前也踩过这个坑,10万张图全走transforms确实扛不住。建议先把预处理后的图像直接存成.pt或者npy格式,训练时只做ToTensor,能快一大截。另外num_workers报错大概率是内存溢出,试试把persistent_workers=True加上,或者把batch_size调小点,worker数降到2看看。还有个小技巧,如果图片尺寸统一,别在transforms里做随机crop,改成固定尺寸能省不少CPU开销。
可以先把图片预处理成缓存文件,训练时直接读缓存,能省一大截时间。另外num_workers报错大概率是内存不够,减到2试试。
我之前也踩过这个坑,10万张图直接读确实要命。你可以试试把预处理后的结果用lmdb或h5py缓存成二进制文件,这样加载就是纯IO,快非常多。另外num_workers报错大概率是内存溢出,可以试试把persistent_workers设为True,或者把每个worker的batch_size调小点,别一次性全塞进去。至于transforms里的随机操作,训练时留一两个必要的就行,太多确实拖速度,验证集的话可以全关掉。
我之前也踩过这个坑,10万图全走transforms确实扛不住。建议先把resize和归一化做好,直接存成.pt或者npy,训练时只做ToTensor这种轻量操作,速度能快好几倍。另外num_workers报错大概率是共享内存不够,试试把persistent_workers=True加上,或者把batch_size调小点,别一次全塞进去。至于tfrecord,PyTorch这边可以用WebDataset或者自己写个LMDB,但前期转换也得花时间,不如直接预处理成tensor省事。
我之前也踩过这个坑,10万张图全走jpg解码确实扛不住。建议先跑个脚本把图片统一resize成小尺寸(比如256)再存成png或者npy,训练时直接读预处理后的文件,能快好几倍。另外num_workers设4报错大概率是内存爆了,可以试试把persistent_workers=True加上,或者把batch_size调小点,再不行就换prefetch_factor=2这种参数。transforms里的随机操作尽量挪到GPU上用torchvision的随机裁剪那套,或者干脆少用,别全堆在CPU端。