最近在跑一个图像分类项目,用ResNet50做backbone,batch size设成32,结果跑不到10个epoch就OOM了。我检查了一下,数据加载用了DataLoader的num_workers=4,也试过pin_memory=True,但显存占用一直飙到24G+(我显卡是RTX 3090)。网上搜了一下,有人说用混合精度训练可以省显存,但我试了torch.cuda.amp还是偶尔报错。另外,我是不是应该把一些中间变量手动释放?还是说ResNet50本身就这么吃显存?有没有前辈指点一下调优方向,或者推荐一些常用的显存优化技巧?先谢过大家了。
楼主
4小时前
用PyTorch训练ResNet50,显存总爆满,是代码问题还是我哪里没设置对?
请 登录 后发表回复
全部回复
共 1 条
2楼
4小时前
3090 24G显存跑ResNet50 batch size 32应该没那么容易爆,除非你输入分辨率特别大或者加了别的开销。我怀疑问题可能出在DataLoader的num_workers上,4个worker不算多,但如果你数据预处理里做了大量transform(比如随机裁剪、色彩抖动这些),每个worker都会占用额外显存来缓存数据,加上pin_memory=True会让内存到显存的传输更积极,反而可能加剧OOM。混合精度训练的确能省一半显存,你说偶尔报错,是不是没加Gradient Scaling?AMP的GradScaler必须配合使用,否则梯度下溢会导致loss变成NaN,然后训练崩掉。另外ResNet50本身不算轻量,但正常训练下24G显存完全够,你可以试着用torch.utils.checkpoint(梯度检查点)来用时间换空间,关键层加个checkpoint能省不少显存。手动释放中间变量意义不大,PyTorch的autograd会自动管理计算图,真正吃显存的是前向传播时保存的中间激活值,你可以考虑减少batch size到16或者8先跑通,再逐步优化。如果数据集的图片尺寸不是必须224x224,缩小到160x160甚至128x128也能明显降低显存占用,分类精度损失通常很小。