最近在做一个小项目,想用ResNet50在自定义数据集上做迁移学习,训练图像分类模型。结果每次跑到第二个epoch,显存就飙到20G左右(我只有24G),然后直接OOM挂掉。数据集每张图是224x224,batchsize已经降到8了,还用了混合精度训练(torch.cuda.amp),也试了梯度累积,但感觉治标不治本。
我怀疑是不是模型本身太大,或者我的数据加载方式有问题?但看网上很多人同样配置都能跑,是不是我哪里写的不对?
想请教一下大家:除了换更小的模型,还有没有其他能稳定跑完训练的方法?或者有没有什么检查工具,能帮我定位显存占用到底在哪个层?谢谢各位大佬了。
用PyTorch跑ResNet50做迁移学习,显存总爆掉,求优化思路
全部回复
共 170 条我之前也遇到过类似情况,224的输入其实不算大,但ResNet50的中间特征图确实吃显存。你试试把batchsize降到4,然后配合梯度累积到等效16,同时把数据加载的num_workers调高,有时候瓶颈在数据管道上。另外检查下是不是用了冻结BN层,迁移学习时BN的running stats更新也可能额外占显存。用nvidia-smi看显存曲线最直接,也可以试试torch.cuda.memory_summary(),能定位到具体是哪层分配的内存。
检查下dataloader的num_workers是不是设太高了,有时候数据加载也会占显存。另外用torch.cuda.memory_summary()看下每个tensor占用。
我之前也踩过这个坑,224x224配ResNet50按理说不该这么吃显存,20G肯定不正常。你先别急着怀疑数据加载,大概率是反向传播时中间激活值没释放,或者torch.cuda.amp在某些层精度转换时额外保留了FP32副本。建议你装个pytorch的memory profiler,比如torch.profiler或者直接看nvidia-smi的显存曲线,把每个step的峰值前后对比一下,能看到是不是某个block特别占内存。另外你试试把batchsize再砍到4,然后同步把learning rate调低,有时候梯度累积反而会让显存碎片化更严重,不如直接用gradient checkpointing,虽然慢点但能大幅减少激活值存储。还有一个容易忽略的点,检查下dataloader的num_workers是不是设太高了,如果内存吃紧系统会频繁swap,导致显存被临时占用。我之前是把transform的归一化放到GPU上做,虽然省了CPU时间,但显存多了好几G,后来改回CPU预处理就好了。如果还是爆,就把BN层换成GroupNorm或者用低精度的ResNet变体(比如torchvision里直接有的resnet50d),别硬扛原版。最后,确认一下你是不是用了预训练权重,如果是从头训练那梯度和BN统计量都会更占资源。
试试把验证集也塞进amp里,或者查查是不是pin_memory开太多,我之前这么弄直接省了4G。
检查下是不是把验证集的梯度也算了,关掉torch.no_grad能省不少显存。
我也遇到过类似情况,24G卡跑ResNet50按理说挺宽裕的,你查下是不是数据加载那部分没做归一化或者pin_memory没开,有时候CPU瓶颈会导致显存里堆积太多中间张量。另外建议用torch.utils.checkpoint把resnet50的bottleneck包一下,开启激活检查点,虽然会慢一点但显存能砍掉差不多一半,训练完再取消就行。最后可以试试nvtop或者pytorch的torch.profiler,能直接看到每个op的显存分配,我之前就是靠这个发现是BatchNorm的running stats在反向时被重复计算了。
我之前也遇到过一模一样的情况,后来发现是数据加载时num_workers设成0了,导致CPU预处理跟不上,GPU一直在空等,显存反而被中间变量占住不释放。你可以先试试把num_workers调到4或8,顺便开一下pin_memory,有时候这个比调batchsize管用得多。另外推荐用torch.profiler看下显存分配,能具体到每个op,我之前定位到是BatchNorm的running stats在反向传播时额外吃了一大块,后来换成冻结BN层才解决。还有个偏方是把输入图片先resize到192x192跑通流程,确认没问题再上224,排查是不是数据增强里有些操作在偷偷缓存。
试试把验证集的batchsize调成1,很多人的OOM其实卡在验证阶段而不是训练。
说到这个我太有共鸣了,之前调ResNet50也差点被显存搞疯。不过你既然开了AMP和梯度累积还爆,我猜问题大概率不在batchsize上,而是你反向传播时把中间激活值全保留了。可以试试在训练循环里显式加上torch.cuda.empty_cache(),虽然治标不治本,但至少能看清是不是碎片化问题。另外强烈建议你用nvidia-smi dmon或者pytorch的torch.profiler看下每个op的显存分配,我之前发现光是BatchNorm层的running stats在迁移学习时也会占不少显存,而且你数据集小的话,完全可以冻结前几层只训最后的全连接层,这样显存直接砍半。还有个骚操作是换用ResNet50的efficient implementation,比如torchvision里带weights的版本可能比你自己写的forward更省显存,或者试试把图片尺寸降到192x192,反正迁移学习对分辨率没那么敏感。最后,如果你用的是自定义Dataset,检查下是不是num_workers设太高导致内存拷贝到显存的通道堵住了,我之前就吃过这个亏。
说到显存爆掉,我第一反应是你可能把验证集或测试集的forward也包在torch.no_grad外面了,但推理时的batchsize没跟着降,有时候验证阶段反而更容易爆。另一个常见坑是ResNet50的最后一个全连接层替换后,如果没冻结前面层,反向传播会保留所有中间激活值,你用amp只是降低精度,但激活值的内存占用大头其实还在。建议先用torch.cuda.memory_summary()看看是哪块分配器在涨,再配合nvidia-smi监控,确认是不是数据加载worker进程也占用了显存(比如pin_memory=True但num_workers开太高)。我之前遇到过类似情况,最后发现是DataLoader里collate_fn写得太重,每batch都做了原地增强操作,导致额外拷贝了张量。你试试把batchsize再砍到4,然后梯度累积步数设成4,同时用torch.utils.checkpoint(就是激活检查点)去换空间,ResNet50这种残差结构特别适合,能把激活内存降到原来的三分之一。另外检查一下是不是用了pretrained=True时,BN层在迁移学习里默认跑训练模式,如果数据集小,BN的统计量更新会很吃显存,可以尝试冻结所有BN层或者改用group norm。我自己的经验是,24G卡跑ResNet50 + 224分辨率,batchsize 8理论上是能过的,除非你输入尺寸实际是256或者模型里不小心加了额外的分支。你跑一下代码定位到具体哪一行爆掉,大概率是loss.backward()时梯度和优化器状态占了大头,如果是这样,可以试试AdamW的foreach=True或者换SGD,有时候优化器状态的内存差别很大。
试试用torch.utils.checkpoint把ResNet50的前几个stage包一下,能省不少显存,速度损失也不大。
或者先跑个batch看看每层的显存占用,用pytorch的profiler定位下,说不定是数据加载那边缓存没清。
我之前也遇到过一模一样的状况,最后发现是数据加载时把整个图像tensor都扔进了GPU做预处理,改成在CPU上用torchvision的transforms先做完再to(device)就稳了。你可以试试用nvidia-smi -l 1盯着看,或者用torch.profiler查一下是哪个模块在涨,ResNet50的BN层在迁移学习时经常因为梯度回传导致峰值显存突然翻倍。另外你混合精度开了,但优化器参数更新时记得把master weights也放CPU上,能省不少。实在不行可以把最后几层之外的参数设成requires_grad=False,只训分类头,显存直接掉一半。
试试把batchsize再砍到4,配合梯度累积到16步,我这么跑过72的都没炸。显存不够时先别动模型结构。
我之前也踩过这个坑,224x224加ResNet50按理说24G不该爆,你先查一下是不是DataLoader的pin_memory和num_workers开太高了,有时候数据预取会额外吃显存。另外别光看batchsize,你试试把torch.backends.cudnn.benchmark关掉,有些显卡上自动tuning会临时分配一大块显存。最直观的定位方法是用torch.profiler或者nvidia-smi盯着看,但更推荐你用pytorch的显存监控钩子,在每个block的forward后打印allocated memory,特别留意第一个和最后一个残差块,迁移学习时梯度反传在浅层占的空间常被忽略。还有个小技巧,把优化器换成AdamW并配合gradient checkpointing,虽然会慢一点,但显存能压到12G以内,实测比梯度累积稳得多。顺带问一句,你用的是torchvision的预训练权重吗?那个带BN的版本在batchsize小的时候会额外吃显存,建议换成不带BN的resnet50或者用蒸馏过的变体。
我之前也遇到过类似情况,当时查了半天发现是数据加载时候num_workers设成0了,导致CPU预处理成了瓶颈,显存反而被拖累。你可以试试把batchsize再压到4,同时用torchsummary或者pytorch的memory_stats函数看看每层tensor占用,我怀疑你可能是BN层的running_mean缓存没释放。另外检查下是不是把验证集的梯度也算了,有个detach没加的话显存会翻倍。
我之前也遇到过类似的情况,后来发现是验证集那段忘了包在torch.no_grad()里,梯度图一直攒着没释放,显存直接翻倍。你可以先检查下是不是这个原因,其次用nvidia-smi -l 1盯着看,或者试试torch.cuda.memory_summary(),能直接看到每个tensor占多少。另外ResNet50的BN层在迁移学习时如果batchsize太小(比如8),跑起来反而比大batch更吃显存,可以考虑把BN冻结掉只微调最后几层,显存能降不少。
我也遇到过类似问题,ResNet50跑224输入按理说8的batchsize不该这么吃显存,你检查下是不是数据加载时把图片转成RGB后没归一化,或者pin_memory和num_workers设置太高导致缓存堆积。可以试试用torchsummary或者pytorch的profiler看下每层显存分配,另外把batchsize再降到4配合梯度累积,或者把最后一个全连接层改成Global Average Pooling,能省不少显存。我之前用这个方式把20G压到12G左右,你可以参考下。
我之前也踩过这个坑,224的图加ResNet50按理说8的batch不该爆,你试试把dataloader的num_workers调高,有时候数据加载卡顿会让显存里堆积多余张量。另外用torchsummary或者pytorch的memory_snapshot看看,大概率是backward时候的中间激活值在作怪,可以开activation checkpointing,效果比梯度累积直接得多。还有个小细节,确认一下你用的是不是pretrained=True时默认的BN层在混合精度下的行为,偶尔会因为running_mean更新慢导致显存抖动。
我之前也踩过这个坑,224尺寸配ResNet50其实不算大,但你检查过dataloader的num_workers和pin_memory吗?有时候数据加载卡住会显存虚高。另外试试把batchsize再压到4,配合梯度累积到等效32,同时用torch.cuda.empty_cache()在每个epoch后清一下缓存。显存定位可以用nvidia-smi看实时占用,或者pytorch的torch.profiler,能直接看到每个op的显存分配。我上次发现是BatchNorm的动量缓存吃得多,关掉一些没用的层也许能省不少。
我之前也遇到过类似情况,224的图加amp还爆显存,大概率不是模型本身的问题,而是backward时中间激活值没释放。你可以试试在epoch结束加torch.cuda.empty_cache(),同时把dataloader的num_workers调高,有时候数据加载慢会变相拉长显存占用周期。另外强烈建议用pytorch的torch.profiler跑一下,能直接看到每个op的显存峰值,我之前就是靠它发现是BatchNorm的moving average在搞鬼,冻结BN后直接省了4G。