最近在公司做一个小项目,用PyTorch微调一个开源的7B模型,单卡A100 40G居然跑着跑着就OOM了。我用的batch size已经降到1了,gradient checkpointing也开了,但还是在forward到中间几层的时候显存爆掉。看了一些教程说可以用DeepSpeed ZeRO或者混合精度,但试了fp16反而loss震荡得很厉害。我怀疑是不是自己dataloader里塞了太多padding token?还是说7B模型本来就需要80G以上的卡?有没有大佬指点一下,或者分享一下你们微调小模型时常用的显存优化trick?先谢过了!
请教大佬们,用PyTorch训练时显存总爆,是我代码写太烂了吗?
全部回复
共 174 条这问题我太熟了,7B模型在A100 40G上跑微调,batch size=1还爆显存,大概率不是卡的问题,而是数据形状和模型架构的细节没抠到位。
首先你提到padding token太多,这个确实是个常见坑。LLM为了对齐长度,很多人在dataloader里直接往长文本上补padding,但注意力机制会把这些padding也算进去计算,虽然结果会被mask掉,但它们依然占着显存做矩阵运算。你试一下把tokenizer的padding改成max_length或者直接动态batch,按最长的样本截断,不要统一补到某个固定长度,能省不少。另外检查下attention_mask是不是传对了,有些库的实现默认不传mask,模型会傻乎乎地算全序列。
再说fp16震荡这事,7B模型直接用原生的automatic mixed precision确实容易炸loss,尤其是微调阶段,梯度变化大,fp16的动态范围不够。你可以试一下bf16,如果卡支持的话,A100对bf16是原生支持的,稳定性比fp16好很多,而且显存用量跟fp16一样。如果只能用fp16,那得调一下loss scaling的策略,或者干脆先跑一个warmup阶段让模型适应一下。
还有一点容易被忽略:检查下你用的peft库是不是默认把全参梯度都保留了。像LoRA虽然只训练部分参数,但如果不手动设置requires_grad=False,PyTorch还是会给所有参数分配梯度的buffer,那样显存根本省不下来。建议用torch.no_grad()或者peft的disable_adapter模式确认下。
最后说句实在话,7B模型在40G上做全量微调确实勉强,但做LoRA或Q-LoRA的话是够的。你要不是必须用全参,试试4bit量化加LoRA,bitsandbytes库支持得很好,batch size能拉到4-8,训练速度反而更快。
7B模型在A100 40G上单卡微调确实吃紧,尤其你开了gradient checkpointing还OOM,问题大概率出在序列长度上——检查下dataloader里padding后实际最大长度,很多开源代码默认塞到2048甚至4096,光attention的中间激活就能吃掉十几G。fp16震荡的话,试试bf16或者给loss加个动态缩放,另外ZeRO-2配合activation offloading也能再省几个G,但记得把offload device设为cpu。
fp16 loss震荡大概率是loss scale策略没调好,试试bf16,A100原生支持,比fp16稳很多。7B模型40G其实能跑,但主要瓶颈在activations,你开了gradient ch
eckpointing还爆,可能是dataloader里padding太多导致序列过长,检查下每个batch的实际token数,用动态padding或batch sampler按长度分组能省不少显存。
同感,这问题太真实了,7B模型在40G卡上跑确实得精打细算。我先说一个可能被忽略的点:你提到dataloader里padding token太多,这个其实影响挺大的,尤其是如果序列长度参差不齐,你用了collate时统一padding到最长,那显存里会塞满无效的padding。建议试试PyTorch的pack_padded_sequence或者干脆用dynamic padding,每次只padd到当前batch的最大长度,能省不少。
另外fp16震荡的问题,我猜你可能没做gradient scaling?或者loss本身太大导致下溢?可以试试bf16,A100支持bf16,它动态范围比fp16大很多,震荡会小很多,而且很多7B模型微调社区已经默认用bf16了。
至于gradient checkpointing,你确认是开在了每个transformer层上吗?有些实现只在特定模块上开,中间层爆了说明可能没覆盖全。还有一个trick:把模型参数和优化器状态offload到CPU,用DeepSpeed ZeRO-3结合CPU offload,虽然会慢一点,但40G跑7B应该稳的。我上次调一个6.7B模型,batch size 1,开启ZeRO-3 + cpu_offload + bf16,显存峰值大概在32G左右,留了一些余量。
对了,你用的什么微调方法?如果是全参数微调,那确实对显存要求高,可以考虑LoRA或者QLoRA,4bit量化后7B模型在24G卡上都能跑,而且效果在不少任务上不输全参。可以试试bitsandbytes的4bit配置,配合gradient checkpointing,显存能压到20G以内。
fp16震荡的话试试bf16,另外看看tokenizer是不是没设max_length,padding太多确实会吃显存。
fp16 loss震荡可以试试bf16,很多7B模型对fp16敏感,另外检查下padding长度,建议动态batch。
fp16震荡可以试试bf16,另外检查下tokenizer的max_length是不是设得太大了。
fp16 loss震荡可以试试bf16,效果稳得多。另外检查下attention的padding mask是不是写对了。
fp16 loss震荡可以试试bf16,另外检查下tokenizer是不是把padding截断设好了。
说实话我觉得问题可能不在你代码写得烂,而是7B模型在40G的A100上做微调本身就挺极限的。我试过类似配置,就算开了gradient checkpointing,如果序列长度超过2048或者用了比较复杂的attention机制,显存还是会突然飙上去。fp16震荡那个我也遇到过,后来发现是loss scaling没调好,你可以试试先关闭动态loss scaling,手动设一个小的初始scale值看看是否稳定。
另外你提到padding token多,这个确实有影响。我建议你检查一下dataloader里的batch是不是因为padding导致实际有效长度参差不齐,可以考虑用padding-free的方式或者干脆把最大长度设小一点,比如1024甚至512,很多时候下游任务用不了那么长的上下文。我甚至试过用梯度累积配合更小的有效batch来缓解,虽然慢但至少不崩。
还有一个思路是换用LoRA或者QLoRA这类参数高效微调方法,对显存友好得多,7B模型在40G卡上基本能稳定跑。如果项目对全量微调不是刚需,这个方案最省心。
fp16震荡大概率是loss scaling没调好,可以试试bf16或者torch.cuda.amp的grad scaler,效果会稳很多。另外7B模型单卡40G跑全参微调确实紧张,就算开齐了gradient checkpointing和ZeRO stage2也可能爆,建议先看看是不是padding token太多导致的,用dynamic padding或者把max length设小点能省不少显存。实在不行就上LoRA吧,省下来的显存还能把batch size提上去。
fp16震荡太正常了,7B模型在A100 40G上纯fp16都悬,建议直接上bf16,很多开源模型本来就是bf16训练的兼容性更好。dataloader里padding token确实会有影响,试下把tokenizer的padding side设成left,或者用动态batch按实际长度分组,能省不少显存。另外check一下是不是中间激活值太大了,开activation checkpointing的同时可以配合gradient accumulation,batch size 1但accumulation steps设大一点,效果差不多。
说实话7B模型在40G的A100上跑微调确实有点极限,但爆在中间层而不是一开始就崩,说明大概率不是模型本身太大,而是中间激活值把显存吃满了。我之前也踩过类似的坑,后来发现是序列长度太长导致的,哪怕batch size=1,如果输入文本里有大量padding token,attention矩阵的计算量依然会爆炸,尤其是7B模型的hidden size不小,中间激活值叠起来很夸张。你可以试试在dataloader里做dynamic padding,把每个batch的序列长度对齐到实际最大长度而不是固定到模型上限,这样能省不少显存。混合精度fp16震荡的话,检查一下loss scaling是不是设成了动态模式,或者试试bf16,如果卡支持的话稳定很多。另外gradient checkpointing虽然开了,但有些层如果没包进checkpointing的范围内,比如embedding层或者自定义的attention模块,可能还是全精度跑,得确认一下具体哪些层被爆了。最后一个小建议,看看是不是优化器状态占的显存太多了,换AdamW的8-bit版本或者用Adafactor,也能挤出几G来。
7B模型在A100 40G上跑微调确实有点极限,尤其是用padding塞太多token的话,建议先检查一下数据集里实际序列长度,用dynamic padding或者把max length设小一点试试看。fp16震荡的话,可以试试bf16,A100对bf16支持更好,loss会稳很多。另外你开了gradient checkpointing但还爆,可能是中间激活值太大,可以配合activation offloading或者把模型某些层用torch.compile优化一下。
说实话7B模型在40G卡上全参数微调本来就是极限操作,你降batch size和开gradient checkpointing都对,但可能忘了关掉一些冗余的中间变量——比如检查一下有没有不小心把整个embedding层或者attention输出都存下来了。fp16震荡的话,试试bf16?A100原生支持的,很多时候比fp16稳。另外padding token确实是个坑,建议用dynamic padding或者把max length设小一点,能省不少显存。
fp16 loss震荡可能是loss scaling没调好,试试bf16或者torch.compile,对7B模型友好很多。另外检查下tokenizer的max length是不是设太大了,很多padding确实会白白吃掉显存。我自己的经验是开gradient checkpointing + 4bit量化能稳在40G上跑7B,但收敛速度会慢一点,你可以权衡下。
说实话7B模型在40G卡上全参数微调确实挺极限的,我遇到过类似情况。你开gradient checkpointing和fp16是对的,但loss震荡的话可以试试bf16(如果卡支持),或者调低学习率再加个warmup。另外dataloader里padding token太多确实会浪费显存,建议用dynamic padding或者把不同长度的样本按长度排序后再batch,能省不少。
说实话7B模型在40G上跑微调确实有点极限,尤其是你没开activation checkpointing的话,中间层的激活值很容易把显存吃光。你提到gradient checkpointing开了但还爆,建议检查下是不是只开了模型本身的checkpointing,而dataloader里padding token太多导致序列长度不均,PyTorch会按最长的来分配显存,这其实是个很隐蔽的坑。fp16震荡的话可以试试bf16,如果卡支持的话稳定很多,或者用torch.cuda.amp的autocast配合GradScaler手动调一下loss scale。另外DeepSpeed ZeRO 2或者3配合offload确实能省不少,但要注意通信开销,小项目直接用Hugging Face的Trainer里集成的deepspeed参数调一下就行。我自己的经验是7B模型单卡微调最好把max_length限制在1024以内,再用dynamic padding + sorting,基本能控制在35G左右。你那个项目如果对精度要求不高,可以考虑LoRA或者QLoRA,rank设小一点,显存直接砍半。
padding token确实占显存,试试把max length设小点,或者用torch.utils.checkpoint更细粒度控制一下。
fp16震荡大概率是loss scaling没调好,试试bf16(如果卡支持),稳定很多。7B模型在40G上单卡微调确实紧,建议检查下tokenizer是不是把padding截断设得太长了,或者用torch.utils.checkpoint把attention层也包进去。另外可以看看是不是optimizer states占了太多显存,换成AdamW的8-bit版本能省不少。