最近在做一个小项目,想用ResNet50在自定义数据集上做迁移学习,训练图像分类模型。结果每次跑到第二个epoch,显存就飙到20G左右(我只有24G),然后直接OOM挂掉。数据集每张图是224x224,batchsize已经降到8了,还用了混合精度训练(torch.cuda.amp),也试了梯度累积,但感觉治标不治本。
我怀疑是不是模型本身太大,或者我的数据加载方式有问题?但看网上很多人同样配置都能跑,是不是我哪里写的不对?
想请教一下大家:除了换更小的模型,还有没有其他能稳定跑完训练的方法?或者有没有什么检查工具,能帮我定位显存占用到底在哪个层?谢谢各位大佬了。
用PyTorch跑ResNet50做迁移学习,显存总爆掉,求优化思路
全部回复
共 161 条我最近也碰到过类似问题,排查下来发现是数据加载时开了太多worker,每个worker都会缓存一些图像预处理结果,累积起来显存就炸了。你可以试试把DataLoader的num_workers设成0或者1,看看峰值有没有降下来。另外用torch.cuda.max_memory_allocated()打印每步的显存占用,能快速定位到是前向还是反向爆的,有时候是loss计算里多了一个不必要的tensor。
试试把num_workers设成0,有时候多线程加载会吃显存,另外检查下是不是梯度没清干净。
试试用torch.cuda.memory_summary()看下显存具体分配,我怀疑是你模型里某些层的中间变量没释放干净。另外检查下DataLoader的num_workers,设太高有时会占额外显存,可以降到2或4看看。还有个小技巧,把batch normalization层设成eval模式,只训练分类头,能省不少显存。
我和你遇到的情况差不多,后来发现是ResNet50的BN层在微调时会额外缓存中间激活值,试试冻结前几层或者改用梯度检查点(activation checkpointing),能省不少显存。另外用torch.cuda.memory_summary()能看到每块显存的分配情况,我之前就是靠这个发现数据加载时图片没释放干净。
你这情况我跑CIFAR-10的时候也遇到过,ResNet50的BN层在迁移学习里其实挺吃显存的,尤其如果冻结了前面层但保留了所有BN参数,反向传播时梯度还是会占不少空间。建议你先用torchsummary或者torchinfo看看模型每层参数和激活值的大小,我怀疑你数据加载时可能没做pin_memory=False的尝试,有时候多进程预加载反而会额外占用显存缓存。另外检查一下你的DataLoader里num_workers是不是设太高了,我一般设成2或者干脆0,不然CPU预处理积压也会让显存虚高。还有一个骚操作是试试把模型输出层的全连接层换成一个更小的分类头,比如先接一个Global Average Pooling再只接两层线性层,这样能省下不少参数。不过说实话,24G跑ResNet50+224x224+batchsize8按理说应该够,你是不是开了什么额外的可视化工具或者同时跑了多个进程?可以用nvidia-smi -l 1实时盯一下显存变化,看看到底是哪个瞬间炸的。
遇到过类似的情况,你试过用torch.no_grad()冻结BN层吗?有时候迁移学习里BN层的统计量会额外吃显存。另外可以看看dataloader的num_workers是不是设太高了,我之前调成0反而更稳。推荐用torch.cuda.memory_summary()打印一下,哪个层的allocated bytes最大一目了然。
看到你这个情况,我第一反应是检查下是不是ImageNet预训练权重把BN层里的running_mean和running_std也带过来了,如果冻结了某些层但又没设置eval模式,梯度会反向传播到那些统计量上,显存会莫名其妙多占不少。另外你试过用torch.utils.checkpoint做梯度检查点吗?虽然会牺牲一点速度,但对ResNet这种深层网络能省下不少中间激活的显存,特别是前几个stage的feature map特别吃空间。数据加载方面也可以排查下,比如worker数量开太多可能导致数据预处理时CPU把GPU的pin_memory buffer撑爆,我一般设成4或者直接用0看看对比。还有一个容易忽略的点是验证集推理时如果没加torch.no_grad(),那梯度图也会被保留下来,第二个epoch刚好是第一次val的话可能正好撞上。建议先用nvidia-smi配合gpustat实时盯一下,或者用torch.cuda.memory_summary()直接打印出每个tensor的分配情况,定位到是哪一层最吃显存再针对性优化。
我之前也踩过这个坑,ResNet50的BN层在迁移学习时特别吃显存,尤其是前几个epoch梯度回传的时候。你可以试试把batchsize再砍到4,然后配合梯度累积,虽然慢点但能稳住。另外检查下是不是数据加载时开了太多num_workers,有时候内存和显存会互相抢资源。最推荐用pytorch的torch.cuda.memory_summary()看一眼,能直接告诉你哪层分配最多,我上次就是发现是最后全连接层的梯度缓存爆的。
我之前也踩过这个坑,ResNet50的BN层在迁移学习时特别吃显存,尤其是batchsize小的时候,统计量更新反而更占资源。你可以试试把BN层冻住,只训练最后的全连接层,显存能降一半还多。另外检查一下是不是开了gradient checkpointing,这个对ResNet特别管用,用一点计算换显存,基本能稳住。数据加载那边建议用num_workers=4以上,有时候CPU瓶颈也会导致显存临时峰值飙升。最后推荐装个pytorch-memlab,能直接打印每层显存占用,比瞎猜强多了。
我之前也遇到过类似的情况,后来发现瓶颈不在模型本身,而是数据加载那一步。你试试把DataLoader的num_workers调大,同时用pin_memory=True,能明显缓解显存碎片化。另外可以用torch.cuda.memory_summary()看下具体是activation还是中间变量在吃显存,多半是反向传播时保存的梯度占了大头,可以试试打开checkpoint机制(torch.utils.checkpoint)来换时间省显存。
还有个小技巧,把batchsize再砍到4,但多开几个梯度累积步数,效果其实差不多,关键是别让显存触顶。我上次这么改完,24G卡跑ResNet50加DeiT都没再爆过。
试试把batchsize再砍到4,配合gradient checkpointing,显存能省一大半,ResNet50真没你想的那么吃显存。
我之前也遇到过类似情况,ResNet50做迁移学习其实挺吃显存的,尤其是如果冻结层没设对,梯度会全量计算。你可以试试把batchsize再压到4,然后配合梯度累积,另外检查一下是不是dataloader的num_workers开太多导致显存碎片化。工具的话,torch.cuda.memory_summary()能看到每个张量的占用,还有nvidia-smi的实时监控也能帮你判断是不是前向传播时缓存没释放。我后来发现我OOM是优化器状态太大,换了AdamW加权重衰减后好很多,你也可以看看是不是这里的问题。
我之前也踩过这个坑,ResNet50的BN层在迁移学习时特别吃显存,尤其是你如果没冻结前几层的话。建议先用nvidia-smi盯着看,或者用pytorch的torch.cuda.memory_summary()打印详细分配,基本能看出是激活值还是优化器状态占大头。另外试试把图片用albumentations做在线增强时顺便归一化,别存float32,直接转成uint8再喂,能省不少。还有个小技巧,backbone用resnet50换成resnet34的权重初始化,精度掉不了多少但显存能降一截。
我之前也踩过这个坑,ResNet50在224输入下其实不算特别吃显存,但如果你用了ImageNet预训练权重,第一层卷积和最后的全连接层在反向传播时梯度会特别占空间。建议你先用nvidia-smi监控一下,看是不是数据加载时候的pin_memory或者num_workers开太多导致CPU内存和GPU显存之间互相挤占,有时候datapipe里to(device)忘了放会隐式拷贝。另外可以试试把BatchNorm层冻结(bn.eval()或requires_grad=False),迁移学习里这招能省不少显存,因为BN的running_mean和running_var在反向传播时不需要梯度。还有个冷门技巧:用torch.utils.checkpoint对残差块做梯度检查点,虽然会慢一点但显存能从20G直接掉到10G以内。你要是确定想定位具体哪层爆的,可以用torch.profiler或者hook每个模块的allocated_bytes,我之前用这个发现瓶颈居然在最后的全连接层,因为自定义数据集类别少但输入特征维度大,那个矩阵乘法特别吃显存。另外检查一下是不是混合精度没生效,amp只在forward和loss计算时用,但如果你没把optimizer包进GradScaler,梯度还是fp32的,等于白开。最后实在不行就把输入尺寸降到192或者用ResNet34做teacher,效果差不了太多但显存压力小一半。
我之前也遇到过类似情况,当时查了半天发现是DataLoader的num_workers设太高,导致CPU和GPU之间的数据拷贝占了不少显存,你试试把workers降到2或者4看看。另外ResNet50的BatchNorm在迁移学习时最好冻结前几层,这样不光省显存,收敛还快,你可以用torchsummary或者nvidia-smi的按进程显存统计来定位是不是数据加载那块有问题,而不是模型本身。还有个小技巧,把输入图像先做一次随机裁剪到更小尺寸(比如192),跑通后再慢慢往上调,也能缓解不少压力。
试试把batchsize再砍到4甚至2,配合gradient checkpointing,显存能省一大截,不够就换AdamW加OneCycle。
我也遇到过类似的情况,当时排查下来发现是数据加载时pinned memory和num_workers开太多,反而把显存挤爆了,你可以试着把DataLoader的num_workers调低到2,顺便用torch.cuda.reset_peak_memory_stats()看下每个step的峰值,别光看总占用。另外ResNet50用amp的话,建议先确认下是不是某些自定义层没走cuda的autocast,导致梯度回传时精度不一致,内存会异常翻倍。还有个笨办法,把batchsize降到4,然后用梯度累积到等效16的batch,虽然慢点但至少能跑完,跑通了再逐步往上加。
我跑过类似的配置,24G卡带ResNet50按理说真不该爆,先查下是不是数据加载那边出了问题,比如num_workers开太大或者pin_memory导致额外显存占用。你可以用nvidia-smi -l 1实时盯着看,或者torch.cuda.memory_summary()看看哪一层分配最多,我猜大概率是BN层的统计或者反向传播的中间激活值在搞鬼。另外试试把优化器换成AdamW加fused版本,有时候默认的SGD动量缓存也吃显存。还有个小技巧,图像预处理别在GPU上做,全挪到CPU,能省出不少。
我之前也遇到过一模一样的情况,后来发现是数据加载时没做pin_memory和num_workers,导致CPU卡在瓶颈上,GPU反而被临时张量占满。你可以先用torch.cuda.memory_summary()看下分配峰值在哪,特别留意BatchNorm的running stats在迁移学习时是不是被设成了requires_grad=True。另外试试把ResNet50的stem和layer1冻结掉,只训后面的层,显存能省下快一半。
我之前也遇到过类似情况,后来发现是DataLoader的num_workers开太高,每个worker都会复制一份模型权重,显存直接翻倍。你把workers降到2或4试试,另外检查下是不是在验证集上也开了梯度,记得用torch.no_grad()包一下。还有个笨办法,用torch.cuda.memory_summary()看每个tensor的分配情况,能直接定位到是哪一层或者哪个操作在吃显存,比瞎猜强多了。