最近在跑一个图像分类的小项目,用的PyTorch 2.0,数据就是几千张图片,不算大。但发现训练时CPU利用率一直上不去,GPU经常在等数据,epoch时间明显比在Linux上长。我按网上说的把num_workers从0调到4、8甚至12,结果不但没变快,有时候还更卡了,甚至会报BrokenPipeError。查了一下发现好像Windows下DataLoader的worker机制和Linux不一样,是spawn而不是fork?那是不是意味着我的预处理代码(比如albumentations)有内存拷贝开销?还是说应该用别的方式加载数据?有经验的前辈能分享一下Windows下的最佳实践吗?
为什么PyTorch的DataLoader在Windows上这么慢?换了num_workers也没用
全部回复
共 24 条Windows下把预处理挪到GPU或主进程试试,worker设成0反而可能更快。
确实,Windows下spawn机制真的是个坑,每个worker都会重新import整个模块,你那些albumentations的转换定义、全局变量全得重新初始化一遍,内存和启动开销自然大。我之前也遇到过,后来直接把预处理改成在dataset的__getitem__里用cv2加numpy写死,绕开那些重型库,速度快了不少,你可以试试把albumentations换成torchvision自带的transform看看。
另外num_workers不是越大越好,Windows上尤其明显,你试过2或者3吗?我这边发现超过4反而因为进程调度和IPC瓶颈更卡。还有个偏方,就是把数据先全部读进内存做成tensor存成.pt文件,训练时直接load,虽然占内存但能彻底绕开IO问题,几千张图完全够用。
至于BrokenPipeError,基本就是主进程和worker通信断了,试试在训练循环外面加个if name == 'main'保护,然后dataloader的persistent_workers=True,有时候能缓解。说实话Windows上PyTorch性能就是不如Linux,如果只是个人项目,装个WSL2跑会省心很多,数据读取这块能差出一倍时间,你值得考虑下。
Windows下spawn确实会重复导入和拷贝数据,试试把num_workers设成0,用内存映射或提前把图片转成npy,比调worker数管用。
Windows下DataLoader慢基本是spawn的锅,每个worker都要重新import一遍库和你的dataset代码,开销比Linux的fork大得多。你可以试试把num_workers设成2到4就行,再往上加反而会因为进程调度和内存复制更卡。另外persistent_workers=True和pin_memory=True能省掉每个epoch重建worker的时间,对Windows挺管用的。如果还是不行,干脆把图片预处理好存成npy或者lmdb,训练时直接读,绕开decode和augment的瓶颈。