最近在跑一个Transformer分类任务,batch size=32,序列长度256,模型大概1.2亿参数。之前同样的代码在A100上稳定训练,显存占用大概18G左右。但这两天重跑,到第3个epoch时直接OOM,报错说“CUDA out of memory”,但前两个epoch都是正常的。
我排查了数据加载、梯度累积,都没发现问题。唯一变化是我把优化器从AdamW换成了SGD+momentum,但应该更省显存才对啊?另外,我用了gradient checkpointing,但只在forward里开了。
有没有可能和CUDA缓存分配策略有关?或者pytorch版本更新后内存碎片化更严重?求遇到过类似情况的大佬指点一下,真心不想降batch size。