最近在跑一个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。
楼主
13天前
PyTorch训练时显存突然爆掉,但之前同样的代码没事,怎么回事?
请 登录 后发表回复
全部回复
共 43 条
2楼
1天前
把优化器从AdamW换SGD确实不该涨显存,但注意SGD+momentum如果开了nesterov,动量缓冲和梯度临时张量的生命周期可能跟AdamW不一样,检查下backward里是否有额外的中间变量被保留。另外gradient checkpointing只开forward的话,反向传播时重算图会临时多出激活峰值,建议用torch.cuda.memory_summary看下是哪个tensor卡在缓存里,大概率是碎片化,试试torch.cuda.empty_cache加max_split_size_mb=128能不能缓解。
3楼
1天前
优化器换SGD后梯度稀疏性变了,显存碎片化确实可能更严重,建议试试torch.cuda.empty_cache()或者调小PYTORCH_CUDA_ALLOC_CONF里的max_split_size_mb。
4楼
22小时前
我最近也踩过类似的坑,换优化器后显存反而涨了,后来发现是SGD的momentum会给每个参数多存一份历史梯度,虽然单看不大,但叠加checkpointing和长序列的中间激活,峰值会卡在某个临界点上。另外建议你盯着第2个epoch结束时的显存曲线,如果接近18G,那第3轮OOM大概率是碎片化导致的,可以试试torch.cuda.empty_cache()加上减少dataloader的num_workers来缓解。不过你提到之前稳定运行过,那pytorch小版本更新导致缓存策略变化也很有可能,建议对比下两个版本的cuda缓存分配器实现。