最近在公司做一个小项目,用PyTorch微调一个开源的7B模型,单卡A100 40G居然跑着跑着就OOM了。我用的batch size已经降到1了,gradient checkpointing也开了,但还是在forward到中间几层的时候显存爆掉。看了一些教程说可以用DeepSpeed ZeRO或者混合精度,但试了fp16反而loss震荡得很厉害。我怀疑是不是自己dataloader里塞了太多padding token?还是说7B模型本来就需要80G以上的卡?有没有大佬指点一下,或者分享一下你们微调小模型时常用的显存优化trick?先谢过了!
请教大佬们,用PyTorch训练时显存总爆,是我代码写太烂了吗?
全部回复
共 174 条fp16震荡大概率是loss scaling没调好,试试bf16,A100支持得很稳。
padding token确实吃显存,但7B全参微调40G本来就紧,建议LoRA先跑通再想别的。
fp16 loss震荡大概率不是精度问题,是你学习率没跟着调,混合精度下loss scale和lr的配合很关键,可以先试试把lr降到原来的十分之一再看看。另外7B模型fp16权重就得14G,optimizer状态再加14G,梯度再14G,光这些就40G出头了,你A100爆掉太正常了,不是代码烂,是数学上就不够。padding token那个确实有影响,但顶多占几个G的显存,不是主因,建议你用attention mask把padding部分真正屏蔽掉,而不是只靠pad_token_id。我自己的经验是直接上DeepSpeed ZeRO stage 2,offload optimizer到CPU,这样40G跑7B微调是够的,但要注意CPU offload会拖慢速度,得权衡。你如果不想用DeepSpeed,还有个土办法是冻结前几层embedding,只微调后面几层transformer block,显存能省一大截,效果对很多任务来说损失不大。最后建议你装个nvidia-smi监控一下每个step的显存峰值,看看是不是某个中间激活值特别大,有时候是attention的seq length太长导致的,可以试试对长序列做截断或分块。
fp16震荡大概率是loss scaling没调好,另外padding确实会白吃显存,试试动态padding+flash attention能省不少。
fp16震荡大概率不是精度问题,先查一下loss scale是不是没调好,或者某些层对精度特别敏感,可以试试bf16。padding token确实会浪费显存,但7B模型40G跑不起来更可能是activation占大头,你开gradient checkpointing后有没有确认是forward还是backward爆的?我建议把序列长度截断到512,再配合ZeRO stage 2,A100上跑7B微调是够用的。另外你用的是HuggingFace Trainer还是手写训练循环?后者的话记得手动清理中间变量。
这题我太熟了,之前微调13B的时候也卡在40G上怀疑人生。你fp16震荡大概率不是精度问题,是loss scale没调好,试试torch.cuda.amp的GradScaler把初始scale调大点,或者干脆换bf16,A100对bf16支持很好基本不掉点。padding token确实是隐形杀手,记得把attention mask传对,同时用动态padding或者干脆把数据集里特别长的样本截断一下,能省不少显存。另外7B全参微调在40G上本来就紧巴巴的,别死磕,上LoRA或者QLoRA吧,4bit量化加LoRA的话30G以内随便跑,效果也不差。如果非要全参,试试ZeRO stage 2加offload optimizer,但注意offload会拖慢速度,你batch size又小,可能变成显存省了时间爆了。最后检查下是不是有变量没detach,比如loss里混了graph,这种隐性显存泄漏最坑人,用torch.cuda.max_memory_allocated看下峰值到底在哪一层炸的。
fp16振荡大概率是loss scaling没调好,试试torch.cuda.amp的GradScaler,padding那边用masked attention能省不少显存。
7B用40G单卡确实紧,但你这配置还爆的话,查下是不是序列长度太长或者中间激活没释放。
fp16震荡大概率是loss scale没调好,试试bf16,7B在40G上真没必要硬扛。
padding确实占显存,但你这情况更像activation峰值问题,查下中间层输出大小。
fp16震荡大概率不是精度问题,先查下loss scaling和梯度裁剪,尤其7B模型用bf16会比fp16稳很多,A100对bf16支持很好。padding token确实会浪费显存,建议动态padding到batch内最长序列,能省不少。另外你试试把optimizer换成AdamW的8bit版,或者直接上ZeRO stage 2,7B单卡40G按理说能塞下,重点看下activation memory。
fp16 loss震荡大概率不是精度问题,先检查下loss scaling是不是没开对,或者试试bf16,A100对bf16支持很好,基本能直接替代fp16。7B模型在40G上跑微调确实很紧,但batch size=1还爆的话,多半是padding token在作怪,把max length设短点或者用动态padding能省不少显存。另外可以看看是不是优化器状态占太多了,AdamW的momentum和variance在fp32下很吃显存,换bitsandbytes的8bit优化器能省一大截。我自己微调6.7B时用这些组合拳,20G都跑得动,你可以先试试把input长度砍到512看还爆不爆。
fp16 loss震荡大概率不是精度问题,先检查下loss scale是不是没调好,或者试试bf16,A100对bf16支持很好,基本无痛切换。padding token确实会浪费显存,你可以在collate_fn里按batch最大长度动态padding,或者直接用attention mask把padding部分屏蔽掉,省下的显存可能比你开gradient checkpointing还多。7B全参微调40G确实紧,但理论上是能跑的,建议用torch.utils.checkpoint把激活重计算粒度调细一点,同时把optimizer换成AdamW的8bit版,能省不少。另外你确认下是不是dataloader的num_workers开太多导致的CPU内存瓶颈,有时候OOM报错其实是CPU侧先爆了。
fp16震荡大概率是loss scaling没调好,试试bf16,A100支持且稳得多。
padding token确实占显存,但你这情况更像activation峰值问题,查查中间层输出大小。
fp16震荡大概率是loss scale没调好,试试bf16,另外padding记得mask掉。
fp16 loss震荡大概率不是精度问题,先检查一下loss scale是不是被设成固定值了,动态loss scaling对7B这种规模挺关键的。另外padding token确实会白白吃掉显存,建议把attention mask配合动态padding一起用,或者直接按长度排序分桶,能省不少。你试过torch.compile吗?A100上配合cudagraphs有时候能压掉30%左右峰值显存,比单纯开gradient checkpointing效果好。最后确认下是不是中间激活值没释放,有些自定义layer会持有tensor引用导致显存泄漏,用torch.cuda.max_memory_allocated对比一下峰值和当前占用就能看出来。
fp16震荡大概率是loss scale没调好,建议试试bf16,A100支持得挺好。
fp16 loss震荡大概率不是精度问题,先查一下dataloader里padding是不是统一到最长序列了,把max_length设成实际长度能省不少显存。7B在40G上其实能跑,你试试把attention的key/value缓存清掉,或者换用FlashAttention,有时候能直接砍掉三分之一占用。还有个小技巧,把优化器状态用bitsandbytes量化到8bit,几乎不影响收敛但显存立省。最后实在不行就上ZeRO stage 2,比stage 3稳定多了。
fp16震荡大概率不是精度问题,你先看看loss曲线是不是一开始就炸,如果是的话多半是学习率没跟着缩放,试试把lr降到原来的十分之一再看看。另外padding token确实会白白吃掉显存,尤其你序列长度不齐的时候,建议用attention mask配合动态batch,或者干脆把最长序列截断到512。7B模型在40G上其实能跑,但你要把optimizer states offload到CPU,或者直接用Deepspeed stage2,光开gradient checkpointing不够的。还有个冷门技巧,把input_ids和attention_mask放到同一个tensor里,能省一点显存,虽然不多但有时候就差那几百MB。
7B模型在A100 40G上全参数微调确实紧巴,但理论上不该OOM这么早。你试试看把padding全部去掉,用动态batch或者按长度排序,有时候光这一项就能省掉30%显存。另外fp16震荡的话,检查下有没有给loss scaler留足更新步数,或者干脆用bf16,A100上稳很多。还有个小技巧,如果只微调LoRA而不是全参数,显存直接砍半,效果差距其实不大。
40G跑7B全参,说实话本来就悬,除非你序列特别短。我怀疑问题不在batch size,而是你max length设太大了,很多padding token其实占着显存不干活。可以先统计下实际有效长度,把max length砍到中位数或者90分位,显存能下来一大截。另外fp16震荡考虑下是不是学习率太高了,降到1e-5以下试试,或者换AdamW加权重衰减。
我遇到类似情况时,最后发现是中间激活值没释放,尤其是用了gradient checkpointing之后,某些层的激活被重复计算但没清干净。你可以用torch.cuda.max_memory_allocated对比下峰值和最后占用,看看是不是内存碎片化。另外试试把dataloader的num_workers调低,有时候CPU加载慢会导致GPU等待,显存反而被临时缓存占满。Deep
7B模型在40G上跑全参微调本来就很勉强,你开gradient checkpointing还爆说明大概率是序列长度+padding的锅,建议把tokenizer的padding策略改成动态padding或者直接截断到统一长度,能省下不少显存。fp16震荡的话试试bf16,A100对bf16支持很好,loss稳定性比fp16强不少。另外如果你只是微调任务,Lora或者QLora真的够用,省下来的显存能开大batch,训练速度反而更快。
fp16震荡大概率是loss scaling没调好,试试bf16或者冻结前几层参数,能省不少显存。
fp16震荡大概率是loss scaling没调好,可以试试bf16,A100对它的支持很友好,基本无损。padding token确实会白白占显存,用attention mask配合动态padding能省不少,或者干脆把dataset里短样本补齐到同一长度时用最省内存的打包方式。另外7B全参数微调40G本来就紧,建议先看看是不是激活值爆了,比如把max_length砍到1024,或者换用LoRA只训一小部分参数,效果差距不大但显存直接掉一个量级。