最近在做一个小项目,想用ResNet50在自定义数据集上做迁移学习,训练图像分类模型。结果每次跑到第二个epoch,显存就飙到20G左右(我只有24G),然后直接OOM挂掉。数据集每张图是224x224,batchsize已经降到8了,还用了混合精度训练(torch.cuda.amp),也试了梯度累积,但感觉治标不治本。
我怀疑是不是模型本身太大,或者我的数据加载方式有问题?但看网上很多人同样配置都能跑,是不是我哪里写的不对?
想请教一下大家:除了换更小的模型,还有没有其他能稳定跑完训练的方法?或者有没有什么检查工具,能帮我定位显存占用到底在哪个层?谢谢各位大佬了。
用PyTorch跑ResNet50做迁移学习,显存总爆掉,求优化思路
全部回复
共 170 条我之前也踩过这个坑,ResNet50跑224的图按理说不该这么吃显存,你先用nvidia-smi盯一下是不是数据加载那边没做pin_memory,或者worker数太少导致CPU卡住,显存反而被缓存占满了。另外可以试试把BN层冻结掉,只训最后的全连接层,显存能省一大截,效果也不差。排查工具的话,pytorch自带那个torch.profiler挺好用的,能看每层内存分配,或者直接hook一下模型的forward,打印每层输出tensor的shape和内存,比瞎猜快多了。
试试用torch.profiler看下显存峰值在哪,另外检查下ResNet50的BN层是不是在冻结参数时没设eval模式。
把数据加载的num_workers调高,pin_memory打开,有时候瓶颈不在GPU而在CPU喂数据太慢。
试试把batchsize再砍到4,224分辨率对ResNet50来说真不小,amp省的是计算不是显存。
或者用torch.profiler跑一下,能直接看每层显存峰值,比瞎猜快多了。
说实话224的输入8的batch按理说不该这么吃显存,你先用nvidia-smi盯着看是不是数据加载那边把显存占满了,比如num_workers开太高或者pin_memory=True有时候反而会炸。另外检查一下是不是把验证集的梯度也保留了,或者模型里有个别层没设成eval模式导致反向传播额外开了一倍显存。我上次也遇到过类似情况,最后发现是dataloader里每张图都做了随机resize,导致缓存了太多中间变量,你试试固定输入尺寸或者用transforms直接归一化应该能缓解不少。
我遇到过一模一样的坑,ResNet50的BN层在迁移学习里特别吃显存,尤其是你如果没冻结前面几层的话。建议先试试把backbone的梯度关掉只训练分类头,显存能掉一半还多,等收敛了再解冻微调。另外检查下你dataloader的num_workers,有时候数据加载线程会复制模型权重,那个也占显存的。可以用nvidia-smi看下显存是不是被多个进程分了,或者试试torch.cuda.memory_summary(),能直接看到每个tensor的占用,比瞎猜强。
我之前也遇到过类似情况,排查下来发现不是模型本身的问题,而是数据加载时pin_memory和num_workers开太多,把显存和内存都吃满了。你可以先试下把这两个参数调小,再把batchsize降到4,配合梯度累积,基本能稳住。另外用torch.cuda.memory_summary()看一下,大概率是激活值占大头,可以试试在forward里加个torch.no_grad()做特征提取,只微调最后一层,显存占用能降一半以上。
我之前也踩过这个坑,ResNet50加224输入,batchsize8按理说不该爆的,你检查过num_workers和pin_memory没?有时候数据加载线程一多,CPU来不及喂数据,GPU反而会缓存一堆临时张量,显存就莫名其妙上去了。
建议你先用torch.cuda.memory_summary()看看分配峰值到底在哪,大概率不是模型参数,而是激活值或者中间变量。另外可以试试把模型切成几个阶段,手动记录每层输出尺寸,用torch.profiler跑一个step,它会按算子给你显存占用排行,比瞎猜强多了。
我怀疑你可能是用了固定的ImageNet预训练权重,然后忘了冻结BN层?迁移学习时如果BN在训练模式下更新,会额外保存每个batch的统计量,显存消耗比想象中高不少,可以试试用model.eval()冻结BN,或者改用GroupNorm替换。
还有个偏方,把输入尺寸临时降到192或者160跑通流程,能确认是不是尺寸导致的激活值爆炸。如果降尺寸后显存正常,那就是ResNet50的stage4输出太大,可以改一下最后一个block的stride,或者用空洞卷积替代下采样。
梯度累积确实治标不治本,因为累积只是延缓了反传频率,但前向和反向的峰值显存没变。你可以试试mixed precision的grad_scaler设置,有时候scaler会额外缓存动态loss的缩放因子,在长训练时积累碎片。
最后实在不行,就切到torch.utils.checkpoint,把stage3和stage4包起来,用计算换显存,虽然会慢个20%左右,但至少能稳定跑完。别急着换模型,ResNet50在24G卡上绝对能跑,肯定是某个隐藏细节没调对。
说实话你这配置跑不动真不是你的问题,ResNet50本身不算大,但迁移学习时如果加载了ImageNet预训练权重,第一层卷积和最后的全连接层会额外占用不少显存,尤其是梯度回传时中间激活值特别吃内存。我之前用torchsummary或者pytorch的memory_stats接口查过,发现大头根本不是模型参数,而是BN层和激活值缓存,你可以试试在forward里关掉一些层的梯度,比如用requires_grad=False冻结前几层,只训练最后几层,这样显存能直接砍半。
另外你提到梯度累积,那个只能缓解batchsize太小带来的BN统计量不准问题,对显存峰值没帮助,真正管用的是减少输入分辨率或者用RandAugment做在线裁剪,我一般把224改成192,精度几乎不掉,显存却能省出3-4G。还有个小坑,检查下你的DataLoader是不是num_workers开太多,或者pin_memory=True,有时候数据预取也会悄悄吃显存,尤其是多进程拷贝时。
工具方面别用torchsummary了,它只看参数不看激活,推荐用torch.profiler或者memory_profiler,能按层输出分配峰值,我之前定位到是resnet的layer3输出张量太大,果断改了空洞卷积替代下采样,直接稳了。最后实在不行,你可以考虑用huggingface的transformers里的AutoModelForImageClassification,它内部做了梯度检查点,牺牲一点速度换显存,24G跑batchsize 16都试过没问题。
我之前也踩过类似的坑,后来发现多半是数据加载的锅,尤其是随机裁剪或增强时开了太多worker,显存会被临时张量占满。你可以先试试把num_workers调成0,再配合pin_memory=False,看OOM会不会推迟。另外,torch.cuda.memory_summary()能看每个张量的分配情况,定位是不是forward里某个中间变量没释放。还有个小技巧,用gradient_checkpointing把ResNet50的中间激活重算一遍,显存能降一半,速度稍微慢点但稳得很。
我之前也踩过这个坑,ResNet50跑224输入按理说不该这么吃显存,24G爆掉肯定有隐藏问题。你先别急着换模型,用nvidia-smi盯着看,大概率不是模型本身,而是数据加载或者计算图没释放。我建议你试试torch.cuda.empty_cache()在每轮epoch后手动清一下,有时候是缓存碎片累积导致的假性OOM。另外,检查下是不是把验证集的梯度也算了,或者不小心把整个数据集都塞进GPU了,我之前犯过把tensor直接.to(device)后忘了detach的错。还有一个关键点,混合精度下要确保所有层都支持fp16,有些自定义层回退到fp32会瞬间吃掉大量显存,你可以用torch.autograd.detect_anomaly()或者profiler看看具体哪块峰值最高。如果还不行,就把batchsize再砍到4,同时把图像预处理里的RandomResizedCrop改成简单的Resize,减少计算图复杂度。说实话,梯度累积确实治标不治本,它只是省了反向传播的峰值,但前向激活值还是占着。最后实在不行就换个DenseNet或者EfficientNet,效果不差,显存压力小一半。
我之前也踩过这个坑,224x224加ResNet50按理说8的batchsize不该爆,你先检查下是不是dataloader的num_workers开太多,或者pin_memory=True把显存当缓存用了。另外可以试试把模型的BN层冻住,再用torch.utils.bottleneck或者nvidia-smi dmon看实时占用,大概率是优化器状态和中间激活值在作祟,梯度累积配合更小的batchsize(比如2)反而能稳住。
试试把batchsize再砍到4,然后开gradient checkpointing,ResNet50这招最管用,显存直接砍半。
我之前也遇到过类似情况,ResNet50加自定义数据集特别容易在BN层上吃显存,尤其是batchsize小的时候。你可以试试把torch.backends.cudnn.benchmark设成True,再手动把输入图像用prefetch加载到固定内存,有时候数据加载线程的瓶颈也会让显存看起来飙升。另外建议用torch.cuda.memory_summary()跑一下,能直接看到每个tensor的占用分布,很多时候问题出在优化器状态或者中间激活值上,不一定是模型本身。
试试把batchsize再砍到4,开amp的同时关掉grad checkpointing,ResNet50吃显存主要在bn层。
用nvidia-smi看下是不是数据加载线程占的缓存没释放,设num_workers=0跑一版对比下。
我之前也遇到过这问题,你试试把ResNet50的BN层换成GroupNorm,或者冻结前几层只训后面,显存能降不少。另外检查下DataLoader的num_workers和pin_memory,有时候数据加载卡住也会导致显存虚高。用torch.cuda.memory_summary()能看到具体分配,但层级别定位可以开torch.profiler,能看到每个op的显存峰值。还有个小技巧,把输入尺寸临时降到192跑通流程,确认不是代码bug再调回224。
试试把验证集的梯度也关掉再查下dataloader的num_workers,大概率是验证阶段没加no_grad偷偷吃显存了。
说实话我也踩过这个坑,ResNet50的BN层在迁移学习里特别吃显存,尤其你如果没冻结前几层的话。建议先用torchsummary或者pytorch的memory_stats接口看看具体哪块分配最多,我猜是激活值缓存而不是模型参数。另外可以试试把输入尺寸临时降到160x160跑通流程,显存能省一半,等调好超参再改回224。还有个小技巧,梯度累积别和amp一起用,有时候反而会触发奇怪的显存碎片化,直接开gradient_checkpointing配合batchsize=16可能更稳。
试试把batchsize再砍到4,然后开gradient checkpointing,ResNet50吃显存主要是中间激活值,这个能省一大截。
试试把resnet50的BN层冻结掉,再用torchsummary看下每层显存,八成是backward时梯度爆的。
我上次也是这么解决的,把batchsize再砍半加梯度累积,稳得很。
我之前也踩过类似的坑,24G看着挺大但ResNet50加个BN层在反传时真的吃显存。你可以先试试把transform里的Normalize挪到GPU上做,或者用torch.utils.checkpoint把ResNet的block包一下,能省不少激活值内存。另外检查下是不是数据加载时每步都做了重复的增广操作,把worker数量和prefetch_factor调低点,有时候CPU瓶颈反而会让GPU缓存堆积。还有个笨办法,直接打印每个module的显存占用,用pytorch的memory_profiler或者torch.cuda.memory_snapshot看看是不是卡在最后的全连接层,之前我就是在fc层前加了个dropout忘了删,白白多占2G。