最近在跑一个文本分类任务,模型用的是BERT-base,数据大概几十万条。单卡A100(40G),batch size调到16就OOM了,试了梯度累积但训练速度慢得离谱。看网上说DeepSpeed的ZeRO能省显存,但配置起来有点复杂,而且我现在用的是PyTorch原生训练循环,不知道值不值得花时间迁过去。另外也试过torch.cuda.amp混合精度,效果有限。想问下各位大佬,这种情况是直接砍batch size硬扛,还是上DeepSpeed?或者有没有其他更简单的显存优化技巧?顺便问下,ZeRO-2和ZeRO-3在实际使用中差别大吗?谢谢!