最近用 Pytorch 训模型的过程中,发现总是训练几轮后,出现显存爆炸 out-of-memory 的问题,询问了 ChatGPT、查找了各种文档。。。
在此记录这次 debug 之旅,希望对有类似问题的小伙伴有一点点帮助。
问题描述:
训练过程中,网络结构做了一些调整,forward 函数增加了部分计算过程,突然发现 16G 显存不够用了。
用 nvidia-smi 观察显存变化,发现显存一直在有规律地增加,直到 out-of-memory。
解决思路:
尝试思路1:
计算 loss 的过程中是否使用了 item() 取值,比如:
train_loss += loss.item()
发现我不存在这个问题,因为 loss 是最后汇总计算的。
尝试思路2:
训练主程序中添加两行下面的代码,实测发现并没有用。
torch.backends.cudnn.enabled = True
torch.backends.cudnn.benchmark = True
这两行代码是干啥的?
大白话:设置为 True,意味着 cuDNN 会自动寻找最适合当前配置的高效算法,来获得最佳运行效率。这两行通常一起是哦那个
所以:
如果网络的输入数据在尺度或类型上变化不大,设置 torch.backends.cudnn.benchmark = True 可以增加运行效率;
如果网络的输入数据在每次迭代都变化,比如多尺度训练,会导致 cnDNN 每次都会去寻找一遍最优配置,这样反而会降低运行效率。
尝试思路3:
及时删除临时变量和清空显存的 cac