最近在做一个小项目,想用ResNet50在自定义数据集上做迁移学习,训练图像分类模型。结果每次跑到第二个epoch,显存就飙到20G左右(我只有24G),然后直接OOM挂掉。数据集每张图是224x224,batchsize已经降到8了,还用了混合精度训练(torch.cuda.amp),也试了梯度累积,但感觉治标不治本。
我怀疑是不是模型本身太大,或者我的数据加载方式有问题?但看网上很多人同样配置都能跑,是不是我哪里写的不对?
想请教一下大家:除了换更小的模型,还有没有其他能稳定跑完训练的方法?或者有没有什么检查工具,能帮我定位显存占用到底在哪个层?谢谢各位大佬了。
用PyTorch跑ResNet50做迁移学习,显存总爆掉,求优化思路
全部回复
共 170 条试试用torch.utils.checkpoint把ResNet50的bottleneck包一下,能省不少显存,代价就是慢点。
24G跑ResNet50按理说绰绰有余啊,我怀疑你OOM可能不是模型本身,而是数据管线或者优化器状态的锅。你试过用torch.utils.checkpoint吗?把resnet50的layer3和layer4包进checkpoint函数,前向时就不存中间激活值了,显存能直接砍掉一半多,代价就是多一次反向重算,但训练速度影响很小。另外你确认下是不是开了torch.backends.cudnn.benchmark,有时候这个会自动找算法导致临时显存峰值特别高,关掉试试。还有个小技巧,用torch.cuda.memory_summary()在epoch结束打印一下,能清楚看到到底是模型参数、梯度还是激活值占了空间,我之前就是靠这个发现是DataLoader的num_workers设太高,导致每个worker都拷贝了一份模型副本。如果数据集不大,试试先把图片全部resize成224存成tensor再喂进DataLoader,省掉on-the-fly解码的临时内存。最后实在不行,把batchsize降到4配合梯度累积,但记得把学习率按比例调低,不然收敛会抖。
我跑过类似的配置,224x224加ResNet50,batchsize 8按理说显存不该飙这么狠。你先用nvidia-smi watch看看是不是数据加载时缓存没清,或者num_workers开太多导致CPU内存爆了,这个经常被忽略。另外检查一下是不是把验证集的梯度也保留了,记得用torch.no_grad()包住验证循环。如果还不行,试试冻结前几层只训练后面的层,显存能省不少。
试试把验证集的shuffle关掉,再用torch.cuda.max_memory_allocated打点看看,八成是backward峰值太高。
学到了,感谢分享!
我之前也遇到过类似情况,后来发现是数据加载时没开num_workers和pin_memory,导致CPU预处理跟不上,GPU一直在等数据但显存却莫名被占满。你可以先用nvidia-smi监控一下,看看是不是数据加载的临时张量堆积了。另外检查下是不是在验证阶段也开了梯度,或者把BatchNorm层设成eval模式了,这些细节挺容易忽略的。
试试把图片尺寸缩到192或者用torch.utils.checkpoint,省显存效果立竿见影。
试试把resnet50的BN层冻住,用torch.utils.checkpoint换显存,能省不少。
检查工具就用nvidia-smi看分时占用,或者pytorch的profiler,能定位到具体层。
试试把workers调成0和pin_memory关掉,有时是DataLoader的预取在吃显存,不是模型本身的问题。
我之前也踩过这个坑,ResNet50的BN层在迁移学习时特别吃显存,尤其是如果冻结了前面层但没关掉BN的统计更新,中间变量会堆得离谱。你可以试试把模型切成几个阶段,用torch.utils.checkpoint把中间激活值丢掉,牺牲点速度但显存能省一半。另外检查下dataloader是不是num_workers设成0了,有时候数据预处理也在占显存,比如ToTensor和Normalize放GPU上跑了。真想定位的话,用pytorch的torch.cuda.memory_snapshot或者nvidia-smi看每个进程的显存分配,不过最直接的还是设个hook打印每层输出大小,基本就能看出是哪个block炸的。
之前我也遇到过类似的情况,后来发现是数据加载时num_workers设太大,CPU预处理成了瓶颈,显存反而被临时张量占满了,调到4之后立刻好转。你可以先用nvidia-smi看看是不是数据加载峰值导致的。另外检查一下是不是把验证集的drop_last忘了,最后一个不完整的batch有时候会特别吃显存。真要定位层的话,可以用torch.profiler,能直接看到每个op的显存分配,我上次就是靠它发现是BatchNorm的running stats在反向传播时占了额外空间。
说实话你这情况我太熟了,之前调一个医学影像分类也这样,ResNet50加224输入按理说真不该这么吃显存,20G明显不正常。你先别急着怀疑模型,我建议用torch profiler或者直接跑一个batch然后看nvidia-smi,重点查一下是不是dataloader的num_workers设太大,或者pin_memory开着导致CPU内存和显存之间拷贝爆炸,有时候数据增强的transform在GPU上跑反而更占显存。另外你检查下是不是把loss和梯度保留在计算图里了,比如每个step后没手动清零optimizer.zero_grad,或者用了梯度累积但累积次数没配合好,导致反向传播的中间变量一直堆积。还有个冷门的坑,如果用了torchvision的预训练模型,记得把BN层设成eval模式,特别是迁移学习时BN的running stats更新会额外保留batch维度的中间量。实在不行就试试冻结前几层只训练最后几层,显存能砍掉一大半,而且小数据集上效果差别不大。你要是能贴一下训练循环代码片段,大家可能更快帮你定位。
试试把resnet50的bn层冻住,还有检查下dataloader的num_workers是不是开太多,内存页表炸了也影响显存。
或者
用torchsummary打印每层输出尺寸,再配合nvidia-smi看峰值,八成是backward时中间激活值爆的,考虑开个gradient checkpointing。
我之前也被这个坑过,224x224加batchsize 8按理说不该爆的,你检查过dataloader的num_workers和pin_memory没?有时候数据加载线程太多反而会占额外显存,尤其是Windows下特别明显。另外你用的是torchvision的预训练模型吗?如果是的话,记得把BN层设成eval模式,只训练最后几层或者用frozen backbone,不然梯度回传会缓存大量中间激活值,那才是显存大头。还有个思路是检查一下是不是优化器状态占的,Adam自带两倍参数量的动量,换SGD能省不少。想定位具体层的话,可以试试torch.cuda.memory_snapshot或者用pytorch的profiler,能看到每个tensor的分配点。我之前是把224输入先resize到160跑通流程,确认没问题再慢慢加回去,虽然麻烦但很稳。另外混合精度记得把grad_scaler的scale_factor调小一点,有时候loss spike会导致scale爆炸然后显存突然飙升,我遇到过一次。你试试把batchsize降到4然后梯度累积step设成2,理论上等效但显存峰值能砍一半。
我之前也遇到过一模一样的情况,后来发现是数据加载那边出了问题,你试试把num_workers调高,或者用pin_memory=True,有时候瓶颈不在模型本身。另外你可以用torch.cuda.memory_summary()看一下是不是有碎片化,或者跑之前先用torch.autograd.detect_anomaly()排查一下,我上次就是被一个中间变量给坑了。话说你用的是ImageFolder还是自己写的Dataset?如果是后者,可能返回的batch里带了很多额外显存开销,检查一下会不会是没做归一化导致数值爆炸了。
我之前也踩过这个坑,ResNet50的BN层在迁移学习时特别吃显存,尤其是batchsize小的时候反而更严重。你可以试试把冻结层的batchsize调大,只对最后几层做反传,或者干脆用FrozenBN,能省不少显存。另外检查下是不是数据加载时num_workers开太多,CPU和GPU争抢也会导致显存峰值异常,建议先用单进程跑一次对比下。还有个小技巧,用torch.cuda.memory_summary()看下峰值在哪分配,有时候是优化器状态和激活值缓存的问题,不是模型本身。
我之前也遇到过类似情况,224的输入+ResNet50按理说8的batch不该爆,你先看看是不是dataloader的num_workers开太多,或者pin_memory=True导致显存和内存互相拖累。另外建议用torch.cuda.memory_summary()打印一下,我猜大概率是backward时候的中间激活值占大头,可以试着把模型换成resnet50的preact版本,或者干脆冻结前几层只微调后面,显存能降一大截。还有个小技巧,输入图片别用transform直接做归一化,放GPU上做能省点临时显存,虽然不多但够你多撑一个epoch了。
我之前也遇到过类似情况,你检查过DataLoader的num_workers和pin_memory没?有时候数据加载占的显存比模型还夸张,尤其是224x224加上多进程,试试把workers调低或者换prefetch_factor。另外建议你跑一下nvidia-smi看下训练时的显存分配,别只盯torch.cuda.max_memory_allocated,有时候是cuDNN的workspace在作祟,设个torch.backends.cudnn.benchmark=False能省不少。还有个偏方,把BatchNorm换成GroupNorm试试,ResNet50在batchsize小的时候BN统计量不稳也容易爆显存。
试试把batchsize再砍到4,然后开gradient checkpointing,能省不少显存,ResNet50用这个很稳。
换个思路查一下dataloader的num_workers,有时候数据加载和GPU计算重叠不好也会显得显存峰值高,调成4或8试试。
我原来也遇到过类似情况,后来发现是数据加载那步没做好,transforms里如果带随机裁剪之类的操作,会额外占不少显存,试试把预处理放到CPU上做,只把tensor丢给GPU。另外你可以用torch.cuda.memory_summary()看下峰值在哪,很多时候是backward时中间变量爆的,开一下cudnn.benchmark再配合gradient checkpointing能省不少。还有个笨办法,把batchsize再压到4,然后多累积几步,虽然慢点但至少不崩,我之前就这么跑通的。