最近在跑一个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。
PyTorch训练时显存突然爆掉,但之前同样的代码没事,怎么回事?
全部回复
共 59 条把优化器从AdamW换SGD确实不该涨显存,但注意SGD+momentum如果开了nesterov,动量缓冲和梯度临时张量的生命周期可能跟AdamW不一样,检查下backward里是否有额外的中间变量被保留。另外gradient checkpointing只开forward的话,反向传播时重算图会临时多出激活峰值,建议用torch.cuda.memory_summary看下是哪个tensor卡在缓存里,大概率是碎片化,试试torch.cuda.empty_cache加max_split_size_mb=128能不能缓解。
优化器换SGD后梯度稀疏性变了,显存碎片化确实可能更严重,建议试试torch.cuda.empty_cache()或者调小PYTORCH_CUDA_ALLOC_CONF里的max_split_size_mb。
我最近也踩过类似的坑,换优化器后显存反而涨了,后来发现是SGD的momentum会给每个参数多存一份历史梯度,虽然单看不大,但叠加checkpointing和长序列的中间激活,峰值会卡在某个临界点上。另外建议你盯着第2个epoch结束时的显存曲线,如果接近18G,那第3轮OOM大概率是碎片化导致的,可以试试torch.cuda.empty_cache()加上减少dataloader的num_workers来缓解。不过你提到之前稳定运行过,那pytorch小版本更新导致缓存策略变化也很有可能,建议对比下两个版本的cuda缓存分配器实现。
我上次也碰到过,重启一下机器就好了,应该是显存碎片化没释放干净。
试试清一下缓存,torch.cuda.empty_cache()放每个epoch开头跑一下。
我遇到过类似的情况,而且也是换了优化器之后出现的。虽然SGD本身不占额外显存,但PyTorch的CUDA caching allocator在不同优化器下分配内存的模式会不一样,尤其是当SGD的动量缓冲区和AdamW的moment/variance在显存池里的排布方式不同,可能触发碎片化问题。你前两个epoch没事,第三个才爆,这多半不是模型或者数据的问题,而是显存池里累积了大量不连续的小块,等到需要一个大块时正好塞不下。gradient checkpointing只在forward里开确实会减少激活显存,但它在反向传播时其实会重新计算并释放,反而可能加剧碎片化。你可以试试在训练循环开始前设置torch.cuda.empty_cache(),或者用torch.cuda.memory_reserved()和max_memory_allocated()对比看看到底是分配还是利用率的问题。另外,如果你用的是新版本PyTorch,可以留意一下2.x里cudaMallocAsync的默认策略变了,有时候改成PYTORCH_CUDA_ALLOC_CONF=backend:cudaMallocAsync能明显缓解碎片。我之前就靠这个参数救回来了,建议你也试试,顺便把batch size临时降到16跑一两个epoch看看是不是还爆,能帮你定位是不是分配策略的锅。
试试把torch.cuda.empty_cache()加在每个epoch结尾,之前我也遇到过这种玄学OOM,多半是碎片化问题。
换SGD后lr和momentum调了吗,有时候优化器变了梯度分布也会影响显存分配策略。
SGD+momentum的动量缓冲在1.2亿参数下其实比AdamW多占显存,试试调小batch size或开max_split_size_mb。
把SGD+momentum换回来之后显存反而涨,这个方向其实挺反直觉的,但很可能是momentum缓冲区的锅——SGD的动量项会额外存一份和梯度同shape的float32状态,1.2亿参数算下来也不小了,而且如果原先AdamW用了混合精度,SGD反而可能没走AMP的路径。另外你提到gradient checkpointing只开了forward,但反向传播时激活值释放和重算的边界可能因为优化器变了而不一样,建议检查一下torch的版本是不是顺手升级过,新版缓存分配器对碎片更敏感,前两个epoch正常第三个爆掉很典型的碎片累积。可以试着在epoch开始前清一下cache,或者把max_split_size_mb调小试试。
SGD换过来之后如果没动学习率或momentum的schedule,loss曲线可能会变陡,某些中间层的梯度范数突然拉高,导致激活值内存临时性暴涨,这跟优化器本身省不省显存是两码事。另外gradient checkpointing只开forward的话,backward还是会存完整中间结果,建议你打印一下每个step的torch.cuda.max_memory_allocated,看看是不是第三个epoch某个batch异常。CUDA缓存碎片化确实会在连续跑几天后出现,可以试试在epoch开头手动torch.cuda.empty_cache,或者把PYTORCH_CUDA_ALLOC_CONF设为max_split_size_mb=128,我遇到过类似情况,多半是某个batch的序列长度实际没填满,pad后形状变了导致临时buffer分配不连续。
哎这个我太有同感了,之前跑BERT-large也遇到过一模一样的幽灵OOM,前两个epoch稳如老狗,第三个直接炸。你提到换SGD这个点我反而觉得是线索,虽然SGD本身省显存,但momentum项在PyTorch里实现时如果没和gradient checkpointing配合好,有时候会触发额外的中间变量保存,尤其是你只在forward里开checkpointing,backward的激活缓存反而可能因为优化器状态变化而重新分配。
另外你说到CUDA缓存分配,这个真的不能忽略,PyTorch的缓存分配器默认会预留一部分显存,但如果你代码里某个动态shape的操作(比如变长序列padding)在第三个epoch才触发某条路径,碎片化就会突然严重。我建议你在训练循环里加个torch.cuda.empty_cache()试一下,虽然治标不治本,但能确认是不是碎片问题。
还有个更阴间的可能:你是不是用了torch.compile或者新版本里默认启用了某些融合优化?我之前升级到2.1之后,同样代码显存占用莫名高了15%,后来发现是inductor对SGD的优化没跟上,反而把梯度checkpointing的recompute逻辑搞复杂了。建议你回滚到之前那个pytorch版本跑一次SGD,如果没问题那就是版本兼容性的锅。
对了,你监控过显存曲线吗?如果是第三个epoch才突然飙升,也可能和某个特定batch的数据有关,比如那条样本的序列实际有效长度特别长,导致某层计算图临时显存需求暴涨。可以试试按batch维度打印一下每个batch的max_memory_allocated,定位到具体是哪个step爆的。
最后想问下你数据加载时shuffle是不是每epoch都做?如果前两个epoch的数据分布和第三个差异很大,也可能影响dropout或者LayerNorm的统计量,间接导致中间激活大小波动——虽然理论上不该差这么多,但为了排查还是可以固定seed对比一下。
换SGD后动量缓冲会额外占显存,而且没开checkpoint的backward,第三个epoch碎片攒够了就炸。
SGD+momentum确实比AdamW省显存,但它对显存碎片更敏感,因为AdamW的优化器状态占大头反而让分配更规整。你这情况很像前两个epoch把缓存撑大后,第三个epoch刚好触发了碎片化导致的OOM,可以试试torch.cuda.empty_cache()或者设PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True。另外gradient checkpointing只开forward的话,backward重计算时显存峰值会往后挪,跟epoch数叠加起来也容易在中间炸。先跑个nvidia-smi盯着看是不是慢慢涨上去的,别急着改代码。
SGD动量会多存一份buffer,epoch间碎片累积到第3轮就炸了,试试torch.cuda.empty_cache()或调PYTORCH_CUDA_ALLOC_CONF。
SGD动量确实会多存一份状态,但更可能是碎片问题,试试 PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True。
SGD+momentum其实不一定比AdamW省显存,AdamW虽然有两个动量状态,但SGD如果momentum没设对或者dampening参数导致额外buffer,反而可能有意外开销。不过更可疑的是第3个epoch才炸,前面都正常,这很像内存碎片累积的问题,尤其你开了gradient checkpointing,反向重算时峰值会更高。可以试试设置PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,或者训练中途手动torch.cuda.empty_cache()看看能不能缓解。另外确认下换优化器后有没有无意中让某些中间变量保留了计算图,这个坑挺常见的。
SGD动量会多存一份动量buffer,显存反而比AdamW大,你试试调小batch或清下缓存。
SGD+momentum其实不一定更省显存,动量缓冲区虽然比AdamW的少,但如果你开了nesterov或者dampening,中间变量还是会多占一点。不过更可疑的是gradient checkpointing只在forward开——有些实现里backward重计算时如果没配合好,反而会在第三个epoch累积碎片。你可以试试在训练循环里每个epoch结束手动调一下torch.cuda.empty_cache(),或者设个PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,碎片问题最近几个版本确实挺烦人的。
SGD+momentum本身不省显存,momentum buffer要占一份参数量的状态,AdamW是两份,所以你是省了一份,但省出来的空间可能刚好被碎片吃掉。重点看第3个epoch才炸,前两个正常,大概率是某个epoch触发了不同长度的序列或者数据分布变了,导致checkpointing的segment划分跟之前不一样,峰值就顶上去了。建议先设PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True试试,碎片问题能缓解不少。另外确认下验证集是不是也在同一个进程里跑,有时候eval没包no_grad,到第三个epoch才轮到它。
SGD动量确实可能是个隐藏坑,尤其配合gradient checkpointing的时候。checkpointing本身会在反向时重新算forward,如果SGD的momentum buffer在第一次更新后驻留了额外状态,加上重计算时的临时激活叠加,显存峰值可能就顶上去了。你可以试试把momentum先设成0跑一轮,或者用torch.cuda.memory_summary看看碎片到底涨在哪。另外PyTorch 2.x的缓存分配器对长序列的碎片确实比1.13敏感,设个PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True有时候能救。