最近在跑一个简单的ResNet50迁移学习项目,数据集大概2万张图片,每张resize到224x224。我设batch size=32,结果跑第一个epoch就报CUDA out of memory。我用的是RTX 3060 12G,按理说应该够用吧?是不是我DataLoader里开了太多num_workers?还是说模型里用了什么隐藏的显存泄露?尝试把batch size降到8倒是能跑了,但训练慢得离谱,而且准确率也不太行。
想请教下各位大佬,这种场景下一般怎么优化显存?混合精度、梯度累积这些方法真的能立竿见影吗?还是说我哪里写得不规范?求指点,孩子快被显存整自闭了。
用PyTorch做图像分类训练时显存炸了,是我代码写错还是batch size太大?
全部回复
共 158 条12G跑ResNet50+224分辨率,batch32确实有点悬,但也不至于秒炸。你先把num_workers降到4试试,这玩意儿经常莫名吃显存。混合精度绝对值得开,显存直接砍半,3060有tensor core不用白不用。另外检查下是不是忘了开eval模式或者没detach梯度,有时候小细节比大参数更致命。
12G跑224的ResNet50,batch32确实紧,先开amp混合精度试试,显存能省一半。
梯度累积只是曲线救国,本质没降峰值,重点还是查下dataloader有没有把缓存堆显存里。
混合精度加梯度累积基本能解决,ResNet50吃显存但12G不该这么惨,检查下pin_memory别开太大。
12G跑ResNet50加224分辨率,batch32按理说真不该炸,先检查下是不是把验证集也塞进DataLoader了,或者模型没切eval模式导致梯度缓存。混合精度我试过,3060上能省一半多显存,但记得关掉AMP的grad scaler报错自动回退,不然偶尔会精度抖动。梯度累积倒是不省显存,只是变相调小batch,你不如直接batch16加AMP跑,速度和精度都能平衡。另外num_workers一般不影响显存,除非你开了pin_memory还疯狂预取,那个会占锁页内存,但跟CUDA OOM关系不大。
12G跑ResNet50加224输入,batch32确实卡在临界点上,这代卡显存带宽也一般。不过你降到8准确率就不行,更像学习率没跟着batch size调,线性缩放原则了解一下。混合精度值得一试,3060的Ampere架构支持得挺好,显存直接砍半,梯度累积倒是能救急但会拖慢训练节奏。另外检查下pin_memory和transform里是不是有重复加载的bug,有时候OOM是预处理线程吃爆内存导致的。
12G跑ResNet50加224图,batch32确实紧,降到16加AMP基本就稳了,梯度累积也能救急。
之前3060跑类似任务,开AMP后显存直接砍半,你试试cast到fp16,比调workers管用。
12G跑ResNet50加224分辨率,batch32按理说真不该爆,你先查下是不是把梯度也存了或者优化器里带了动量缓存,有时候显存是这么悄悄吃满的。混合精度确实立竿见影,开AMP后显存能省近一半,而且3060的Tensor Core不用白不用;梯度累积我建议先别上,它只解决batch大小问题,不治显存病根。另外你这准确率不行八成是lr没跟着调,batch从32降到8,学习率也得相应降,不然收敛肯定受影响。
12G跑224的ResNet50加32的batch确实会紧,我3070ti也是12G,之前一样爆过。你先把num_workers降到4试试,这个影响不大但能省点显存,另外检查下是不是开了gradient checkpointing,那个反而更吃显存。混合精度建议直接上AMP,基本能省一半显存,训练速度还快,准确率一般不会掉。梯度累积适合你这种降batch后想保持等效batch size的情况,但真不是首要优化手段,先把AMP开了再说。还有个小技巧,把图片用albumentations做预处理,别在DataLoader里堆太多transform,也能省点。
12G跑ResNet50加224分辨率,batch32按理说真不该爆,你先排查下是不是在验证阶段也把梯度开着,或者输入没归一化导致显存峰值异常。混合精度确实立竿见影,3060的Ampere架构支持得不错,能省一半左右显存,梯度累积倒是更治本但会拖慢收敛。另外num_workers只影响CPU内存和加载速度,跟显存基本没关系,别在这上面纠结。建议先用torch.cuda.max_memory_allocated定位一下峰值节点,再决定是开AMP还是换224以下分辨率,准确率掉太多的话,可能就是学习率没配合调。
12G跑32的resnet50确实紧张,建议开AMP混合精度,显存直接砍半,还能白嫖一波速度。
梯度累积只解决batch size不够大的问题,对显存占用没帮助,你这情况先开AMP试试。
12G跑ResNet50加224输入,batch32按理说真不至于直接爆,你先把num_workers调成0试试,很多时候是DataLoader预取把显存占满了。混合精度肯定要开的,AMP能省将近一半显存,而且速度还能提一截,我3060跑类似任务都是这么干的。梯度累积其实不太能解决你的问题,它只是模拟大batch,显存占用该多少还是多少。另外你检查下是不是PyTorch版本太老,旧版的内存管理确实有点问题,更新到2.x会好很多。
12G跑ResNet50加224的输入,batch32按理说确实不该爆,你先查下是不是在loss.backward()之前把整批图像都堆在GPU上没释放,比如可视化或者保存梯度那类操作。混合精度对显存改善挺明显的,尤其你这种单卡场景,AMP开起来基本能省一半,梯度累积倒是不省显存,只是等效放大batch。另外准确率不行可能跟batch size关系不大,先确认下预训练权重有没有冻结对层,还有学习率要不要跟着有效batch调。
12G跑ResNet50+224分辨率,batch32按理说真不该爆,你先把num_workers调成0试试,有时候多进程加载反而会复制一份模型显存。混合精度大概率能救你,amp加上之后显存直接砍半,我3070跑类似任务从16涨到32都没压力。梯度累积适合你这种想保batch又怕爆的情况,但准确率掉的话先检查下是不是学习率没跟着调。另外确认下你用的预训练权重是不是真的冻结了BN层,我之前就栽在这上面过。
12G跑ResNet50加224分辨率,batch32按理说真不该爆,你先看看是不是在loss.backward()之前把整批图片都堆在显存里没释放,比如循环里累积了tensor。另外混合精度我个人感觉是立竿见影的,能省将近一半显存,而且3060的Ampere架构支持得挺好,代码改动也就几行。梯度累积适合你这种想保持大batch又显存不够的情况,但注意要同步调整学习率,不然收敛会变慢。还有个容易被忽略的点,验证阶段记得包一下torch.no_grad(),不然反向传播图会一直攒着。
3060 12G跑ResNet50开32确实容易爆,混合精度加梯度累积直接解决,别怀疑代码。
12G跑224的ResNet50开32确实紧,换混合精度加梯度累积试试,保准稳。
3060跑这个配置不该炸,八成是验证集没关梯度或者数据加载堆积了,先查下这两块。
RTX 3060 12G跑ResNet50 batch size 32按理说确实不该炸,我之前用11G的2080Ti跑类似配置都没问题,所以大概率不是硬件上限的问题。你可以先检查一下是不是在训练循环里把loss或者输出累加到了某个list里没释放,这种隐藏的显存泄漏特别常见,跑几个batch就爆了。另外num_workers本身不占显存,它影响的是内存和CPU,别被这个带偏了。混合精度确实立竿见影,AMP一开显存能降三成左右,速度还能快一点,但要注意某些层可能需要手动处理。梯度累积是把大batch拆成小batch再累加梯度,显存降了但训练时间不会省,适合你这种想模拟大batch又跑不动的情况。降到batch size 8能跑但准确率不行,可能是BN层在太小batch下统计量不稳,试试换成GroupNorm或者冻结BN。建议先用torch.cuda.memory_summary()看一下到底哪块占了大头,比瞎猜强。
RTX 3060 12G跑ResNet50 batch size 32按理说不该炸,我之前用11G的2080Ti跑类似配置都没问题,所以大概率不是显存本身不够,而是哪里没配对。先确认一下你是不是没加torch.no_grad()就跑了验证集,或者训练循环里loss没.item()就直接累积计算图了,这种低级错误特别容易让显存悄悄涨上去。num_workers开太多确实会占一部分显存,但一般也就几百兆,不至于直接OOM,你可以先设成4试试。混合精度基本是立竿见影的,amp一开显存能省三成左右,速度还能快一点,配合梯度累积凑等效batch size很香。另外检查下是不是每轮都存了带梯度的中间变量,或者没清cache,torch.cuda.empty_cache()偶尔手动来一下也有用。batch size降到8能跑但准确率不行,很可能是BN层在小batch下统计量不稳,这时候换GroupNorm或者冻住BN会好很多。