最近在跑一个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训练时显存突然爆掉,但之前同样的代码没事,怎么回事?
全部回复
共 43 条试试把torch.cuda.empty_cache()加在每个epoch结尾,这问题八成是碎片化,和优化器关系不大。
我之前也踩过类似的坑,换优化器后显存反而涨了,因为SGD的momentum会额外缓存一份梯度历史,虽然单看比AdamW少,但配合gradient checkpointing的释放策略可能会打乱显存复用节奏。另外你检查过cudnn.benchmark有没有被自动改吗?有时候输入长度不变但batch内padding变化会触发重新搜索算法,导致临时缓存暴涨。还有个小技巧,试试在epoch开始前手动调torch.cuda.empty_cache(),虽然治标不治本,但能确认是不是碎片化问题。如果还不行,建议对比一下PyTorch版本更新日志,最近几个版本在缓存分配器上改过好几次,降级到之前稳定版可能就恢复正常了。
我之前也踩过类似的坑,换优化器导致OOM真不一定是因为显存占用变高,而是内存碎片化加剧了。SGD+momentum虽然省了AdamW的额外状态,但PyTorch的缓存分配器对小块内存的复用策略更敏感,尤其是你开了gradient checkpointing,它会在forward里反复释放和重新分配激活值,碎片一多,第三epoch刚好撞上某个峰值就爆了。你可以试试在训练循环里加torch.cuda.empty_cache(),虽然治标不治本,但能缓解峰值压力。另一个思路是检查一下PyTorch版本,我遇到过2.1升到2.2后,同样的代码显存占用涨了10%的情况,官方后来也修了几个缓存分配相关的bug。另外,你确认下SGD的momentum实现是不是真的比AdamW省,我印象里momentum的缓存也是按参数shape分配的,如果模型里有大量小tensor,碎片化会更严重。最后,建议用torch.cuda.memory_summary()看下OOM瞬间的分配详情,对比一下前两个epoch的峰值,大概率能看到碎片化的痕迹。
这问题我上周刚踩过类似的坑,也是莫名OOM但改小batch就没事。你换SGD这个点其实很关键,虽然SGD本身省显存,但momentum项会额外保存一份梯度历史,加上你开了gradient checkpointing,它在反向传播时重新计算前向的显存峰值和常规训练完全不同,优化器切换后这个峰值位置可能就变了。我怀疑你前两个epoch没事是因为缓存还没吃满,到第三轮正好撞上碎片化分配失败,A100的40G看着大但连续显存被分割后反而容易触发cudaMalloc失败。你可以试试在训练循环里加两句torch.cuda.empty_cache()和torch.cuda.synchronize(),在epoch结束位置调用,看能不能缓解,我这么做后峰值降了差不多3G。另外强烈建议你监控一下第2.5个epoch时的reserved memory和allocated memory差值,如果reserved远大于allocated就是典型的缓存碎片问题,用PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True这个环境变量重启试试,我换了这个配置后基本告别这种玄学OOM了。还有个小细节,gradient checkpointing配SGD时,如果用了momentum=0.9,建议把checkpoint的granularity调粗一点,比如按transformer层为单位而不是每个子层,能减少很多临时张量。
1.2亿参数配AdamW本来就会在优化器状态上吃掉两倍于SGD的显存,你换回SGD后显存占用应该明显下降才对,但既然前两个epoch正常第三个才爆,大概率不是模型本身的问题,更像显存碎片化在第三个epoch触发了某个峰值分配。建议你试试在训练循环里定期调torch.cuda.empty_cache(),或者干脆把环境变量PYTORCH_CUDA_ALLOC_CONF设为max_split_size_mb=128,这招对碎片化引起的偶发OOM特别管用。另外你确定gradient checkpointing只在forward里开对了吗,SGD的momentum如果用了带buffer的实现,反向传播时可能会额外保留一些中间变量,这点可以重点查一下。最后,如果pytorch是最近更新的,可以对比下之前版本的显存分配行为,我遇到过2.1到2.2后同样代码峰值内存涨了5%的情况。
我之前也踩过类似的坑,换优化器之后显存曲线变得特别诡异,SGD虽然省了AdamW的动量缓存,但有时候梯度裁剪或者loss缩放会额外申请临时张量,正好撞上碎片化就炸了。你可以试试在dataloader里加个non_blocking=True,或者干脆把batch size临时降到28跑一个epoch对比下,能快速定位是不是缓存分配问题。另外PyTorch 2.x的缓存分配器确实改过策略,如果之前是1.x升上来的话,建议设一下PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,这招对我上次解决类似OOM挺管用。
跟优化器关系不大,八成是显存碎片化,试试torch.cuda.empty_cache或者调低batch size看是否稳定复现。
感觉像是优化器切换后显存碎片化变严重了,试试设个PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True看看。
我遇到过类似情况,换SGD后梯度缓冲变了,检查下是不是checkpointing的中间激活没释放干净。
看到你说换SGD反而爆显存,我第一反应也是挺反直觉的,但仔细想想还真有可能是优化器状态的问题。AdamW虽然有两个动量项,但PyTorch实现里对SGD+momentum的momentum buffer有时会额外分配一块和梯度同shape的连续内存,特别是在你开了gradient checkpointing的情况下,反向传播时梯度释放和重算的节奏变了,反而容易让CUDA缓存碎片化更严重。我之前遇到过类似情况,把SGD的momentum设成0.9时显存峰值比AdamW还高,后来用torch.cuda.empty_cache()在每个epoch结尾手动清一下,再把dataloader的num_workers调低点,居然就稳了。另外你提到PyTorch版本更新,我怀疑是不是新版的CUDA caching allocator对分段内存的复用策略改了,之前爆掉后重启进程,用PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128跑,就没再出现过中途OOM,你可以试试这个环境变量。还有个细节,gradient checkpointing如果你只在forward里开,而没在backward里配合使用torch.utils.checkpoint.checkpoint包裹整个block,实际上重算的中间激活可能没被正确释放,建议检查一下是否每个transformer层都用了checkpoint。最后想问你一下,前两个epoch正常,第三个才爆,是不是某个特定batch的长度或mask特别不规则?我之前遇到过序列里有异常长的样本导致临时张量峰值飙升,虽然你说了长度固定,但保不齐padding逻辑有隐藏bug。
我之前也碰到过类似灵异事件,后来发现是PyTorch版本小更新后,CUDA caching allocator的默认行为变了,碎片化更严重。你试试在代码开头加个torch.cuda.empty_cache(),或者把PYTORCH_CUDA_ALLOC_CONF设为max_split_size_mb=128看看。另外SGD+momentum虽然省显存,但如果你开了gradient checkpointing且没配合offload,某些中间变量释放时机可能和AdamW不一样,导致峰值反而更高。建议用torch.cuda.memory_summary()打印一下第3个epoch开始时的内存快照,对比前两个epoch看哪里爆的。
我遇到过类似的情况,而且也是换优化器之后出现的。SGD+momentum虽然理论上省显存,但它的momentum buffer在PyTorch里可能和AdamW的exp_avg/exp_avg_sq分配方式不一样,尤其是你用gradient checkpointing的时候,前向重算会临时持有更多中间变量,这时候SGD的state反而可能和checkpoint的临时显存撞在一起,导致峰值比AdamW还高。你可以试试在第三个epoch开始前手动调一下torch.cuda.empty_cache(),或者干脆把gradient checkpointing改成在backward里也生效,虽然慢点但能压峰值。另外你说到CUDA缓存碎片化,这个确实存在,特别是如果前两个epoch里某些tensor的shape是动态的(比如attention mask长度变化),PyTorch的caching allocator会留下很多小块碎片,到第三个epoch某个大tensor申请连续内存时就直接OOM了。建议你跑的时候把PYTORCH_CUDA_ALLOC_CONF设成max_split_size_mb:128或者expandable_segments:True,后者对碎片化效果很明显。还有个小细节,SGD的momentum实现是in-place更新,如果你数据加载那边用了pin_memory并且num_workers>0,有时候worker进程的预取缓存会和主进程的显存分配产生微妙的时序冲突,我之前就是加了prefetch_factor=2之后突然OOM,去掉就好了。你可以看看这两个epoch的迭代时间有没有变慢,如果变慢了大概率就是碎片化在累积。
我前几天也踩过类似的坑,换优化器看着显存公式是省了,但SGD的momentum缓冲在PyTorch里有时会和gradient checkpointing的激活重计算机制打架,尤其是在序列长度这种维度上,缓存块分配策略会突然变得很激进。你试试把torch.cuda.empty_cache()放在每个epoch结束的地方,或者直接用torch.cuda.set_per_process_memory_fraction限制一下峰值,看能不能把OOM往后推。另外我怀疑跟PyTorch版本也有关系,2.1之后CUDA caching allocator对碎片化的处理改过一版,你如果最近升级过,可以回退到之前跑通的版本对比下。还有个思路是检查下SGD的momentum参数,如果设了0.9,它其实会额外维护一份和模型参数同形状的动量缓冲,1.2亿参数算下来也有快500MB,虽然不多但加上碎片化可能就刚好卡在临界点。最绝的是我上次遇到类似情况,最后发现是DataLoader的num_workers在第三个epoch时某个worker崩了,导致显存没被正确释放,你最好也盯着nvidia-smi看下是不是有残留进程占着显存。
我之前也碰到过类似玄学,换SGD按理说省显存,但优化器状态和AdamW差异挺大的,momentum的buffer有时候反而会挤占碎片空间。gradient checkpointing只开forward的话,backward还是会额外存激活值,试试把recompute也用在backward上?另外建议监控一下第三步的显存分配曲线,很可能就是CUDA缓存碎片化到临界点了,可以试试在epoch开头清一下缓存或者调大PYTORCH_CUDA_ALLOC_CONF的max_split_size。
说实话看到你换SGD这个细节我第一反应也是不该更费显存,但后来转念一想,问题可能出在优化器状态和梯度累积的交互上。SGD+momentum的momentum buffer虽然比AdamW少一组参数,但PyTorch在混合精度或某些版本下会为每个参数额外分配临时张量,尤其是你开了gradient checkpointing,反向传播时重算的中间激活和优化器更新时的临时变量容易撞在一起,导致峰值显存比理论值高出一截。
另外你说的CUDA缓存分配策略确实值得深挖,PyTorch的缓存分配器在显存紧张时不会立即释放碎片,而是保留在池子里,如果前两个epoch刚好把显存“喂”到某个临界点,第三个epoch某个突发张量申请就会触发OOM,哪怕总占用没到极限。我遇到过类似情况,解决方法是加一行torch.cuda.empty_cache()在每个epoch结束手动清一次,或者干脆把PYTORCH_CUDA_ALLOC_CONF设为max_split_size_mb=128,强制限制碎片大小。
还有个小坑,你确认换优化器后学习率和momentum参数都调对了吗?SGD如果momentum设得偏大,某些实现会额外保存历史梯度副本,显存占用比默认值高不少。我上次就是被这个坑了,最后发现是momentum=0.9时PyTorch内部为每个参数多开了一个float32备份。建议你用nvidia-smi盯着看前两个epoch每个step的峰值变化,如果第三个epoch突然跳涨,大概率是某个batch的序列长度或mask模式触发了非均匀计算路径,跟优化器关系反而不大。
试试把torch.cuda.empty_cache()加到每个epoch开头,SGD的momentum buffer偶尔会触发更碎的内存分配。
我之前也踩过类似的坑,换优化器后显存反而异常,大概率是SGD的momentum缓存和AdamW的exp_avg/exp_avg_sq在内存布局上不一样,碎片化更严重。你试试在训练循环里加个torch.cuda.empty_cache(),或者把max_split_size_mb调小一点,能缓解不少。另外gradient checkpointing只在forward开的话,backward时中间激活还是会重建,建议确认下是否覆盖到了整个模型,有时候某些层没包进去也会导致峰值暴涨。
我之前也踩过类似的坑,换优化器之后显存曲线完全变了。SGD+momentum虽然本身不存一阶二阶动量,但它的梯度更新方式和AdamW不一样,可能导致中间激活值的生命周期变长,尤其是配合gradient checkpointing的时候,checkpoint的触发粒度可能会因为优化器改变而变得不均衡,我猜你第三个epoch的某个batch刚好触发了碎片化最严重的临界点。
另外你提到pytorch版本更新,这个很关键,我遇到过2.1之后默认的缓存分配器行为改了,虽然总显存占用看起来差不多,但连续训练几个epoch后,旧的张量释放和新的分配反复交错,很容易在某个时刻撞上碎片瓶颈,A100上尤其明显,因为显存大但分配粒度也更敏感。你可以试试在训练循环里手动调一下torch.cuda.empty_cache(),或者用cudaMallocAsync那个异步分配策略,有时候能立竿见影。
还有个思路,既然前两个epoch正常,第三个才爆,可以怀疑是不是和梯度积累的步数或者学习率调度有关,SGD的momentum会累积历史梯度,如果某个batch的梯度范数特别大,可能会让优化器状态临时膨胀,虽然理论上不占显存,但实际实现里某些cuDNN或者cuBLAS的workspace会跟着变。我建议你在第三个epoch开始前打印一下每层的梯度shape和当前缓存占用,把分配器内部的block信息dump出来看,比盲猜快多了。
我之前也踩过类似的坑,所以看到你这个帖子特别有共鸣。你提到换优化器后显存反而炸,我怀疑问题不一定在优化器本身,而是SGD+momentum的momentum buffer在PyTorch里可能和梯度checkpointing的释放机制有冲突。我遇到过gradient checkpointing只在forward里开,但backward时中间激活被重新计算后,显存碎片化会突然加剧,尤其到第三个epoch正好是缓存分配器开始复用碎片的时候。另外你查过CUDA的PYTORCH_CUDA_ALLOC_CONF设置吗?如果环境变量里没配garbage_collection_threshold,默认情况下缓存块不会及时合并,跑几个epoch后碎片积累到临界点就会OOM。我上次就是靠把max_split_size_mb调小到4,再加上expandable_segments:True解决的,你可以试试。还有个疑问,你用的是不是最新版PyTorch?2.1之后分配器改过策略,对长序列任务确实更敏感,我同事降回2.0.1就稳定了。不过话说回来,前两个epoch正常第三个才爆,更像是某个隐藏的显存泄漏,建议你用nvidia-smi监控每个step的峰值,看看是不是在第2个epoch结束时有未释放的临时tensor。
我遇到过类似的坑,SGD+momentum虽然省显存,但它的momentum buffer在PyTorch里是单独分配的,如果之前用AdamW时那些状态没释放干净,换优化器后可能反而触发碎片化。你可以试试在换优化器后加torch.cuda.empty_cache(),或者把PYTORCH_CUDA_ALLOC_CONF设成max_split_size_mb:128看看。另外gradient checkpointing只在forward里开的话,backward时激活值还是会重新计算,如果显存峰值刚好卡在那个临界点,前两个epoch侥幸过了,第三个epoch可能因为缓存碎片就爆了。
我之前也踩过类似的坑,换SGD后显存反而涨了,因为momentum项会在反向时额外缓存梯度历史,而AdamW的exp_avg其实复用了一部分显存。另外gradient checkpointing只在forward里开的话,反向传播时激活值还是会重新计算,如果恰好某个batch的输入导致激活峰值变大,就会OOM。建议你查一下是不是某个batch的序列长度分布变了,哪怕同是256,padding多的样本显存波动也很大。还有个思路,跑之前先清一下CUDA缓存,或者设置PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,能缓解碎片问题。