最近在跑一个简单的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分辨率,batch 32理论上真不该炸,我怀疑你代码里有个隐蔽的坑。先检查下是不是在训练循环里把每张图都detach到CPU或者重复存了梯度,还有看下是否用了torch.no_grad()做验证但没切回train模式,这种小细节很容易爆显存。num_workers影响的是CPU内存不是显存,所以大概率不是它的锅。混合精度确实立竿见影,我3060上直接能省一半显存,而且速度还快,你试着加个torch.cuda.amp包装下前向和loss,基本无痛。梯度累积虽然能救急,但会拖慢收敛,准确率下降可能跟这个有关,不如先试试把batch降到16加AMP,应该能跑得动。另外你确认下DataLoader的pin_memory开没开,开了能省点显存碎片,还有模型里如果用了BatchNorm的track_running_stats=False也会额外占显存。最后建议你监控下显存曲线,用nvidia-smi看是不是每个step都持续增长,如果是那就是泄露,得查下有没有在循环里不小心创建了计算图。
12G跑ResNet50加224尺寸,batch32确实紧,先开混合精度试试,能省一半显存。
梯度累积解决不了显存峰值问题,本质还是得降batch或者换小模型,你用8太保守了。
3060跑224的resnet50开32确实勉强,降到16加AMP基本就稳了。梯度累积对显存没帮助,纯调batch就行。
混合精度yyds,12G开16batch加amp能跑,准确率低可能是学习率没跟着调。
12G跑ResNet50加224分辨率,batch32按理说不会直接爆,我觉得你大概率不是显存不够,是碎了一地。PyTorch默认的缓存分配器经常把释放的显存块留着不还给驱动,下次申请更大块就直接OOM,你试试点一下torch.cuda.empty_cache()再跑,或者把dataloader的persistent_workers和prefetch_factor都关掉看看。另外ResNet50的BN层在迁移学习时如果开了track_running_stats,反向传播会额外存每层的中间统计量,这个在12G上虽然不算致命但很占空间,建议改momentum或者直接冻结前几层。混合精度肯定有效,AMP能直接砍掉一半显存,但前提是你得把loss scaling处理好,不然数值波动会很烦。梯度累积我也常用,但你这情况更像是单步峰值太高,先试试gradient_checkpointing,把ResNet的block包一下,用时间换空间,效果比降batch明显多了。至于准确率下降,batch从32掉到8,BN的统计量会变吵,你可以试试把BN换成GroupNorm或者用大点的学习率warmup,别急着怀疑代码写错了。对了,你确认一下是不是在backward之前忘了把optimizer.zero_grad()放对位置,有时候梯度叠加会让显存持续增长,看起来像泄漏,其实只是没清干净。
12G跑224的ResNet50开32必爆,正常,换混合精度加梯度累积,稳得一批。
试试AMP加梯度累积,3060这卡就这么点显存,batch16加混合精度基本能稳。
说实话12G跑ResNet50加224分辨率,batch32按理说没那么容易爆,你这情况我怀疑是pin_memory或者数据预处理那块在搞鬼。num_workers开太多确实会吃额外显存,尤其是每个worker都预加载一批图,建议先调成4试试,还有确认下transform里有没有 inadvertently 把图片转成浮点再归一化,那玩意会多占好几倍内存。
混合精度是真能立竿见影,AMP一开显存直接砍半,而且3060的Tensor Core不吃亏,训练速度还能提一截。梯度累积的话适合你这种小显存场景,但注意别跟BN冲突,最好配合虚拟batch size一起调,不然准确率波动会很大。
另外你说batch8训练慢得离谱,这不太正常,除非你learning rate没跟着调,小batch需要更小心地调lr,不然收敛就是慢。要不你试试batch16加AMP,我估计能稳定跑,速度也比你现在的batch8快不少。
最后提一句,看看有没有把验证集的gradient也存下来了,有时候无意中开了model.eval()但没包torch.no_grad(),那显存一样会被吃光。先排查这些,别急着上太复杂的优化。
12G跑ResNet50加224分辨率,batch32按理说真不该炸,你这情况我第一反应是pin_memory和num_workers的锅,尤其Windows上num_workers开多了反而会吃显存,因为每个worker都会复制一份数据缓存。我之前用3060跑类似任务,batch64都能稳住,后来发现是DataLoader里用了transform的随机增强,每次迭代都会重新申请内存碎片,把persistent_workers=True加上,再把num_workers压到4,瞬间就稳了。混合精度确实立竿见影,你直接torch.cuda.amp包一下forward和loss计算,显存能砍掉将近一半,而且3060的安培架构对fp16支持很好,几乎不掉点。梯度累积就是个折中方案,相当于变相扩大batch,但没法解决你单次前向的峰值占用,建议先排查下是不是模型里忘了设.eval()或者有梯度图没释放。还有个冷门技巧,把图片先预处理成张量存磁盘,训练时直接load tensor而不是走PIL,能省掉一大块临时内存。你降到batch8准确率不行大概率是BN层统计量太抖,跟显存优化是两码事,先解决显存再考虑调学习率吧。
12G跑resnet50加224的batch32确实紧,混合精度加梯度累积基本能解决,num_workers影响不大。
同款卡,我之前直接开AMP显存直接砍半,你可以先试试这个,不行再调batch。
12G跑224的resnet50开32确实悬,先开AMP试试,基本能省一半显存。
12G跑ResNet50加224输入,batch32按理说确实不该爆,你先看下是不是在loss.backward()之前把整个dataset都塞进GPU了,或者模型里有个没关梯度的BN层。混合精度基本是必开的,显存直接砍半,3060对fp16支持也还行,梯度累积对显存没帮助但能让你用更小的batch模拟大batch,先别急着骂代码。另外你降到8准确率不行大概率是学习率没跟着调,batch变了lr也得按比例缩放,不然收敛慢很正常。
12G跑ResNet50+224尺寸,batch32确实有点紧,但也不至于直接爆。你试试把pin_memory关了,或者把transform里的Normalize放到GPU上做,有时候数据加载那块的显存占用比想象中大。混合精度建议直接上,AMP在3060上能省不少显存,而且速度还快,梯度累积倒是能解燃眉之急但训练时间会拉长。另外检查下是不是有个地方不小心把梯度retain_graph了,我上次就是这问题白折腾半天。
12G跑ResNet50加224分辨率,batch32按理说确实不该爆,但你先检查下是不是在训练循环里把每个batch的loss和输出都存下来了,或者验证阶段忘了关梯度计算,这两个坑我踩过好多次。另外num_workers影响的是CPU数据加载,跟显存关系不大,别在这个上面纠结。混合精度(AMP)和梯度累积在这种场景下确实管用,尤其AMP基本是白嫖显存和速度,强烈建议加上,反正代码改动就几行。还有个细节,你迁移学习的话记得把BN层也设成eval模式,有时候冻结了主干但BN还在跑训练模式,也会多吃显存。
12G跑224的ResNet50按理够,先开混合精度,能省一半显存,还不行就梯度累积。
batch降到8太亏了,不如保持32加梯度累积,效果一样还省显存。
12G跑224的ResNet50开32确实勉强,先开AMP试试,能省一半显存。
梯度累积本质不省显存,只是模拟大batch,建议先把num_workers调低排查下。
12G跑ResNet50+224的batch32确实有点紧,但也不是完全没救。你先把num_workers降到4试试,这个影响不大,关键看是不是pin_memory开着导致显存碎片化。混合精度(AMP)是真有用,直接能省一半显存,而且训练速度还能提一截,准确率基本不掉。梯度累积适合想保住大batch效果但显存不够的情况,不过你batch8都慢的话,不如先改AMP+把图片预处理挪到GPU之前,说不定能撑住batch16。另外检查下是不是不小心把验证集的梯度也算了,或者有变量没detach,这种隐藏bug也会爆显存。
12G跑ResNet50加224输入,batch32按理说真不该爆,先看看你是不是在训练里同时保留了梯度图或者验证集没关梯度,我之前就栽在这上面。混合精度值得开,能把显存砍一半,但准确率崩的话检查下loss scaling。梯度累积就是个曲线救国,本质没省显存,但能让等效batch变大,训练慢的问题可以试试加大num_workers或者用pin_memory。另外你降到8准确率不行,很可能是学习率没跟着batch size调,这个比显存更值得注意。
12G跑ResNet50加224的输入,batch32按理说真不至于爆,我怀疑你pinned memory开太多或者验证阶段没关梯度,先检查下有没有写model.eval()。混合精度挺管用的,我换了之后显存直接砍半,速度还快了,就是loss缩放那块得调一下。梯度累积倒是能救急,但准确率波动会大点,建议先用AMP试试,如果还爆就看看是不是num_workers设太高挤占了显存。
12G跑ResNet50加224输入,batch32按理说真不该爆,你可以先排除下是不是开了gradient checkpoint或者模型里偷偷加了什么大tensor。混合精度确实立竿见影,3060的显存带宽优化很明显,梯度累积也能救急但会拉长训练时间。另外num_workers一般不影响显存,真正吃显存的是优化器状态和中间激活值,你可以试试把验证集也放进DataLoader时shuffle关掉,偶尔会遇到缓存没释放的情况。如果降到8才能跑,建议查下是不是PyTorch版本和CUDA版本没对齐,之前遇到过类似诡异问题。
12G跑ResNet50加224的输入,bs32按理说真不至于爆,你先排查下是不是在验证阶段忘了关梯度计算,或者DataLoader的pin_memory开太高了。混合精度绝对是第一优先级,开了之后显存能省一半还多,训练速度反而更快,精度基本无感。梯度累积适合你这种小显存但想保持大batch的场景,不过注意BN层会受影响,最好配合同步BN或者调整累积步数。另外你降到bs8准确率不行可能不是batch的锅,先确认下学习率有没有跟着batch size一起调,不然模型收敛会很挣扎。
12G跑ResNet50加224分辨率,batch32确实有点悬,我3070ti 8G用batch16都勉强,你降到8能跑说明不是代码问题。混合精度记得开一下,显存能省一半还提速,另外把pin_memory关掉试试,num_workers影响不大。梯度累积加上效果也明显,但准确率不行可能得看看学习率要不要跟着调。