最近在做一个小项目,想用ResNet50在自定义数据集上做迁移学习,训练图像分类模型。结果每次跑到第二个epoch,显存就飙到20G左右(我只有24G),然后直接OOM挂掉。数据集每张图是224x224,batchsize已经降到8了,还用了混合精度训练(torch.cuda.amp),也试了梯度累积,但感觉治标不治本。
我怀疑是不是模型本身太大,或者我的数据加载方式有问题?但看网上很多人同样配置都能跑,是不是我哪里写的不对?
想请教一下大家:除了换更小的模型,还有没有其他能稳定跑完训练的方法?或者有没有什么检查工具,能帮我定位显存占用到底在哪个层?谢谢各位大佬了。
用PyTorch跑ResNet50做迁移学习,显存总爆掉,求优化思路
全部回复
共 161 条建议先检查下dataloader里num_workers设太大,有时多进程缓存也会爆显存。
试试把数据加载的num_workers设成0,有时候多线程加载会偷偷吃显存。
试试用torch.cuda.max_memory_allocated看峰值在哪,也可能是DataLoader的pin_memory在占显存。
感觉可能是ResNet50的BN层在迁移学习时默认开了track_running_stats,显存占用会随着epoch增加而累积,试试冻结前几层或者用gradient checkpointing,能省不少显存。另外检查一下dataloader的num_workers是不是设太高了,有时候多进程加载也会偷偷吃显存。可以用torch.cuda.memory_summary()看每步的分配,或者装个pytorch_memlab逐层分析。
试试开gradient checkpointing,能省不少显存,虽然慢点但稳。
写得挺好,建议补充一些性能数据。
试试把num_workers设成0,有时候多进程加载反而会缓存额外的显存。
这个问题我前段时间也遇到过,除了混合精度和梯度累积,可以试试用torch.utils.checkpoint来给ResNet50的Bottleneck模块做梯度检查点,能省不少显存。另外检查下你的DataLoader里num_workers是不是设太高了,有时候数据预读取线程也会占显存。如果还想精确定位,用torch.cuda.memory_summary()或者nvidia-smi看逐层占用都挺直观的。
看到这个情况我第一反应是检查一下数据加载那边是不是有内存泄漏,比如每个epoch开新的DataLoader没有释放,或者num_workers设太高导致CPU内存溢出到显存里了。我之前遇到过类似问题,后来用torch.cuda.max_memory_allocated()和torch.cuda.memory_summary()一查,发现是中间特征图没释放,建议你跑完第一个epoch后手动加个torch.cuda.empty_cache()试试。
另外ResNet50的BN层在迁移学习时默认会跟踪running stats,如果你开了eval模式但忘记冻结某些层,反向传播时梯度计算会吃掉大量显存。可以试试把前几层的requires_grad设为False,只训练最后的全连接层,这样参数少很多,显存压力能降一大半。
还有一个小技巧是检查一下你的图像预处理,是不是在GPU上做数据增强?比如用torchvision.transforms在GPU上跑RandomHorizontalFlip之类的,这些操作会双倍占用显存。建议全部挪到CPU端,用DataLoader的prefetch_factor调小一点,配合pin_memory=False试试。
如果还是不行,可以试试梯度检查点(checkpointing),用torch.utils.checkpoint把ResNet50的一部分层包起来,用计算换显存,虽然慢一点但能稳定跑完。我之前在24G卡上跑ResNet152就是这么干的,batchsize能撑到16。
你试试把DataLoader的num_workers调小一点,或者pin_memory=False,有时候多进程预加载会偷偷多吃显存。另外可以装个pytorch_memlab或者torchinfo,跑之前print一下模型每层的参数量,看看是不是全连接层插错了导致参数爆炸。我之前也是224的图跑ResNet50,batchsize设4加梯度累积,16G卡都能稳,你24G不应该啊,检查下是不是把验证集的梯度也保留了?
试试用torch.cuda.memory_summary看详细分配,或者检查下Dataloader的num_workers是不是设太高了。
我也遇到过类似的问题,ResNet50的BN层在反向传播时显存占用会突然暴增。可以试试把torch.backends.cudnn.benchmark设为False,或者手动检查一下DataLoader里是不是开了太多workers导致内存碎片。另外推荐用torch.cuda.memory_summary()打印各层占用,可能发现是中间变量没释放干净。
检查下dataloader的num_workers和pin_memory,开太高容易爆显存,我之前调低就稳了。
你这情况我遇到过类似的,ResNet50本身其实不算特别大,但224x224的图在混合精度下显存冲到20G确实不太正常。我怀疑问题可能出在数据加载或者梯度计算上——比如你用的自定义数据集是不是在每次迭代时都做了额外的数据增强或预处理,导致计算图膨胀了?可以试试把数据预处理全放到DataLoader的num_workers里做,别在__getitem__里写太多变换逻辑。
另外,建议你装个torch.cuda.memory_summary()或者用nvidia-smi实时盯着,跑一个batch就打印一下显存分配,看看到底是模型本身占得多还是中间激活值占得多。如果确认是激活值的问题,可以试试在forward里用checkpointing(梯度检查点)来牺牲一点速度换显存,或者手动释放一些中间变量。
还有一个细节:你用的预训练模型是不是把BN层也冻住了?有时候迁移学习时如果不小心把BN的running_mean和running_var也加入梯度计算,会额外吃显存。我一般先用torch.no_grad跑一个epoch看看基线显存,再逐步排查。你那个梯度累积如果用的是accumulation_steps=4,batchsize其实相当于32,显存压力可能反而更大,可以考虑降到2或1试试。
你这问题我遇到过,ResNet50的BN层在训练时会缓存中间激活值,显存消耗比想象中大。建议试试torch.utils.checkpoint.checkpoint,把每个残差块包进去,用计算换显存,batchsize能翻倍。另外检查下dataloader的num_workers是不是设太高了,有时候多进程加载会额外占显存,2或者4就够了。可以用torch.cuda.memory_summary()打印详细分配,或者nvtop实时看显存变化,定位起来很快。
你这情况我遇到过,24G显存跑ResNet50按理说是够的,建议先检查一下DataLoader的num_workers是不是设太高了,有时候多进程加载反而会缓存一堆中间张量。另外可以试试显存分析工具,比如torch.cuda.memory_summary()或者nvidia-smi看具体哪一步在涨,我上次发现是CrossEntropyLoss的reduction默认sum比mean吃显存多很多。
试过用torch.cuda.memory_summary()打印显存快照没?能直接看到每个tensor的占用,我之前就被DataLoader的pin_memory坑过,开了之后显存莫名其妙多占几个G。另外检查一下是不是把验证集的梯度也保留了,或者反向传播后忘了清空optimizer的梯度,这些细节很容易漏掉。
看到这个情况我第一反应是检查下你数据加载时的num_workers和pin_memory设置,有时候这两个参数没调好会导致显存里缓存太多没释放的数据,尤其是pin_memory=True配合大一点的num_workers容易在DataLoader里堆积张量。另外你试过把梯度累积的batch拆得更细吗?比如batchsize=2然后累积4步,虽然慢点但能彻底解决瞬时显存峰值,比单靠amp稳定得多。我自己的经验是ResNet50的BN层在迁移学习时会保留原数据集统计量,如果你加载的是预训练权重但忘了设置model.train()或者手动冻结某些层,反向传播时的中间变量会异常膨胀,可以用torch.utils.checkpoint对部分残差块做梯度检查点,用时间换空间。要定位具体哪层吃显存,我习惯在训练循环里插几行torch.cuda.max_memory_allocated()打log,或者用torchinfo这个库直接看每一层的输出张量大小,比猜要准。还有个小技巧是检查下你transform里有没有用RandomResizedCrop这类随机操作,它们虽然不占显存但会影响数据预读取的流水线效率,间接拖慢训练。
我最近也踩过这个坑,其实ResNet50加BN层在训练时会缓存中间激活值,这才是吃显存的大头。你可以试试把torch.no_grad()加到eval模式下的前向传播里,或者用torch.utils.checkpoint把部分层的激活值丢掉换计算时间。另外检查一下DataLoader的num_workers是不是设太高了,有时候多进程加载也会偷偷占显存。
我遇到过同样的问题,试了好多方法后发现是数据加载时没关pin_memory和num_workers导致的,你可以检查下dataloader这两个参数,有时候默认设置会多占不少显存。另外可以用torch.cuda.memory_summary()打印一下各层显存分布,或者试试torch.utils.checkpoint对ResNet的某些block做梯度检查点,能省不少显存。还有个小细节:如果用了pretrained的BN层,可以设成eval模式,也能减少一些显存占用。