最近用 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