最近在跑一个简单的ResNet50迁移学习项目,数据集大概2万张图片,每张resize到224x224。我设batch size=32,结果跑第一个epoch就报CUDA out of memory。我用的是RTX 3060 12G,按理说应该够用吧?是不是我DataLoader里开了太多num_workers?还是说模型里用了什么隐藏的显存泄露?尝试把batch size降到8倒是能跑了,但训练慢得离谱,而且准确率也不太行。
想请教下各位大佬,这种场景下一般怎么优化显存?混合精度、梯度累积这些方法真的能立竿见影吗?还是说我哪里写得不规范?求指点,孩子快被显存整自闭了。
用PyTorch做图像分类训练时显存炸了,是我代码写错还是batch size太大?
全部回复
共 152 条12G跑ResNet50 batch size=32确实有点紧,尤其224x224输入+3通道+Adam优化器,显存峰值很容易超。你可以先检查下DataLoader的pin_memory=True是不是开了,那个在某些情况下会额外吃显存。混合精度(amp)绝对是立竿见影的,能省接近一半显存,而且几乎不影响精度;梯度累积也能救急但会拖慢训练节奏。另外建议确认下是不是用了ImageNet预训练权重,那个本身显存占用就比随机初始化高不少。
12G跑224的ResNet50加batch size 32确实有点紧,但第一个epoch就炸不太正常,我怀疑是DataLoader那边num_workers开太多导致CPU预处理占用了额外显存,或者你用了什么默认的pin_memory=True没注意。混合精度绝对是立竿见影的,torch.cuda.amp包上就能省将近一半显存,而且对精度影响微乎其微,强烈建议你试试。梯度累积也能救急,相当于用时间换空间,比如设accumulation_steps=4就能模拟出batch size 32的效果,但收敛速度会慢一点。不过准确率不行可能跟batch size关系不大,更多是学习率没调对或者迁移学习的冻结策略有问题,你降到8后学习率也得跟着降,不然梯度更新太震荡。另外检查下模型里有没有不小心把BatchNorm设成eval模式或者用了太多dropout层,有时候这些小细节比显存更坑。总之先上混合精度和梯度累积,再调一下num_workers到4以内,应该就能稳在12G上跑了。
12G跑batch size=32确实有点紧了,试试梯度累积加混合精度,效果立竿见影。
batch size 32按说3060能抗住,检查下transform里是不是不小心加了随机resize或者多余cache。
12G跑224的ResNet50,batch size开32确实有点赌,混合精度和梯度累积能救,试试看。
12G显存跑ResNet50加224的图,batch size=32按理说是能跑的,我之前用同样的卡试过类似的设置,大概率不是batch size本身的问题。你检查下DataLoader里num_workers是不是设太高了?一般设4到8就够,开太多反而会占用额外显存来缓存数据,而且有时候还会导致CPU和GPU之间数据搬运的瓶颈。另外,迁移学习时如果你把整个ResNet50的梯度都打开了,光是模型参数加中间激活值就会吃掉不少显存,可以试试冻结前面几层或者用梯度检查点(checkpointing)来省显存。混合精度训练确实立竿见影,尤其是torch.cuda.amp,基本无痛能省一半左右,配合梯度累积还能撑起更大的有效batch size。不过你提到batch size降到8准确率不行,这个有点奇怪,8的batch size对迁移学习来说也不算特别小,可能是学习率没跟着调?建议先把学习率按batch size比例缩放试试,或者用warmup策略。最后检查下代码里有没有不小心把验证集或者额外张量留在GPU上,有时候一个detach()没加就会炸。
12G跑ResNet50加224的图,batch size=32确实有点极限了,尤其是如果输入是3通道浮点,加上中间层的激活和梯度,显存占用很容易就奔着10G+去了。我个人经验是,ResNet50在这种分辨率下,batch size=16左右是比较稳的,你可以先试试16,然后逐步往上加,用torch.cuda.max_memory_allocated()看一下峰值。混合精度(AMP)绝对立竿见影,fp16能把显存砍掉将近一半,而且现在torch.cuda.amp用起来很简单,加个autocast和GradScaler就够了,建议必开。梯度累积也是个好办法,相当于用小batch模拟大batch,但注意要调一下学习率或者warmup,不然收敛会慢。另外,num_workers开多了确实会额外占用一些CPU内存,但对显存影响很小,主要还是模型和batch本身的问题。至于准确率下降,batch=8训练慢可能是因为没有同步BN或者学习率没调好,建议你把lr跟着batch size等比例缩放一下,比如batch=32时lr=0.01,那batch=8时lr可以先降到0.0025试试。最后检查一下你的DataLoader里pin_memory和non_blocking设置,有时候这些细节能省点显存碎片。别自闭,这问题刚入门都会碰到,调参就是个试错的过程。
讲真3060 12G跑ResNet50加224的图,batch size=32按理说确实不应该直接爆,我怀疑问题可能出在DataLoader的num_workers上——开太多的话每个worker都会预加载一批数据进显存,叠加起来就吃紧了,可以试着调成4或者2看看。另外还有个坑是PyTorch的pin_memory,如果你设了True但系统内存不够,反而会跟GPU显存抢资源,可以先关掉试试。混合精度(AMP)是我觉得最立竿见影的方法,显存直接砍半,而且对准确率影响很小,建议你加上torch.cuda.amp试试,代码改动也不大。梯度累积的话,你这种batch size降到8的情况很适合,设个accumulation_steps=4就能模拟出32的效果,训练速度虽然慢点但比单卡显存强。至于准确率不行的问题,batch size小的话建议把学习率也相应调低,不然梯度震荡太厉害。说实话你这配置跑迁移学习完全够用,就是初始化设置没优化好,稍微调一下就能流畅跑起来的。
12G跑ResNet50加224输入,32的batch确实有点紧,换混合精度基本能解决,别太纠结代码问题。
12G跑ResNet50加224的输入,batch32按理说真不该爆,先查下是不是pin_memory开太多或者验证集里混了没resize的图。不过你这降到8准确率就不行有点怪,建议先排除学习率没跟着batch调的问题。混合精度和梯度累积确实能省不少,但梯度累积本质是模拟大batch,你batch8累积四步效果应该和32差不多才对。还有个小技巧,把transform里的Normalize换成在GPU上做,也能省点显存。
ResNet50跑224输入,12G显存batch32按理说绰绰有余,你这情况八成是代码里哪块没写对,比如DataLoader的drop_last没设或者模型里不小心加了别的层。混合精度是最直接的解决方案,amp一开显存直接砍半,3060的TensorCore也能吃满。梯度累积我建议你配上,batch8累积4步效果等同32,但注意BN层会有偏差,最好用syncBN或者调低累积步数。另外准确率不行大概率是学习率没跟着batch size缩放,你把lr按比例调一下试试。
12G跑ResNet50加224输入,batch32按理是不会爆的,你查下是不是把ImageNet的BN统计量也一起fine-tune了,那个很吃显存。混合精度确实立竿见影,直接开AMP能省将近一半,梯度累积对显存没用,只是变相减小batch。另外num_workers跟显存没关系,但建议设成4或者8,别开太多。我怀疑你是不是不小心把验证集的梯度也开了,或者模型里加了什么额外的分支,先排除这些再降batch吧。
12G跑ResNet50加224分辨率,batch32按理说确实不该爆,你检查下是不是把梯度也存到显存里了,或者优化器状态没清。混合精度配合torch.cuda.amp值得试,12G能轻松翻倍batch,梯度累积倒是不怎么省显存只是变相增大batch。另外num_workers影响的是内存不是显存,别在这上面纠结。准确率不高可能跟学习率没调有关,batch变小了lr也得跟着降,不然收敛会不稳定。
12G跑ResNet50加224的输入,batch32按理说真不该爆,我怀疑你transform里是不是带了normalize的GPU操作,或者模型没切train模式导致BN统计量累积。不过说实话,3060的12G带宽跑ResNet50确实有点勉强,我自己的经验是混合精度加梯度累积能直接省一半还多的显存,而且准确率几乎不掉,你试试AMP包里的GradScaler,代码改动很小。至于num_workers,其实它影响的是CPU预处理和内存,跟显存关系不大,你开8个也没问题,别超过CPU核数就行。另外你说batch8准确率不行,这个大概率不是batch大小的问题,可能是学习率没跟着调,或者你迁移学习时把BN层也冻结了,很多新手会踩这个坑。建议你先用torch.cuda.max_memory_allocated()打印一下峰值显存,看看是不是有隐藏的中间变量没释放,我之前遇到过forward里不小心保留了大tensor的引用,导致显存只增不减。最后想问下,你用的是ImageFolder还是自定义Dataset?有时候是数据加载时pin_memory=True配num_workers高反而容易爆,试试把pin_memory关掉,说不定有惊喜。
12G跑224的ResNet50按理说batch32不该爆,你查下是不是在backward之前保留了整批的中间激活值,或者优化器里加了权重衰减却没开foreach。混合精度我觉得最直接,amp一开显存能省一半还多,训练速度也能提上来;梯度累积适合你这种显存紧但batch又不想降的情况,等效batch大一点对收敛也有好处。另外num_workers跟显存没啥关系,准确率不行大概率是学习率没跟着batch size调,降batch后lr也得相应降。
12G跑ResNet50加224的输入,batch32确实有点极限,但也不至于炸得这么狠,你先把num_workers降到4或者干脆设0试试,有时候数据加载线程太多会占额外显存。混合精度绝对是首选,开AMP之后显存基本能砍一半,batch32大概率能跑起来,梯度累积倒是没必要,除非你想把batch再撑大。另外检查下有没有把模型和输入都显式.cuda(),有时候不经意间把验证集的loss或者中间变量留在图上也会导致显存不释放。准确率掉跟batch size关系不大,你降到8可能只是学习率没跟着调,先试试AMP加原batch,应该能救回来。
3060 12G跑ResNet50加224的图,batch32确实紧,但降到8也不该那么慢,检查下是不是在CPU端做预处理卡住了。
混合精度直接开AMP能省一半显存,梯度累积先弄清原理再用,别一上来就堆参数。
迁移学习的话冻结前面几层试试,显存能省不少,你这种情况大概率不是泄露。
12G跑ResNet50加224的输入,batch32按理说真不该爆,我之前用2080Ti 11G跑类似的配置还能剩不少余量。你重点检查下DataLoader的pin_memory是不是设成True了,这个配合num_workers高的时候会额外吃显存,还有确认下模型是不是真的在GPU上,别把数据也堆进显存里。混合精度确实是立竿见影的,AMP一开显存直接砍半,你试试看能不能回到batch32,而且现在torch的amp很成熟,基本不影响精度。梯度累积的话能解决显存问题但会拖慢训练节奏,建议你先把amp开了,再把num_workers调到4试试,另外确认下你的ResNet50是不是从torchvision直接加载的预训练权重,有些自定义实现会把BN层搞得很吃显存。还有个坑,如果你用了albumentations做数据增强,有些操作是在GPU上跑的,那显存也会莫名其妙涨。实在不行就把batch设到16配amp,速度应该比你batch8快不少,准确率的问题大概率是学习率没跟着batch size调,你降batch后记得把lr也按比例降下来。
12G跑ResNet50加224分辨率,batch32按理说真不该爆,你先确认下是不是在训练前把验证集的梯度也算了,或者模型没切eval模式,这两个坑比num_workers坑多了。num_workers只是占CPU内存,显存基本不沾边,除非你pin_memory=True加上数据加载太慢导致GPU等待时缓存堆积,但一般不至于直接OOM。混合精度确实立竿见影,显存直接砍一半还多,而且3060的Ampere架构对fp16支持很好,准确率基本不掉,你装个torch.cuda.amp包几行代码就搞定,强烈建议先试这个。梯度累积我也用过,效果跟调小batch差不多,但注意BN层会受影响,你迁移学习的话可以冻结BN或者用同步BN,不然累积梯度会让BN统计量飘。还有个野路子,把图片resize从224降到192或者176,ResNet50输入尺寸不敏感,精度损失很小,显存能省30%。你降到8能跑但准确率不行,我猜是学习率没跟着batch size调,线性缩放原则,batch从32降到8,学习率也得对应降,不然收敛不稳。最后检查下代码里有没有把整个验证集一次性塞进GPU做评估,很多新手会犯这个,改成循环评估或者用no_grad加batch预测。实在不行就上A100租个云GPU,一小时几块钱,别折磨自己了。
12G跑ResNet50加224的输入,batch32按理说真不该炸,我怀疑你迁移学习时是不是把主干和全连接层的requires_grad都开着,然后反向传播把整个图的梯度都攒下来了?建议先冻结backbone只训分类头试试,显存能掉一大截。另外num_workers跟显存没关系,那个只影响CPU数据加载,别甩锅给它。混合精度是真的立竿见影,amp跑起来显存直接砍半,你3060有Tensor Core不用白不用,代码也就加两三行。梯度累积我平时只在batch实在调不小的时候用,它本质是拿时间换空间,但你要是batch8都卡,那还是先查查是不是有个地方不小心把输入stack成list了,或者模型里有个多余的dropout层在推理时也占显存。还有个小技巧,你用torch.no_grad()跑一下验证集,如果显存还是涨,那八成是DataLoader里pin_memory=True配合num_workers>0在搞鬼,改成False试试。最后说句,准确率不行跟batch size关系真不大,你不如看看学习率是不是没跟着batch调,迁移学习一般用1e-4往下的lr,别拿默认的0.01硬跑。