最近在公司做一个小项目,用PyTorch微调一个开源的7B模型,单卡A100 40G居然跑着跑着就OOM了。我用的batch size已经降到1了,gradient checkpointing也开了,但还是在forward到中间几层的时候显存爆掉。看了一些教程说可以用DeepSpeed ZeRO或者混合精度,但试了fp16反而loss震荡得很厉害。我怀疑是不是自己dataloader里塞了太多padding token?还是说7B模型本来就需要80G以上的卡?有没有大佬指点一下,或者分享一下你们微调小模型时常用的显存优化trick?先谢过了!
请教大佬们,用PyTorch训练时显存总爆,是我代码写太烂了吗?
全部回复
共 174 条fp16震荡这个问题我也遇到过,7B模型在40G卡上硬跑确实容易卡在中间层,跟padding关系不大,主要是activation memory占大头。建议你试下torch.compile+gradient checkpointing组合,有时候比ZeRO省显存,而且对fp16稳定有帮助。另外检查下你的模型有没有未冻结的embedding层,那个也特别吃显存,能锁住就锁住。
7B模型在40G卡上单卡跑确实勉强,试试ZeRO-2加gradient checkpointing,应该能稳住。
跟你遇到的情况差不多,7B模型在40G卡上单卡跑确实挺吃力的,尤其前向中间层显存峰值高。fp16震荡厉害的话,试试bf16或者torch.compile,有时候比fp16稳定很多。另外检查下tokenizer有没有把padding搞太长,这个堆起来也挺占显存的。
fp16震荡很可能是loss scaling没调好,试试bf16或者torch.cuda.amp自动混合精度,能省不少显存。另外7B模型在A100 40G上单卡微调确实紧张,可以检查下tokenizer是不是把padding设得太长,或者用梯度累积来模拟更大batch,同时把optimizer状态切到ZeRO-2。我之前用LoRA微调类似大小的模型,只训练部分参数,显存直接降到20G以下,效果也没差太多。
fp16震荡大概率是损失缩放没调好,可以试试bf16,A100对它的支持很稳。
fp16震荡大概率不是精度问题,你试试给loss scale加个动态调整,或者干脆用bf16,A100对bf16支持很友好,能省一半显存还稳。padding token确实会拖累显存,但7B模型在40G上跑微调本来就紧张,建议你查一下forward的时候有没有把中间变量全存下来了,比如用torch.no_grad包住不需要梯度的部分。另外DeepSpeed ZeRO2比ZeRO3省通信,你这规模其实够用了,别一上来就上ZeRO3。
看到你说fp16 loss震荡,我猜你很可能是在用AdamW的时候忘了给beta2调参,或者loss scaling没稳住,这个比显存爆还坑。7B模型在A100上理论上是能跑的,你开gradient checkpointing之后显存瓶颈很可能不在模型权重,而在activation,尤其是序列长度长的时候,中间层的hidden state会非常吃显存。你可以试试把max sequence length限制在2048或者1024,很多情况下padding token确实会浪费大量显存,因为attention的计算和存储都是按最长序列来的。另外DeepSpeed ZeRO Offload是个办法,但单卡上ZeRO Stage 2就够用了,别直接上Stage 3,那玩意儿会频繁和CPU换数据,慢得你怀疑人生。还有一个容易被忽略的点,你dataloader里如果num_workers开得太多,每个worker都会复制一份模型副本的显存占用,这个在调试时可以先设成0试试。最后我建议你写个简单的显存监控脚本,每跑一步打印一下memory allocated和cached,看看到底是哪个阶段爆的,比瞎猜强多了。
fp16震荡大概率不是精度问题,你先查查loss spike是不是出现在某几个特定batch,如果是的话八成是padding token太多导致attention计算无效,试试dynamic padding或者把max length砍到实际长度分位数。另外7B在A100 40G上全参数微调本来就吃紧,就算ZeRO Stage 3也得开offload,建议直接上LoRA或QLoRA,显存能压到20G左右,效果一般任务上跟全量微调差距不大。你dataloader里要是没做token-level mask,padding对梯度的影响也会被放大,这个比显存更值得先排查。
7B全参微调在40G上确实很极限,但你的怀疑方向对了一半——padding token太多会白白占用显存,建议把max length设成数据分布的分位数,比如95%截断,能省不少。fp16震荡大概率是loss scale策略问题,可以试试bf16(A100支持)或者用SGD+warmup替代AdamW,收敛会稳很多。另外你开了gradient checkpointing但没提offload,可以试试ZeRO-3加CPU offload参数,虽然慢点但能兜底。最后一个小技巧,把输入里的attention_mask和token_type_ids显式传进模型,别让框架自动广播,有时候能省几个G。
fp16震荡大概率不是精度问题,你试试给loss缩放加个动态调整,或者直接换bf16,A100对bf16支持很好,基本能稳下来。padding token确实会占显存,但7B模型40G理论上够用,你检查下是不是max length设太长了,把序列截断到512试试。另外可以开torch.compile加activation offload,能省不少显存,我之前就是这么救回来的。
fp16震荡大概率不是精度问题,你检查下loss scaling是不是没开对,或者某些层本身对精度敏感,试试bf16会稳很多。padding token确实会白白吃显存,但7B模型在40G上就算全量微调也不该爆,重点看下是不是激活值没释放,建议用torch.cuda.memory_summary()蹲一下峰值在哪一层。另外ZeRO stage 2对你的场景够用了,stage 3反而可能因为通信开销拖慢速度。最后查下dataloader的num_workers,有时候数据加载也会占一部分显存,调成0试试。
fp16震荡大概率是loss scaling没调好,试试bf16,A100支持得很稳。padding token确实白吃显存,用attention mask配合动态padding能省不少。
fp16 loss震荡大概率不是精度问题,你查一下loss scaling是不是被关掉了,或者某些层在fp16下梯度下溢,这种情况可以试试bf16,A100对bf16支持很好,基本能直接替换而且省显存。padding token确实会白白占显存,你可以在collate_fn里把attention mask设好,或者干脆把batch内序列按长度排序再分组,能省不少。7B全参微调在40G上确实极限,但不至于完全跑不动,你检查下是不是把优化器状态和梯度都算进去了,AdamW的momentum和variance每参数要8字节,加上fp16的master weight,光优化器就吃掉2倍模型大小。如果模型结构允许,可以考虑freeze掉embedding和部分底层layer,只训后面的层,显存瞬间下来。另外你确认一下是不是在forward里创建了不必要的中间变量,比如把整个序列的hidden state都存下来做hook,这种隐形开销很烦人。最后提个野路子,用torch.utils.checkpoint把每个transformer block都包一下,虽然慢点但能撑住。我上次微调一个6.7B模型,4卡3090用ZeRO-2加offload,batch size开到4也没爆,你可以参考下配置。
7B全参微调在40G上确实很极限,但fp16震荡大概率是loss spike,可以试试bf16或者给关键层单独开fp32。padding token影响没你想的那么大,真正吃显存的是激活值,建议用torch.utils.checkpoint配合gradient accumulation,再不行就上LoRA或者QLoRA,效果差不了太多。你用的是HuggingFace的Trainer还是手写训练循环?有时候默认的data collator会偷偷把序列padding到最长,检查下实际batch里的max_length。
fp16震荡大概率不是精度问题,先看看loss scale是不是没调好,或者某些层对精度特别敏感,可以试试bf16,A100支持得挺好。padding token确实会浪费显存,建议把序列长度统一到实际最大长度,或者用attention mask配合动态padding,能省不少。另外7B模型用40G单卡其实够呛,除非你只训LoRA或者冻结大部分层,全参数微调的话ZeRO stage 2或3基本是必须的。我自己的经验是开gradient checkpointing之后把batch再拆成micro-batch,配合梯度累积,能稳很多。你检查下是不是中间激活值没释放,有时候torch的缓存机制会骗你,用torch.cuda.empty_cache()看看真实占用。
fp16震荡大概率不是精度问题,先看看loss scale是不是没调好,amp的初始scale设太高了很容易这样,我一般直接关掉动态loss scaling手动设个固定值。padding token确实会白白占用显存,但7B模型就算输入很短,单卡40G吃紧也正常,毕竟微调时激活值开销比推理大得多。你试试在dataloader里按长度动态batch,或者把序列截断到512,能省不少。另外可以看一眼是不是把整个模型都放进了device,有时候embedding层单独留float32也会拖后腿。
7B全参数微调在40G上本来就紧,你开fp16震荡大概率是loss scale没调好,可以试试bf16,A100支持得更好,基本没精度问题。padding token确实会白占显存,建议把dataset里按长度分组然后动态padding,能省不少。另外你只提到gradient checkpointing,但offload optimizer到CPU试过没?ZeRO stage 2加offload能把优化器状态挪走,峰值能降一大截。最后检查下forward里有没有创建大的中间变量,比如attention mask或者position id,有时候隐式广播会突然爆显存。
说实话7B模型在40G上全参数微调本来就很极限,你开gradient checkpointing还爆说明问题大概率不在显存总量,而在临时激活值峰值。我遇到过类似情况,最后发现是序列长度太长加上padding没mask好,attention矩阵的显存占用是二次增长的,你把max_seq_len砍到2048或者1024试试,效果立竿见影。
fp16震荡这个太典型了,多半是学习率没跟着调,混合精度下loss scale也会影响稳定性,你可以试试bf16,A100对bf16支持很好,不用改太多代码,数值范围也更大。另外dataloader里padding token确实是个隐藏杀手,不光浪费算力,还会让attention算到无效位置,建议你写个collate_fn动态pad到batch内最大长度,再把attention mask传进去,显存能省下不少。
DeepSpeed ZeRO Stage 2在这种单卡场景其实帮助不大,它主要解决多卡通信和参数分片,单卡你不如直接开torch.utils.checkpoint配合显存优化器,比如8bit Adam。我之前微调7B用40G卡,batch size开4都没问题,关键是把输入序列截断到2k,加上梯度累积模拟更大batch。你要是实在不想折腾,直接上QLoRA,4bit量化后7B模型权重才4G左右,留出大量空间给激活值,效果也不差。
fp16震荡大概率不是精度问题,你看下是不是loss scaling没设置好,或者某些层对精度特别敏感,可以试试bf16,A100对bf16支持很好基本无损。padding token确实会浪费显存,建议用attention mask或者直接把序列按长度分组再padding,能省不少。另外7B全参微调40G确实紧,但如果你只是做LoRA的话,显存占用能降一大半,效果也不差。我上次调7B就是LoRA加gradient checkpointing,batch size开到4都没爆过。
顺便问下你用的是HuggingFace的Trainer还是自己写的训练循环?有时候官方实现里有些隐藏的缓存没清理,也会导致显存越涨越高。
fp16震荡大概率不是精度问题,先查一下loss scale是不是没调好,或者某些层对精度特别敏感,可以试试bf16,A100支持得很好。padding token确实会白吃显存,但你这个情况更可能是7B模型本身在40G上就很极限,尤其序列长度一上来,中间激活值爆炸很正常。我建议你直接上梯度累积,batch size=1但累积个8步,效果差不多,显存还能匀出来给序列长度。另外看看是不是有地方不小心把梯度传到了embedding层,那个显存占用也很夸张。实在不行就考虑LoRA吧,微调7B真没必要全参硬刚。