最近在做一个小项目,想用ResNet50在自定义数据集上做迁移学习,训练图像分类模型。结果每次跑到第二个epoch,显存就飙到20G左右(我只有24G),然后直接OOM挂掉。数据集每张图是224x224,batchsize已经降到8了,还用了混合精度训练(torch.cuda.amp),也试了梯度累积,但感觉治标不治本。
我怀疑是不是模型本身太大,或者我的数据加载方式有问题?但看网上很多人同样配置都能跑,是不是我哪里写的不对?
想请教一下大家:除了换更小的模型,还有没有其他能稳定跑完训练的方法?或者有没有什么检查工具,能帮我定位显存占用到底在哪个层?谢谢各位大佬了。
用PyTorch跑ResNet50做迁移学习,显存总爆掉,求优化思路
全部回复
共 170 条试试把ResNet50的BN层冻住,再用torchsummary看下每层显存,八成是激活值占大头。
试试把数据加载的num_workers调大,另外检查下是不是验证集没关梯度,用torch.no_grad包一下能省不少显存。
之前跑类似任务也踩过这个坑,ResNet50本身不算大,但BN层的显存占用在迁移学习里特别容易被忽略,尤其是你开了amp之后,模型参数和梯度虽然是半精度,但BN的running_mean和running_var还是全精度,而且每个batch都会更新,这部分开销很固定。建议先试试把torch.cuda.amp的grad_scaler换成自定义的,只在forward里用autocast,backward手动做梯度缩放,有时候默认实现会额外缓存一些中间变量。另外检查一下你的数据加载,是不是用了pin_memory=True但num_workers设得太大,导致CPU预处理和GPU拷贝的峰值重叠,可以试试num_workers=4甚至2,同时把dataset里的transform挪到GPU上做(比如用torchvision.transforms的GPU版本),减少CPU到GPU的传输缓冲。还有个思路是冻结前几层,比如只训练最后两个block,显存能砍掉一半,迁移学习本来就该这么干,全量微调在小数据集上反而容易过拟合。关于定位工具,torch.cuda.memory_summary()能看每个tensor的分配情况,但太罗嗦,推荐用pytorch_memlab的LineProfiler,能逐行显示哪行代码分配了显存,我之前用它发现是CrossEntropyLoss的label smoothing在内部生成了one-hot矩阵,换成了sparse版本直接解决。最后如果实在不行,试试把输入尺寸从224降到192,或者用F.relu(inplace=True)配合torch.utils.checkpoint,虽然慢一点但能稳稳跑完。
我跟你遇到过一模一样的情况,后来发现是数据加载时pin_memory和num_workers没调好,导致CPU来不及喂数据,GPU空转但显存碎片暴涨。你试试把DataLoader的num_workers设成4或8,pin_memory开true,顺便检查下transform里有没有无意中把图像复制成多个tensor。另外可以用torch.cuda.memory_summary()看下是模型参数还是激活值占大头,ResNet50按理说224输入+bs8用amp不该这么吃显存。
试试把resnet50的BN层冻结时顺便开cudnn.benchmark,有时候数据加载线程数拉满也会偷显存。
试试把batchsize再砍到4,还有检查下dataloader的num_workers是不是开太多,显存碎片也能卡死你。
试过冻结前面所有层只训fc吗?显存能省一大截,你这batchsize和amp按理说不该爆。
用nvidia-smi看下是不是数据加载线程占的显存,我之前就是num_workers开太多直接爆。
说到这个我太有同感了,之前调一个检测模型也是被OOM折磨到怀疑人生。你怀疑数据加载方式,其实可以先用torch.utils.data.DataLoader的pin_memory和num_workers调优试试,但我觉得更可能的问题是ResNet50的BatchNorm在迁移学习时开了training模式,导致反向传播的中间激活值全被保存了,尤其224x224输入下,每个特征图的显存开销比你想的大得多。建议你把BN层冻结掉,只训练最后的全连接层,或者用requires_grad=False冻结前面所有层,只解冻最后几个block,这样显存直接掉一半以上。另外你可以用torch.cuda.memory_summary()或者nvidia-smi dmon实时看每个时刻的显存分配,再配合torch.profiler看哪一层峰值最高,我之前就是这么定位到是某个残差块的中间张量爆了。还有个偏方,试试把输入尺寸临时降到196x196跑一个epoch,对比显存曲线,如果下降明显就说明确实是特征图尺寸的问题,那就得靠梯度检查点(torch.utils.checkpoint)来换空间了,虽然会慢点但能稳定跑完。最后确认下你的混合精度是不是真的生效了,有时候模型里某些算子不支持fp16会悄悄回退到fp32,导致显存没省下来,可以看看日志里有没有warning。
我之前也遇到过类似的情况,最后发现是数据加载的num_workers设太高,加上pin_memory=True,把CPU和GPU之间的传输缓冲也占了,改成4就稳了。你可以先用nvidia-smi看下显存是不是被数据加载占的,另外试试在dataloader里加个prefetch_factor=2,有时候比梯度累积管用。如果还想深挖,pytorch有个torch.cuda.memory._record_memory_history可以看每个tensor的分配点,但跑起来会更慢,建议先排除数据问题再上这个。
我之前也踩过这个坑,resnet50的bn层在迁移学习时如果没冻结,显存占用会比想象中高不少,你可以试试把前几层参数requires_grad设为False,只训练最后几层和fc,显存能降一截。另外用torch.utils.bottleneck或者nvidia-smi dmon实时看下显存曲线,我怀疑是dataloader的num_workers开太多导致内存碎片,而不是模型本身的问题,毕竟224x224加amp正常不该这么吃显存。你检查下是不是把验证集的梯度也保留了,val模式忘开torch.no_grad?我之前就这样白烧了十几个G。
试试把图像预处理改成随机裁剪加缩放,比直接resize省显存,我之前这么搞直接降了4G。
我跑过类似的配置,24G卡带ResNet50按理说很宽裕,你这情况八成不是模型本身的问题。建议先查一下dataloader的num_workers是不是设太高了,有时数据预处理会偷偷吃显存。另外可以用torch.cuda.memory_summary()看下分配细节,或者试试把batchsize调到4然后开梯度累积,看是不是峰值显存的问题。我之前遇到类似情况是optimizer的state占了不少,换AdamW加zero_grad(set_to_none=True)有时能省点。
你这情况我太熟了,之前做医学图像分类也卡在24G上,ResNet50其实不算大,问题多半不在模型本身。建议你先用nvidia-smi dmon或者pytorch的torch.profiler看一下,是不是数据加载时把整个batch的预处理图都留在计算图里了,比如ToTensor和Normalize没走GPU算子而是挂在CPU上等同步。另外,224x224的图在ResNet50里中间特征图其实挺占显存的,尤其第一个stage的channel是256,如果输入没做归一化或者用了太大的weight decay,也可能让梯度计算时额外保留中间激活。我有个笨办法:把batchsize压到4,开梯度累积步数设成4,但关键是配合torch.utils.checkpoint(梯度检查点)重算激活,显存能省一半还多。还有个小细节,你试试把模型的BN层冻住,只训练最后一层全连接,或者用lr层衰减,这样反向传播时不更新BN的running_mean,能省不少缓存。最后强烈建议用torch.cuda.memory_summary()打印一次,看是不是有个别tensor异常大,我之前就发现是dataloader的num_workers没设好,导致每个worker都复制了一份模型权重。
我之前也遇到过类似情况,后来发现是数据加载时NumPy转Tensor那步没释放内存,加上DataLoader的num_workers设太高,导致CPU和GPU之间传输堆积。建议先用torch.cuda.memory_summary()看下峰值在哪,顺便把pin_memory关掉试试。另外ResNet50的BN层在迁移学习时最好冻结前几层,这样显存能省不少,我那时候把batchsize调到16都没再爆过。
试试把图像预处理和增强挪到CPU上做,别让GPU管这块,能省不少显存。
试试把batchsize再砍到4,然后开gradient_checkpointing,能省不少显存。
用torch.cuda.memory_allocated每步打点看下,ResNet50真不是显存大户,八成是数据加载或优化器状态的问题。
说实话我第一反应也是怀疑你数据加载那边有问题,224x224的ResNet50按理说不该这么吃显存,除非你数据集里的图片没做resize就直接喂进去了,或者transform里带了什么奇怪的操作比如随机裁剪后没归一化,导致实际输入尺寸比224大不少。你可以先用torchsummary或者直接打印一下每个batch的shape,排除一下这个可能。
另外我强烈建议你试试显存监控工具,比如pytorch的torch.cuda.memory_summary(),或者用nvidia-smi配合watch命令看实时占用,这样能明确到底是模型参数、激活值还是优化器状态在占空间。我之前跑类似任务时发现是backward时的中间激活值爆了,后来开了gradient checkpointing,配合AMP和梯度累积,显存直接从20G降到11G左右,效果还挺明显的。
还有个容易被忽略的点,就是你迁移学习时如果冻结了前面所有层,只训练最后的fc层,那其实显存占用会小很多,反向传播只需要保存最后一层的梯度。你要是想保留特征提取层的微调,可以考虑只解冻最后几个block,这样既保留迁移学习的效果,又不会让激活值成倍增加。
检查工具方面,我推荐用PyTorch的torch.profiler或者captum,能按层输出显存分配情况,比单纯猜快多了。你试试把batchsize调成2跑一个epoch看会不会OOM,如果还爆,那就是代码问题,如果稳了,那就把问题定位到显存增长,再逐步排查是哪个操作在累积。
我之前也踩过这个坑,224的图配ResNet50其实不算大,问题多半出在backbone的梯度回传上。你可以试试把BN层冻住,只用更低的学习率微调最后几层,显存能省下不少。另外检查下数据加载那边,num_workers设太低会让CPU来不及喂数据,反而拖慢训练节奏,但不至于爆显存。真要定位的话,用torch.profiler看下每个op的内存分配,或者直接打印model.parameters()的requires_grad,大概率是某些层没冻结还在吃显存。
试试把验证集的准确率计算也挪到训练循环外,或者用torch.profiler看看是不是数据加载那块把显存占了。
说实话你这个配置和batchsize按理说真不该爆,我怀疑问题压根不在模型本身。ResNet50跑224输入,fp16下batch8的显存占用大概也就4-5G,你20G肯定是有东西在偷偷吃显存。建议先查一下是不是验证阶段或者数据增强里用了什么奇怪的操作,比如把整张图resize成多个尺度再堆batch,或者某些transforms在GPU上执行了。我遇到过一次坑是DataLoader的num_workers设成0,导致每个step都在主进程里做预处理,峰值显存直接翻倍,你检查下这块。
另外torch.cuda.amp的scaler要配合GradScaler用,而且loss.backward()之前确保没把loss的item()拿出来,不然会打断自动缩放。还有个思路是用torch.utils.checkpoint把ResNet的block换成梯度检查点,虽然会慢20%左右,但显存能降到原来的三分之一,对于单卡24G来说很划算。工具的话推荐pytorch的torch.cuda.memory_snapshot(),能输出每个tensor的分配栈,或者直接用nvidia-smi的PID跟踪看是不是有个显存泄漏的僵尸进程。
最后说个野路子,把batchsize减到4,同时把优化器换成LARS或者LAMB,这俩对大batch不敏感,但小batch下收敛也能稳住,配合梯度累积到等效batch32,显存反而比你现在更稳。我上次就是这么跑通EfficientNet-B5的,虽然慢点但至少不OOM。