最近在做一个图像分类项目,数据集是自己爬的,大概有10万张图片。我按网上教程写了个自定义Dataset类,里面用了torchvision的transforms做预处理(resize、归一化这些),但训练时每次load数据都特别慢,一个epoch要跑快两个小时。我看别人说用DataLoader的num_workers能加速,但我设了4个worker之后反而报错了,好像是内存炸了。想问下各位大佬,这种自定义数据集的情况,有没有什么常规的优化技巧?比如是不是应该先把图片转成Tensor存起来,还是说transforms里少用点随机操作?或者PyTorch有没有像TF那样可以直接从tfrecord读数据的方式?先谢过大家了。
楼主
22小时前
用PyTorch写自定义Dataset时,数据加载太慢怎么办?
请 登录 后发表回复
全部回复
共 1 条
2楼
21小时前
10万张图用默认的ImageFolder加transforms确实容易卡,尤其是resize这种操作很吃CPU。你试试把图片预处理提前做好,比如用albumentations或者PIL先统一resize成固定尺寸,然后直接存成numpy或者pickle,DataLoader里只做to tensor和归一化,这样能省不少时间。num_workers报错可能是内存不够,可以降到2或者1试试,同时把pin_memory=True打开,另外检查下你的transform里有没有用RandomCrop这类随机操作,太频繁也会拖慢速度。torch没有tfrecord的等价物,但可以考虑用lmdb或者h5py把数据打包成二进制文件,随机读取会比文件系统快很多。