最近在做一个小项目,想用ResNet50在自定义数据集上做迁移学习,训练图像分类模型。结果每次跑到第二个epoch,显存就飙到20G左右(我只有24G),然后直接OOM挂掉。数据集每张图是224x224,batchsize已经降到8了,还用了混合精度训练(torch.cuda.amp),也试了梯度累积,但感觉治标不治本。
我怀疑是不是模型本身太大,或者我的数据加载方式有问题?但看网上很多人同样配置都能跑,是不是我哪里写的不对?
想请教一下大家:除了换更小的模型,还有没有其他能稳定跑完训练的方法?或者有没有什么检查工具,能帮我定位显存占用到底在哪个层?谢谢各位大佬了。
用PyTorch跑ResNet50做迁移学习,显存总爆掉,求优化思路
全部回复
共 170 条我之前也踩过这个坑,ResNet50的BN层在迁移学习时特别吃显存,你试试把batchsize再砍一半,然后用GradScaler配合梯度累积,效果立竿见影。另外检查下是不是开了梯度检查点,torch.utils.checkpoint能省不少显存,就是会慢一点。显存定位的话,可以用pytorch_memlab或者torch.profiler看每层占用,我上次就是靠这个发现是激活值爆了,不是模型本身问题。
我之前也踩过这个坑,ResNet50的BN层在迁移学习里特别吃显存,尤其是你开了混合精度后,如果没把BN层也转成fp16,它还是会用fp32跑,占的地方比想象中多。建议你检查下是不是所有层都用了amp的autocast,或者试试把BN层冻结掉,只训练最后的全连接层,显存能降一大截。另外可以用torch.cuda.memory_summary()看下分配详情,或者装个pytorch_memlab,能直接定位到具体哪一行代码申请的显存。我之前还遇到过数据加载时worker数设太大,CPU内存爆了反而拖累GPU的情况,你顺便看下DataLoader的num_workers是不是调太高了。
我之前也遇到过类似的情况,后来发现瓶颈不一定在模型本身,很可能是数据加载那边没处理好,比如num_workers设太低导致CPU来不及喂数据,GPU干等显存反而被中间变量占满。
你可以试试用torch.utils.checkpoint来对ResNet的残差块做梯度检查点,用计算换显存,效果立竿见影,基本能省一半以上。
另外排查工具的话,直接上pytorch的profiler或者简单点用nvidia-smi盯每个step的峰值,先确认是不是优化器状态和激活值占了大头,我猜你八成是卡在激活值上了。
我之前也遇到过类似情况,排查下来发现是数据加载那步的pin_memory和num_workers没配好,导致CPU来不及喂数据,GPU空转但显存峰值反而更高,你可以先用nvidia-smi看下是不是数据加载瓶颈。另外建议开一下torch.profiler,能直接看到每个op的显存占用,我上次定位到是BatchNorm的running stats在反向时额外吃显存。还有个偏方,把resnet50的stem和layer1先用冻结的,只训练后面几层,显存能降一半,等收敛差不多了再解冻微调。
我之前也遇到过一模一样的情况,排查下来发现是数据加载那步的pin_memory和num_workers开太大,反而把显存挤爆了,你试试把这两个调低或者直接关掉。另外建议用torchsummary或者pytorch内存分析工具看看每层激活值,ResNet50的block4输出特别占显存,可以在forward里临时加个hook观察下。还有个偏方,把图片尺寸先缩到192跑通流程,确认没问题再改回224,这样能快速定位是不是模型结构的问题。
试试关掉优化器里的momentum缓存,或者用torch.utils.checkpoint换显存,能省不少。
先用nvidia-smi盯着看,八成是数据加载的worker线程没设对,num_workers调成0试试。
试试把batchsize再砍到4甚至2,配合gradient checkpointing,ResNet50没那么吃显存。
用torch.cuda.memory_summary()瞄一眼,八成是backward时激活值爆炸了。
试试把batchsize再砍到4,配合gradient checkpointing,省下的显存够你跑完整个训练。
我之前也踩过这个坑,ResNet50在224输入下其实不算大,但如果你用了ImageNet预训练权重,BN层的running mean和running stat在迁移学习时也会占额外显存,尤其是梯度累积时,batchsize小但累积步数多,等于每个step都保留了多份中间激活。你可以先试试把torch.no_grad()包在验证阶段,同时检查一下是不是DataLoader的num_workers开太多,导致内存换页频繁,显存碎片化严重,这个在Windows上特别明显。另外,我强烈建议你用torch.cuda.memory_summary()看具体分配,或者用nvidia-smi dmon实时监控,大概率问题出在模型输出的logits后面接的loss计算上——如果你用了label smoothing或者自定义loss,中间变量没释放也会堆积。还有个偏门技巧:把模型里所有BN层换成GroupNorm,虽然会掉一点精度,但显存能省30%左右,而且对迁移学习反而更稳。最后,你试过用torch.utils.checkpoint吗?把resnet的layer3和layer4包进checkpoint,能用计算换显存,配合amp基本能压到12G以内。如果还不行,可以看看是不是优化器里用了LARS或LAMB,这两个的动量和方差缓存也很吃显存,换回SGD加warmup试试。
检查下是不是验证集也开了梯度,或者DataLoader的num_workers太高,把pin_memory关掉试试。
试试用torchsummary看每层显存占用,八成是BN层在作怪,把syncBN关了能省不少。
我之前也踩过这个坑,ResNet50加AMP其实挺吃显存的,尤其是BatchNorm层在反向传播时会存中间变量。你可以试试把torch.cuda.amp的GradScaler关掉,有时候混合精度反而让显存碎片化更严重。另外用torch.profiler跑一下,看每个层到底分配了多少显存,我之前发现是最后全连接层的梯度占了大头,换成全局池化加小FC会好很多。还有一个笨办法,把输入尺寸临时降到160x160试试,如果显存占用断崖式下降,那就是输入分辨率的问题,224对显存敏感的话可以配合随机裁剪做数据增强来弥补。
我之前也遇到过一模一样的情况,24G卡跑ResNet50迁移学习,batchsize压到4才勉强不炸。你可以先试试把输入图片尺寸临时缩到192x192跑一个epoch,如果显存明显降下来,那就说明数据加载那边可能没做resize或者归一化,导致实际进网络的tensor比预期大。另外torch.cuda.amp其实对显存帮助有限,真正吃显存的是激活值,你可以在每个block后手动插hook打印激活的shape,或者直接用torch.profiler看内存分配,很容易定位到是哪一层爆的。还有个土办法,把dataloader的num_workers调成0,有时候多进程预取会额外占显存。
试试把ResNet50的BN层冻结,用torchmetrics看下显存分布,大概率是backward时梯度峰值爆的。
我之前也踩过这个坑,ResNet50的BN层在迁移学习时特别吃显存,尤其如果冻结了前面层但没关grad,反向传播还是会算。建议先用torchsummary或pytorch的memory_profiler看看每层占用,大概率是最后一个全连接层或输入张量问题。另外你试试把图像预处理放到GPU上做,或者用torch.utils.data.DataLoader的pin_memory=True,有时候数据加载慢导致显存碎片化也会爆。最直接的办法是把batchsize再砍到4,然后配合梯度累积到等效batchsize 32,我这么跑过7B模型都没炸。
说实话,看到你说batchsize降到8还爆显存,我第一反应是怀疑你哪里没配置对,因为ResNet50在224分辨率下,纯训练模式(不带BN的momentum更新和dropout那些)单卡24G跑batchsize 32都绰绰有余,除非你同时开了验证集的forward或者把梯度也存了。你可以先用torch.cuda.max_memory_allocated()去打印一下峰值分配,再用torch.cuda.memory_snapshot()做一下堆栈分析,基本能看出是不是某个中间变量没释放。我猜大概率是你数据加载那一侧出了问题,比如你用了DataLoader的num_workers但没做pin_memory,或者你在每个step里手动调了model.train()和model.eval()导致BN的running stats反复计算,但更常见的是你无意中把整个验证集也塞进了显存——很多人会在每个epoch结束后直接跑验证集,如果验证集没设batchsize,就会一次性加载全部数据。另外你提到梯度累积,但如果是累加后仍OOM,说明你其实没找到峰值点,因为梯度累积只是降低每步的batchsize,但显存峰值还是由单次forward+backward决定的。还有个很隐蔽的坑:如果你用了自定义的Dataset并且返回了额外的meta信息(比如文件名、box坐标),这些Python对象不会被自动释放,会一直占着CPU内存,但不会直接爆显存,不过会拖慢速度让你误判。建议你先跑一个纯dummy输入(torch.randn(8,3,224,224))的完整训练循环,如果这个都爆,那问题在模型结构或优化器状态,如果这个不爆,那问题就在你的数据管线里。最后实在不行,可以试试torch.utils.checkpoint把ResNet的block包一层,用计算换显存,代价是训练慢个30%,但能稳定跑完。
跑20G有点离谱了,224输入ResNet50按理说显存占用没那么夸张。你可以先试试把batchsize降到4,然后开gradient checkpointing(torch.utils.checkpoint),这招能省不少激活显存,代价是慢一点但稳。另外检查下是不是数据加载时pin_memory=True加num_workers太多,或者模型里无意间保留了梯度,用nvidia-smi配合torch.cuda.max_memory_allocated()看看峰值到底在哪层,别光看总占用。
我之前碰到过类似情况,最后发现是dataloader里transform做了太多冗余操作,比如每次迭代都重新读原图再resize,缓存住预处理好的tensor能快很多。你还可以试试把ResNet50的BN层全冻结,只训练最后的全连接层,显存直接掉一半,效果也不差。如果还不行,考虑用torch.profiler跑一遍,它能精确到每个op的显存分配,很快就能定位是backbone还是loss那块在吃内存。
我之前也踩过这个坑,ResNet50的BN层在迁移学习时如果冻结了前面层,反向传播的显存反而更吃紧,建议你先用torchsummary或pytorch_memlab看下具体哪层占得多。另外224x224+bs8其实不算大,你检查下是不是dataloader的num_workers开太多,或者pin_memory=True导致CPU和GPU之间拷贝卡住了。还有个野路子:把输入尺寸降到192x192,准确率掉不了多少,但显存能省一大截。
我之前也遇到过类似情况,最后发现是验证集和训练集一起跑的时候,验证阶段的梯度没清干净,导致显存峰值特别高。你试试把验证循环里也包上torch.no_grad(),然后单独记录一下每个阶段的显存占用,用torch.cuda.max_memory_allocated()对比下训练和验证的峰值,基本就能定位了。另外你用的预训练权重如果是ImageNet的,可以考虑把BN层冻结(bn.eval()),能省不少显存,而且迁移学习效果通常不会变差。
试试把验证集也放进amp上下文,或者用torch.profiler看下峰值在哪,我上次是BatchNorm的running stats吃的显存。
我之前也遇到过类似情况,后来发现罪魁祸首是数据加载时每个step都做了随机resize和增强,显存碎片化特别严重,建议你排查下transform里有没有不必要的操作。另外可以试试把输入图像直接resize到192x192,ResNet50对分辨率没那么敏感,显存能省不少。检查显存占用的话,用pytorch的torch.cuda.memory_summary()看每个tensor的分配,或者装个nvidia-ml-py盯实时占用,基本能定位到是激活值还是梯度炸了。