最近在做一个小项目,想用ResNet50在自定义数据集上做迁移学习,训练图像分类模型。结果每次跑到第二个epoch,显存就飙到20G左右(我只有24G),然后直接OOM挂掉。数据集每张图是224x224,batchsize已经降到8了,还用了混合精度训练(torch.cuda.amp),也试了梯度累积,但感觉治标不治本。
我怀疑是不是模型本身太大,或者我的数据加载方式有问题?但看网上很多人同样配置都能跑,是不是我哪里写的不对?
想请教一下大家:除了换更小的模型,还有没有其他能稳定跑完训练的方法?或者有没有什么检查工具,能帮我定位显存占用到底在哪个层?谢谢各位大佬了。
用PyTorch跑ResNet50做迁移学习,显存总爆掉,求优化思路
全部回复
共 170 条八成是ResNet50的BN层在作怪,微调时把BN冻住或用同步BN试试,显存能省不少。
先检查下是不是数据加载时偷偷缓存了整图集,用torch profiler看下哪层最吃显存。
我之前也踩过这个坑,后来发现多半不是模型本身的问题,而是数据加载和验证阶段没处理好。你可以先检查下是不是每个step都调用了torch.cuda.empty_cache(),或者验证集里也开了梯度,把torch.no_grad()加上试试。另外建议用nvidia-smi监控一下,看峰值是不是出现在forward还是backward,如果前向就爆,那可能是图像预处理时把数据都堆到GPU上了,试试把transforms里的归一化放到CPU端做。还有个小技巧,用torch.utils.checkpoint对ResNet的layer3和layer4做梯度检查点,能省不少显存,代价就是慢一点,但至少能跑完。
我之前也踩过这个坑,ResNet50加AMP其实挺吃显存的,问题不一定在模型本身,你试试把pin_memory关掉,或者num_workers调成0,数据加载那部分经常会悄悄占显存。另外别光看总显存,用nvidia-smi盯着每个进程,或者装个pytorch_memlab,能按行打印出每层的显存分配,我之前就是靠它发现是BatchNorm的缓存没清掉。还有个小技巧,把优化器改成AdamW加fused,配合amp能省一点,如果还不行就检查下输入图像是不是意外转成了FP32,有时候数据预处理那儿会漏掉。
我之前也踩过这个坑,ResNet50跑224输入其实不算大,问题多半不在模型本身。你先用nvidia-smi盯着看,是不是数据加载那步把CPU内存拷贝到显存时爆的,或者pin_memory开太狠了。另外检查一下是不是每步都调用了backward之后没清梯度,或者BN层在迁移学习时跑成train模式导致激活值缓存太大,试试冻结前几层只用BN的running stats。实在不行用torch.utils.bottleneck或者PyTorch自带的profiler,能直接看到每个op的显存分配,我当时就是这么发现是某个卷积层的输入缓存没释放。
我之前也踩过这个坑,ResNet50在224输入下其实不算特别吃显存,20G这个数字不太正常。你试试把torch.no_grad()包在验证阶段,很多人是训练验证一起算梯度才爆的。另外检查一下是不是开了drop_last=False,最后一批batchsize太小反而容易让显存碎片化,我遇到过类似情况。
关于定位工具,torch.cuda.memory_summary()能看到每个张量的分配,但更推荐用pytorch的CUDA caching allocator,直接打印memory history。我之前发现是backward时checkpoint的中间激活值在作祟,ResNet50的bottleneck层如果开了gradient checkpointing,能省一半以上显存。
其实你降到batchsize=4再加8步梯度累积,效果和batchsize=32差不多,但单步显存压力会小很多。还有个偏方,把图片先做随机裁剪到192x192再resize回224,相当于数据增强还能省点显存,很多人没意识到输入尺寸对激活值的影响是平方级的。
最后建议你用nvidia-smi监控一下是不是其他进程占了显存,我试过跑别的代码没清干净,单独看PyTorch占用其实才12G。如果还不行,就把BN层换成GroupNorm,batchsize小的时候BN统计量不稳定也会拖累显存。
第二个epoch才爆显存,大概率不是模型本身的问题,更像是验证阶段没加torch.no_grad(),或者训练循环里把loss、outputs这些带计算图的变量存进了list里。建议先用torch.cuda.memory_summary()看下峰值分配,再确认下dataloader的num_workers和pin_memory有没有踩坑。另外ResNet50最后那个fc层如果没冻结,梯度占的显存也挺可观的,可以先冻住backbone只训分类头试试。
显存涨到20G确实有点夸张,ResNet50+bs8理论上12G左右就够了。建议先用torch.cuda.memory_summary()看看是谁在涨,大概率是某个中间变量没释放,比如你验证集忘了加torch.no_grad(),或者loss累加时把计算图也存下来了。另外检查下dataloader的num_workers和pin_memory,有时候cpu端张量没及时释放也会拖累显存。
第二个epoch才爆显存,八成是缓存没清或者验证阶段没加no_grad,训练循环里loss累积或者保存了带梯度的tensor也会这样。你可以用torch.cuda.memory_summary()看下各阶段分配,再配合torch.cuda.memory_allocated()打印每步增量,基本能定位到是哪块在涨。另外检查下DataLoader的num_workers和pin_memory,有时候是数据侧的问题。ResNet50本身24G跑batch8不该炸,大概率还是代码里有隐藏的显存泄漏。
你这情况大概率不是模型本身的问题,ResNet50加AMP在24G卡上跑batchsize 8不该爆。我猜是验证阶段没加torch.no_grad(),或者loss累加时把计算图也留住了,第二个epoch才崩很符合这个特征。可以用torch.cuda.memory_summary()看看是哪个阶段涨的,再检查下dataloader的num_workers和pin_memory设置。另外迁移学习记得冻结前面的层,只训练fc和后面几个block,能省不少显存。