最近在做一个小项目,想用ResNet50在自定义数据集上做迁移学习,训练图像分类模型。结果每次跑到第二个epoch,显存就飙到20G左右(我只有24G),然后直接OOM挂掉。数据集每张图是224x224,batchsize已经降到8了,还用了混合精度训练(torch.cuda.amp),也试了梯度累积,但感觉治标不治本。
我怀疑是不是模型本身太大,或者我的数据加载方式有问题?但看网上很多人同样配置都能跑,是不是我哪里写的不对?
想请教一下大家:除了换更小的模型,还有没有其他能稳定跑完训练的方法?或者有没有什么检查工具,能帮我定位显存占用到底在哪个层?谢谢各位大佬了。
用PyTorch跑ResNet50做迁移学习,显存总爆掉,求优化思路
全部回复
共 160 条24G显存跑ResNet50还OOM,batchsize=8加混合精度还崩,这不太正常。你检查下dataloader的num_workers是不是设太高了,或者pin_memory=True导致内存泄漏?另外可以用torch.cuda.memory_summary()打印分配细节,或者挂nvidia-smi看看是不是有残留进程占着显存。还有检查下你的预处理是不是无意中把图像尺寸放大了,或者训练时同时开了验证集的梯度计算。
这情况我也遇到过,ResNet50吃显存确实猛,尤其backbone没冻结的时候。建议先检查下是不是开了梯度checkpointing,PyTorch自带torch.utils.checkpoint能省不少显存,就是会慢点。另外可以试试用torch.cuda.memory_summary()打印下具体哪块爆的,说不定是DataLoader里pin_memory和num_workers开太高了。
这个问题我前阵子也遇到过,其实ResNet50本身还好,问题大概率出在DataLoader的num_workers开太多或者缓存了中间变量。可以试试torch.cuda.empty_cache()在每个epoch结束后清一下缓存,另外检查下是不是用了太大的pin_memory,或者把梯度检查点(activation checkpointing)打开,能省不少显存。至于定位工具,推荐用torch.cuda.memory_summary()或者nvidia-smi盯着看,有时候是数据增强那一步偷偷存了图没释放。
同感,我之前用ResNet50跑224的图,batchsize开到16都能跑,你这8就爆了确实奇怪。会不会是DataLoader的num_workers设太高,或者预处理里无意中做了多余的数据拷贝?可以先跑一个batch看看显存占用,用torch.cuda.memory_summary()能打印出每个tensor的分配情况。另外检查一下是不是把验证集的梯度也存了,关掉requires_grad能省不少。
这问题我碰到过好几次,ResNet50其实不算特别大,但迁移学习的时候梯度回传和中间激活确实能把显存吃满。你batchsize降到8还爆,我猜可能是数据加载那边无意中把图像尺寸搞大了,或者预处理里插入了什么额外的tensor没释放。建议先跑一个简单的profile:用torch.cuda.memory_summary()或者nvidia-smi盯着看,哪个step之后显存突然跳变。另外检查一下你的DataLoader,是不是num_workers开太多导致内存复制开销,或者pin_memory=True在部分显卡上反而有bug。
还有一个容易忽略的点——你是不是把整个ResNet50的backward梯度都保留了?迁移学习通常只需要更新最后几层,前面层可以设requires_grad=False,这样反向传播时中间激活图就不会全量保留,显存能省将近一半。用torch.no_grad()包住特征提取部分,或者先把backbone的权重冻结跑一遍,等收敛得差不多了再解冻微调。
如果还是不行,可以试试activation checkpointing(梯度检查点),PyTorch直接有torch.utils.checkpoint,它用计算换显存,训练速度会慢一点点,但24G卡跑224x224的ResNet50,batchsize 8应该稳如老狗。还有,检查下你混合精度里是不是忘了用GradScaler,有时候amp自动混精没生效,float32跑起来显存直接翻倍。
最后说个坑:自定义数据集路径里如果有中文或者特殊字符,PIL加载图像时可能会反复申请内存不释放,我就被这个坑过一天……
你这情况我太熟了,ResNet50按理说224的图batch 8不至于24G撑不住,除非你输入尺寸不是224x224,或者中间有什么奇怪的预处理把图变大了。建议先检查一下transform里有没有不小心resize到更大的尺寸,或者dataloader的num_workers设得太高导致内存碎片。
不过我更怀疑是优化器状态和中间激活值的问题。ResNet50虽然不算轻量,但batch 8的激活值撑死也就几个G,关键是Adam优化器会额外存一阶二阶动量,加上混合精度下的master weights和梯度,这些叠加起来很容易超。你可以试试把优化器换成SGD(momentum=0.9,weight_decay设小点),这样显存能省下3-4G。如果你的任务对收敛速度不太敏感,效果差不了太多。
另外推荐用torch.cuda.memory_summary()直接打印当前分配和预留的显存块,能看出来到底是模型参数、激活值还是优化器状态在吃显存。更细的可以用torch.cuda.set_per_process_memory_fraction(0.9)提前限制最大使用量,避免突然OOM。如果还想省,可以考虑gradient checkpointing,PyTorch有现成接口,对ResNet50能省30%左右的激活显存,代价就是慢一点。
最后提醒一下,混合精度训练下loss scaling可能会把某些层炸掉,检查一下是不是用了torch.cuda.amp.GradScaler,并且确认没在forward里手动对某些层做fp32强制转换。有时候一个小细节没注意,显存就翻倍。
同感,ResNet50在24G卡上跑迁移学习还爆显存确实挺头疼的。我上周也遇到过类似情况,batchsize降到4才勉强稳住,但训练速度直接崩了。
先帮你排查几个常见坑:第一,检查一下是不是把整个ResNet50的梯度都打开了?迁移学习通常只训练最后一层或最后几层全连接层,前面特征提取层要冻结掉(requires_grad=False)。这样反向传播时只更新少量参数,显存占用会直线下降。第二,数据加载里有没有做额外的数据增强?有些增强操作(比如RandomResizedCrop)在GPU上做会额外占用显存,建议全放在CPU上预处理。
关于定位工具,推荐用torch.cuda.memory_summary()打印详细分配,或者装个pytorch_memlab库,可以在每个iteration前后输出当前显存峰值。还有个骚操作:在你怀疑的层前后加torch.cuda.reset_peak_memory_stats()和torch.cuda.max_memory_allocated()对比,能精确定位哪块儿最吃显存。
另外你提到混合精度和梯度累积都试了还不行,会不会是DataLoader的num_workers开太多?我之前设成8直接多吃了2G显存,降到4就正常了。还有就是图像尺寸虽然是224x224,但ResNet50第一个卷积层输出通道数很大,可以试试把模型的stem部分替换成更轻量的结构(比如用3x3卷积替代7x7),网上有现成的ResNet50变体。
最后检查一下是不是把验证集的梯度也打开了?或者有没有在训练循环里意外保留了中间变量?比如计算loss时没有用loss.item()而是直接保留tensor,这种小细节经常被忽略。如果还不行,可以考虑用DeepSpeed的ZeRO优化器,它对显存占用优化非常暴力。
试试把num_workers调成0,有时候多进程加载数据也会偷偷占显存。
说实话你这情况我遇到过类似的,224x224的图用ResNet50按理说batchsize8不应该爆20G,除非你反向传播的时候中间变量没释放干净。建议你先用torch.cuda.memory_summary()看看具体哪个操作在吃显存,很多时候问题出在loss计算或者自定义的data augmentation上,比如用了太大的随机裁剪缓存。另外检查下你是不是把验证集的梯度也保留了,或者dataloader的num_workers开太多导致内存碎片化。混合精度虽然好但如果你模型里用了某些不支持的op,它可能会回退到fp32,这样显存反而更崩。还有个偏方:试试把模型的第一层卷积或者最后的全连接层换成更小的变体,比如用ResNet50但把stem改成7x7步长2的普通卷积,显存能省不少。实在不行就上梯度检查点(checkpointing),用时间换空间,虽然慢点但至少能跑完。
这情况我也遇到过,ResNet50按理说24G显存不至于这么容易爆,除非你输入尺寸或者数据预处理环节有隐藏的内存泄漏。可以先试试用torch.cuda.memory_summary()和nvidia-smi配合跑一个epoch,看显存是不是每个batch都在持续增长,如果是的话,大概率是优化器状态或者中间变量没释放干净,比如某些hook或者自定义loss里不小心保留了计算图。另外检查下DataLoader的num_workers是不是设得太高,有时多进程加载图片如果用了pin_memory=True但没配合好,也会导致显存碎片化。还有一个冷门技巧:把torch.backends.cudnn.benchmark设为False,偶尔能缓解动态shape带来的显存波动。如果这些都不行,可以试试把模型的第一层卷积输入改成更小的尺寸(比如160x160),毕竟迁移学习下游任务往往不需要那么高的分辨率,效果损失很小的同时显存能降一大截。
试试用torch.utils.checkpoint把中间激活值丢掉,能省不少显存。另外检查下dataloader的num_workers是不是设太高了。
我之前也遇到过类似的情况,后来发现是DataLoader的num_workers设太高了,缓存了大量预处理图像导致显存暴涨,降到4或者2试试看。另外可以试试torch.utils.checkpoint,把ResNet的某些大层做梯度检查点,能省不少显存。显存占用可以用torch.cuda.memory_summary()看个大概,或者装个nvidia-smi实时盯着,哪个阶段涨得厉害就排查哪块。
说真的,ResNet50在224x224输入下,按理说不应该这么吃显存,24G都快吃满了肯定哪里不对劲。你试试在训练循环里加个torch.cuda.reset_peak_memory_stats(),然后用torch.cuda.max_memory_allocated()看看到底哪个阶段内存峰值最高,我遇到过类似情况,最后发现是DataLoader的num_workers开太多,每个worker都在复制一份模型参数。另外检查一下你是不是把验证集的梯度也保留了,很多人会在val_loop里忘记加torch.no_grad(),那显存直接翻倍。还有个小技巧,把模型里的BatchNorm换成可选的SyncBatchNorm,虽然对显存没帮助,但有时候会触发一些奇怪的缓存问题。如果这些都试过了,可以试着用torch.utils.checkpoint来激活重计算,虽然慢一点,但能把激活值显存压到原来的三分之一,对ResNet这种结构效果挺明显的。
这个问题我也遇到过,ResNet50加上BN层和全连接真的挺吃显存的。你可以试试把输入尺寸降到192x192,或者用torch.utils.checkpoint做梯度检查点来换空间,虽然慢点但能跑稳。另外检查一下数据加载时是不是每个epoch都在重复复制图像,用pin_memory=True和num_workers>0能省不少临时显存。
试试梯度检查点(gradient checkpointing),能省不少显存,就是慢点。另外检查下dataloader是不是num_workers设太高了。
讲真,24G显存跑ResNet50按理说不至于这么惨,我猜你可能是把验证集或者测试集的梯度也保留了?很多人容易忽略model.eval()下忘记包torch.no_grad(),或者DataLoader的num_workers设得太大导致显存缓存。可以先用torch.cuda.memory_summary()直接看每个tensor的分配情况,爆在哪一层一目了然。另外试试把输入图像的通道数检查一下,有时候加载了四通道图或者16位深度图,显存会翻倍。还有个小技巧:用gradient_checkpointing,在ResNet的forward里把某些层设为checkpoint,用计算换内存,虽然慢一点但能稳住。你用的优化器是不是Adam?它的动量和方差缓存也挺吃显存,换成SGD+余弦退火能省出一截。最后建议把batchsize再降一档到4,配合梯度累积步数调大(比如8步),其实收敛速度并不会慢太多,主要先跑通再调优。
感觉你这情况不完全是模型或batchsize的问题,ResNet50本身没那么吃显存。建议先检查下数据加载部分,是不是把整图或者增强后的图全缓存到显存了,可以试试pin_memory=False或者用DataLoader的num_workers稍微调低点看看。另外可以用torch.cuda.memory_summary()打印下分配情况,或者装个nvidia-smi的top工具实时盯一下,有时候是反向传播时中间变量没释放干净。
我最近也踩过这个坑,排查下来发现是dataloader的num_workers设得太高,加上pin_memory=true,GPU显存里会多存好几批预处理好的数据。你可以试试把num_workers降到0或者2,关掉pin_memory看看有没有缓解。另外建议用torch.cuda.memory_summary()打一下显存分配详情,经常能发现是中间变量或者优化器状态占了大头。
试过把数据加载的num_workers调高或者调低吗?有时候多进程加载会额外占显存,尤其在你显存已经吃紧的情况下。另外可以检查一下是不是在计算loss的时候把整个模型的输出都保留了梯度,用detach()或者把不需要梯度的参数冻结掉应该能省不少。建议用nvidia-smi配合torch.cuda.memory_summary()实时看每步的显存变化,比瞎猜靠谱。
试试把数据加载的num_workers设成0或者4,有时候worker太多也会造成显存泄漏。另外检查一下验证集是不是也跑了梯度计算,记得用torch.no_grad()包一下。可以用torch.cuda.memory_summary()看看具体哪块占得多,我遇到过是BatchNorm层的缓存问题,清空一下缓存或者用torch.backends.cudnn.deterministic=True能缓解。