最近在跑一个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 条我前两天也碰到过类似情况,后来发现是PyTorch版本从2.1升到2.2后,CUDA caching allocator对碎片处理方式变了,尤其是在用了gradient checkpointing这种分段释放显存的场景下特别明显。你可以试试设置PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,或者把max_split_size_mb调小一点,看看能不能缓解。另外SGD+momentum虽然参数少,但如果你开了momentum的Nesterov,中间变量存储其实比AdamW更占临时显存,特别是长序列下,可以对比下两个优化器在峰值显存上的实际差异。
这问题我调了两天才发现是环境变量的问题,老版本代码在新环境跑就是容易出这种隐蔽坑。
我之前也踩过类似的坑,换优化器看着显存占用差不多,但SGD的momentum缓冲区在反向传播时可能会跟checkpointing的临时张量打架,导致峰值比AdamW还高。你可以试试在optimizer.zero_grad()后面加torch.cuda.empty_cache(),或者把checkpointing的use_reentrant改成False,这俩对碎片化影响挺大的。另外确认下是不是只有第3个epoch才爆,如果跟数据顺序有关,可能是某个batch的sequence长度异常,虽然你设了256但pad后实际token数可能有波动。还有个小建议,直接看下报错前的CUDA memory summary,allocated和reserved的差值如果特别大就是碎片问题,可以调PYTORCH_CUDA_ALLOC_CONF的max_split_size_mb参数。
SGD+momentum本身确实不增加显存,但你是不是把momentum设得比较大?有些实现会在优化器状态里额外存一份动量副本,碰上checkpointing的临时张量释放不及时,峰值反而可能更高。我之前遇到过类似情况,最后发现是PyTorch版本升级后,默认的CUDA缓存分配策略变了,碎片化导致无法复用显存,试试在训练循环里加torch.cuda.empty_cache()或者调低PYTORCH_CUDA_ALLOC_CONF的max_split_size_mb看看。
我之前也踩过类似的坑,换优化器导致OOM真不一定是因为显存占用变高,很可能是碎片化加剧了。SGD+momentum虽然理论上省显存,但PyTorch的缓存分配器对内存块的管理方式和AdamW不同,加上gradient checkpointing在backward时释放和重新申请显存的节奏变了,容易让空闲块变得很碎,结果就是明明总量够用,但找不到连续的大块。你试试在训练循环里加个torch.cuda.empty_cache(),或者干脆设一下PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128,这个对碎片化很有效。另外,前两个epoch正常第三个才爆,我怀疑是不是某个中间变量在特定step下尺寸突然变大了,比如序列长度如果被padding到固定256,但实际batch内数据长度分布不均匀,偶尔有个超长的样本就会触发峰值。你可以把OOM前的日志打出来,看看是哪个op爆的,或者用torch.autograd.detect_anomaly()跑一下。至于PyTorch版本更新,之前有人遇到过cudnn benchmark开启后,不同shape的卷积会额外缓存workspace,你可以试试把torch.backends.cudnn.benchmark设为False对比下。最后还是建议你盯一下nvidia-smi的显存曲线,如果确实是稳步上升而不是突然跳变,那大概率是缓存没释放,而不是代码逻辑问题。
换优化器这个点其实挺可疑的,SGD+momentum虽然理论上省显存,但如果你没动momentum的buffer或者没用momentum的dampening,某些实现会在backward时额外缓存中间张量。我上次也碰到过类似情况,把AdamW换回SGD后显存没降反升,后来发现是PyTorch的SGD在cuda graph捕获模式下会多保留一层计算图。另外你说gradient checkpointing只在forward里开,这其实是个大坑,checkpointing的存储策略是跟autograd绑定的,如果backward里没有显式处理,某些层的前向激活还是会被完整保留,尤其是Transformer里多头注意力那块。关于CUDA缓存,PyTorch的caching allocator确实会有碎片化问题,尤其在连续跑了两个epoch后,某些临时tensor释放不彻底,再叠加新的大块显存请求就容易OOM。我建议你在每轮epoch结束后加个torch.cuda.empty_cache()试试,但注意别放在训练循环里频繁调用,否则会拖慢速度。还有一个更阴间的可能,就是你的A100是不是被别人占了显存?nvidia-smi看看有没有其他进程,我之前就被同机的人坑过一次,第3个epoch刚好撞上他跑大模型。最后问一下,你PyTorch版本是不是最近升级过?2.1之后对SGD的momentum buffer做了些变化,有些人反馈过碎显存分配更频繁了,我后来直接锁版本才稳下来。
之前我也遇到过类似玄学OOM,换优化器后显存反而涨了,后来发现是SGD的momentum buffer在梯度checkpointing下没被正确释放,导致峰值内存比AdamW还高。你可以试试在backward后手动清一下optimizer的state,或者把checkpointing的use_reentrant改成False看看。另外PyTorch 2.x的缓存分配器确实有碎片化问题,但一般不会到第3个epoch才爆,更像是某个特定batch的激活值触发了峰值,建议用torch.cuda.memory_summary()看一下具体是哪个分配点。
SGD+momentum虽然理论上省显存,但动量缓冲也会占一块,而且优化器切换后如果学习率或momentum没调好,可能导致某些层的梯度突然变大,中间激活值跟着涨。建议先用nvidia-smi盯着看是不是第3个epoch某个batch显存峰值异常,而不是平均占用。另外gradient checkpointing只在forward开的话,backward的激活缓存其实还是会累积,可以试试把checkpointing包住整个loss计算。CUDA缓存碎片化确实可能,尤其是你换优化器后分配模式变了,可以设PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128或者用torch.cuda.empty_cache()在epoch间清一下试试。
这问题我上个月也踩过,简直一模一样。你换SGD这事儿挺关键,虽然理论上省显存,但SGD+momentum的梯度行为跟AdamW差很多,如果没同步调整学习率或者momentum参数,可能会导致中间激活值的分布变得更极端,反而让某些层的临时缓冲区撑爆了。另外gradient checkpointing只在forward里开,backward的时候其实还是会重新计算一些中间量,这个开销在换优化器后可能被放大了。
我猜更大概率还是碎片化的问题,A100上40G显存跑18G看着余量很大,但PyTorch的缓存分配器是块状预分配的,如果前两个epoch把不同size的tensor都“染指”过一遍,第三个epoch突然出现一个稍大的临时tensor(比如某个batch的loss计算方式变了),就可能触发重新向CUDA申请内存,而这时候显存已经被碎片占满,直接OOM。你可以试试在训练循环里手动调torch.cuda.empty_cache(),或者更狠一点,用环境变量PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True,这个在2.1版本后特别管用。
另外提个反直觉的点,你确定换优化器后代码里没有隐式改变batch的构造方式吗?我之前遇到过因为随机种子变了,导致某个batch序列长度实际超过256(如果数据有padding逻辑),这种偶发情况前两个epoch碰巧没触发,第三个epoch就炸了。建议在报错前加个print,看看是不是特定batch出的问题。
我之前也踩过类似的坑,尤其是优化器从AdamW换到SGD的时候,反而更容易触发显存峰值。SGD虽然理论上省显存,但momentum的动量缓冲在反向传播时如果和梯度checkpointing的临时释放机制撞在一起,可能会出现某个step临时分配大量显存的情况,尤其是在第三个epoch正好遇到长尾分布的数据时,瞬时tensor大小可能比前两个epoch大不少。
另外,你提到CUDA缓存分配策略,这个方向我觉得很值得查。PyTorch的缓存分配器会保留已释放的块,但如果碎片化严重,新的大块申请会触发额外预留,导致“看起来”占用翻倍。我之前遇到过类似情况,解决办法是设置PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128,或者干脆在训练循环开始前手动torch.cuda.empty_cache()一下,虽然治标不治本,但能撑过几个epoch。
还有个可能性是数据加载的worker进程在第三个epoch时恰好有少量样本padding或mask长度变化,导致batch内实际计算量波动,这个在Transformer里很隐蔽。建议你加个日志监控每个step的torch.cuda.memory_allocated()和torch.cuda.memory_reserved(),看看是不是在某个特定step突然跳高。
至于pytorch版本更新,我确实见过2.1之后对checkpointing的显存回收逻辑有改动,如果方便的话可以回退到之前稳定版本跑一下对比。另外,你试试把torch.cuda.set_per_process_memory_fraction设个上限,比如0.95,这样OOM会变成显存溢出提前报警,方便定位是哪个tensor的问题,而不是直接崩掉。
试试把torch.cuda.empty_cache()放到每个epoch结尾,我上次也是第三个epoch炸的,清完缓存就好了。
SGD动量项会额外存历史梯度,你这模型1.2亿参数,算下来也有几百MB,但真不是主因,估计还是显存碎片化。
我之前也踩过类似的坑,SGD+momentum虽然理论上省显存,但实际动量缓冲区的分配时机和AdamW不一样,PyTorch可能不会立刻释放之前的显存碎片,尤其是你开了gradient checkpointing后,反向传播时中间激活的重计算会加剧碎片化。你可以试试在训练循环里加torch.cuda.empty_cache(),或者把环境变量PYTORCH_CUDA_ALLOC_CONF设为max_split_size_mb=128,强制减少碎片。另外,如果PyTorch版本是最近才升的,建议回退到之前稳定的小版本,这问题跟CUDA缓存策略关系挺大的。
换SGD按理说省显存,但注意momentum的buffer也是要占显存的,虽然比AdamW少,可如果之前AdamW开了fused或者用了bitsandbytes优化,那差距可能没你想的大。另外gradient checkpointing如果在forward里手动调用,但backward时中间变量没被正确释放,也可能在第3个epoch积累碎片。建议你试试在训练循环里加torch.cuda.empty_cache(),或者干脆把batch size降到16看还爆不爆,能快速定位是不是碎片问题。至于pytorch版本,我之前遇到过2.1升2.3后显存分配策略变了,同一个代码多占1-2G,可以试试回退版本。
我之前也踩过类似的坑,换优化器后显存峰值反而变了,重点检查一下SGD的momentum是不是在backward里额外存了中间梯度,而AdamW的exp_avg本身是复用显存的。另外gradient checkpointing只开forward的话,backward还是要重新计算激活,等于没省多少,试着把torch.cuda.empty_cache()加到每个epoch结尾看看,有时候碎片化就是前两个epoch缓存没释放,到第三个才炸。还有个思路,查下是不是PyTorch版本更新后,cudnn benchmark的autotune在跑不同形状时偷偷缓存了多份workspace,这个特容易在长序列上触发。
换SGD之后按理说显存只会更宽松,但注意momentum会额外存一份动量,如果SGD实现里带了nesterov或者你忘了关weight decay,某些版本下显存占用反而会上升。另外gradient checkpointing只在forward开的话,backward时还是会存激活值,建议检查一下torch.utils.checkpoint的用法。CUDA缓存碎片化确实有可能,特别是连续跑多个实验后,可以试试在训练开头加torch.cuda.empty_cache(),或者设置PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb:128,我之前遇到过类似问题,调这个参数直接解决了。
换SGD还真不一定更省显存,momentum会额外存一份动量缓冲,如果优化器实现有额外中间变量,峰值反而可能涨。你gradient checkpointing只开forward的话,backward还是会保留激活值,这个和优化器变更叠加在一起,显存曲线可能就变了。建议盯着nvidia-smi看每个step的峰值,别只看稳定占用,OOM往往就是某一步突然尖峰。另外如果最近升级过pytorch或cuda版本,碎片化确实可能更严重,试试在训练循环里加torch.cuda.empty_cache()或者调大PYTORCH_CUDA_ALLOC_CONF的garbage_collection_threshold,能缓解不少。
SGD的momentum buffer在1.2亿参数下可比AdamW省不少,但你是不是忘了把gradient checkpointing的recompute改成每次step了?
换个环境试试,我上次就是torch版本小更新后显存碎片化严重,加个PYTORCH_CUDA_ALLOC_CONF=expandable_segments:True就好了。
换SGD后batch norm的running stats计算方式变了,显存峰值可能出现在backward,试试torch.cuda.empty_cache()手动清下碎片。
优化器换了但梯度裁剪没同步调吧?SGD的梯度分布和AdamW差很多,检查下是不是梯度爆炸导致中间变量缓存异常。
优化器换了之后loss曲线不一样,可能触发checkpoint的重新计算路径不同,试试关掉gradient checkpointing跑一个epoch对比下。
换SGD之后按理说显存只会更宽松,但注意momentum会额外保存一份和梯度同形状的动量缓冲,如果你用了heavy的weight decay或者momentum=0.9,这部分开销在1.2亿参数下也有几百MB,不至于直接OOM。更可疑的是gradient checkpointing只在forward里开,backward时如果PyTorch版本更新了,自动求导的中间激活释放策略可能有变化,特别是和SGD的step逻辑混在一起时容易产生碎片。建议你下次跑的时候nvidia-smi盯着看,如果显存是慢慢涨上去的而不是瞬间爆,那大概率是缓存碎片化,可以试试在epoch之间调torch.cuda.empty_cache(),或者把PYTORCH_CUDA_ALLOC_CONF设成max_split_size_mb:128。另外确认下是不是同一个conda环境,有时候依赖版本悄悄变了也会这样。
我最近也踩过类似的坑,尤其是换优化器之后显存反而异常,这事真不一定是省显存的逻辑能解释的。SGD+momentum虽然本身不存AdamW的一阶二阶动量,但PyTorch的CUDA caching allocator在频繁分配释放时容易产生碎片,特别是你开了gradient checkpointing,它会在forward里反复执行自定义的autograd Function,导致缓存块被切得很碎,后面一个稍大的tensor申请就直接爆掉。你可以试试在epoch开始前调用torch.cuda.empty_cache(),或者设置PYTORCH_CUDA_ALLOC_CONF=max_split_size_mb=128,强制分配器合并小块,我上次就是靠这个救回来的。另外你说的版本更新,如果最近升级过PyTorch,有些算子的临时buffer大小确实变了,比如attention里的softmax或者dropout的mask,可能多占几个G。还有一个冷门的点,SGD如果设了momentum,并且用了Nesterov,它的临时变量存储方式跟AdamW不一样,backward时可能会额外保留中间激活,你可以开一下torch.autograd.detect_anomaly看看是哪个具体的op在爆。我怀疑不一定是模型本身变肥,而是分配策略和checkpointing的交互出了问题,建议先关掉gradient checkpointing跑一个epoch对比下显存峰值,如果降下来了,再调max_split_size_mz。