最近在公司做一个小项目,用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震荡大概率不是精度问题,先检查一下dataloader里padding有没有设成ignore_index,以及attention mask是不是传对了,我之前也被这个坑过。7B在40G上理论能跑,但你要是序列长度拉到2k以上,光激活值就很吃紧,试试gradient accumulation配合batch size=1,还有把input_ids和labels的padding分开处理。另外ZeRO stage 2其实对单卡也有帮助,主要是省下优化器状态的内存,你开了checkpointing之后可以再叠加试试。
fp16震荡大概率不是精度问题,是你loss scaling没调好,或者模型里有不稳定的层,可以先试试bf16,A100对bf16支持很好,基本无损。padding token确实会白白吃掉显存,建议自己写个collate_fn把序列截断到实际长度再动态padding,能省下不少。另外7B全参微调40G确实紧,就算能塞进去也基本没余量,不如直接上LoRA,省显存不说,效果也不差。你那个OOM是在forward中间层,也可能是activation峰值太高,试试torch.utils.checkpoint里把输入也切成小块,或者手动调一下max_seq_len,别让模型吃满上下文。
fp16震荡大概率是loss scaling没调好,试试bf16或者先查下padding有没有超过序列长度的10%。
fp16震荡大概率不是精度问题,是你loss scaling没调好,试试bf16或者torch.cuda.amp的GradScaler,7B在40G上其实能跑,关键看序列长度和attention计算量。padding token确实会占显存,建议把数据packing成固定长度或者用attention mask过滤掉无效位置,能省不少。另外你查下是不是activation显存峰值出现在中间层,可以配合torch.utils.checkpoint把transformer层单独包一下,别整个模型一起开。实在不行就上ZeRO stage 2,比纯fp16稳得多。
fp16震荡大概率是loss scale没调好,试试bf16或者torchao低精度,能省不少显存。
padding token确实占显存,把attention mask用上,或者动态padding到batch内最长长度。
fp16震荡大概率是loss scaling没调好,或者某些层对精度太敏感,可以试试bf16,A100对bf16支持很好,基本无损。padding token确实会浪费显存,尤其是序列长度参差不齐的时候,用attention mask配合动态padding能省不少。7B模型40G其实够跑,但要看你的序列长度和隐藏层大小,建议先打印每层的显存占用定位一下瓶颈。另外ZeRO stage 2搭配offload optimizer能再挤出一块空间,但速度会慢一些。
fp16震荡大概率是loss缩放没调好,建议试下bf16,A100上稳得多。padding确实费显存,用动态padding把序列对齐到batch内最长就行。
说实话fp16震荡这个太典型了,你试试看是不是某些层对精度特别敏感,尤其是embedding和最后的lm_head,很多人直接全层混精就翻车。我一般会先用torch.autocast配合GradScaler,然后单独把容易出问题的层强制回fp32,或者调一下loss scaling的初始值和增长频率,很多时候能解决。
至于padding token这个思路,你倒是提醒我了,如果序列长度差异特别大,塞太多pad确实会让attention的计算量虚高,显存自然就爆了。你可以试试把batch里的样本按长度排序,然后用动态padding或者干脆用packed sequence,哪怕只是把padding比例降到20%以下,省下来的显存都很可观。
7B在40G上做微调其实完全可行,我甚至见过用24G卡跑起来的,关键还是看你怎么分配显存。你开了gradient checkpointing但还在中间层爆,那大概率是激活值峰值出现在某些特别大的线性层或者attention上,建议用torch.profiler看下具体是哪一层占的峰值,然后针对性处理,比如把那个层拆成多个小矩阵乘法。
另外你提到DeepSpeed ZeRO,如果只是单卡的话ZeRO其实没啥用,不如直接上offload或者把优化器状态挪到CPU,但那样训练速度会掉一点。还有个偏方,你可以把batch size设成1然后梯度累积,虽然慢但至少能跑,先确认模型本身能不能过forward,再一步步调。
最后问一下,你的输入序列最长大概多少?如果超过2048,那可能得考虑截断或者用FlashAttention,这玩意儿能省不少显存,而且对长序列特别友好。
说实话7B模型在40G上微调完全不至于这样,你fp16震荡大概率不是精度问题,是学习率和warmup没配合好,尤其是用AdamW的时候。padding token确实会浪费显存,因为attention的计算和缓存是跟着序列长度走的,你试试把dataset里所有样本按长度排序然后bucket动态padding,能省下不少。另外你开了gradient checkpointing但显存还是在forward中间爆,我怀疑是activation checkpoint的粒度问题,可以考虑用torch.utils.checkpoint把每个transformer block单独包一下,别整个model一起checkpoint。还有个很实用的trick是给optimizer加offload,或者直接用bitsandbytes的8bit Adam,能省好几G。DeepSpeed ZeRO在这个场景其实没必要上,Stage 2配合offload就够,但你要注意ZeRO和gradient checkpointing有时候会有冲突。最后建议你装个nvidia-ml-py或者用torch.cuda.memory_summary()看下到底是activation还是梯度占大头,别瞎猜。我自己的经验是7B用fp16加batch size 4在40G上是能跑通的,你肯定有地方配置没对。
说实话你这配置和设置已经挺到位了,7B模型在40G上全参数微调本身就非常极限,别说A100了,我拿80G跑都得精打细算。padding token塞太多确实会浪费显存,因为attention矩阵是按最长序列算的,你试试把dataloader里按长度动态batch或者用pad to max length的策略,能省不少。fp16震荡大概率是loss scaling没调好,或者某些层对精度太敏感,你可以试试bf16,A100支持得很好,稳定性比fp16强一大截。另外DeepSpeed ZeRO Offload是个思路,但CPU offload会把训练速度拖慢到怀疑人生,建议先开ZeRO-2,把优化器状态切分掉,不要急着上offload。还有个冷门trick,把输入序列截断到256或者512,很多任务其实不需要那么长的上下文,收益立竿见影。最后别迷信gradient checkpointing,它只是用算力换显存,开的时候把checkpoint的粒度调到每层或者每两块Transformer块,别全开,反而可能更省。你要是还不行,就直接上LoRA或者QLoRA吧,4bit量化加低秩适配,几行代码的事,效果也不差,别跟自己过不去。
说实话fp16震荡大概率不是精度问题,你先看看loss曲线是不是前期就炸,如果是的话考虑下是不是学习率没跟着scale。padding token确实会白吃显存,建议把attention mask利用起来或者干脆动态padding到batch内最长序列。另外7B单卡40G理论够,但如果你用了全参数微调那肯定紧,试试LoRA或者只冻住大部分层训练,显存能省一半还多。最后查下是不是activation checkpointing没包对地方,有时候只包了transformer block但embedding和norm层还是全量存,那照样爆。
fp16震荡大概率是loss scale没调好,试试bf16或者torch.cuda.amp的GradScaler,能省不少显存。
padding token确实会白吃显存,用attention mask加动态padding,7B单卡40G完全能跑。
fp16震荡大概率是loss scale没调好,试试bf16或者给关键层保留fp32。另外7B用40G确实紧,padding和attention mask检查下吧。
padding token确实会拖累显存,建议动态padding到batch内最大长度,能省不少。
fp16震荡大概率不是精度问题,你试试给loss scaler加个dynamic策略,或者看看是不是某些层本身就对精度敏感,比如embedding和最后的lm head。显存这块,7B在fp16下光权重就14G了,加上激活值、gradient和optimizer state,40G确实紧巴巴的,但降到1batch还爆肯定是中间激活值太大,建议先排查一下序列长度是不是被padding到特别长,把attention mask和pad token处理好能省不少。另外DeepSpeed ZeRO stage 2配offload optimizer到CPU,能再挤出几G,先跑通再说。
fp16震荡大概率不是精度问题,而是loss scaling没调好,试试torch.cuda.amp的GradScaler开动态loss scaling,顺便把模型里容易爆的norm层留在fp32。padding token确实会浪费显存,但你这情况更像是激活值峰值太高,7B在40G上全参数微调本来就紧,建议先跑个profiler看看具体哪层爆的。另外可以试试把attention的seq len限制到实际长度,或者用flash attention,能省不少显存。我微调6.7B时用这些方法40G勉强能跑满,但batch size也只能到2,你再看看是不是优化器状态占了太多。
fp16震荡大概率不是精度问题,你查一下loss scaling是不是没开对,或者某些layer的梯度溢出被clip得太狠了,我之前用bf16在A100上反而更稳,你可以试试。7B全参微调40G确实紧,但也不是完全没戏,你开了gradient checkpointing还爆,那问题很可能出在输入序列长度上,padding token太多会导致attention矩阵按最长序列算显存,试试把dataloader里按长度bucket动态batch,或者干脆把padding截断到固定长度。另外你检查下是否用了flash-attention,这个对显存和速度的提升非常明显,普通attention在长序列上是真的扛不住。如果还不行,建议上LoRA或者QLoRA,4bit量化加低秩适配器,7B微调显存能压到16G以内,效果也不差,公司项目没必要死磕全参。还有个小细节,optimizer的momentum状态也吃显存,换成Adafactor或者SGD加动量能省不少,但得调参。最后,如果你非要用DeepSpeed,ZeRO-3配offload到CPU能救急,但速度会慢到怀疑人生,建议还是从模型和序列长度下手。
fp16震荡大概率是loss缩放没调好,试试bf16,A100支持而且稳很多。
padding token确实吃显存,把attention mask加上,长度统一截断到512试试。
fp16震荡大概率是loss scale没调好,试试bf16或者torch.cuda.amp的GradScaler。另外padding确实吃显存,attention mask配合动态batch能省不少。
fp16震荡大概率是loss scaling没调好,试试bf16或者torch.cuda.amp的GradScaler。另外padding别硬塞,用attention mask配合动态打包能省不少显存。
fp16 loss震荡大概率不是精度问题,你试试给loss scaler加个dynamic clipping,或者干脆用bf16,A100对bf16支持很好,基本无损。padding token确实会占显存,但7B模型40G应该够用,你查下是不是attention mask没传对,导致模型把padding也算了进去。另外可以试试把optimizer状态offload到CPU,或者用torch.utils.checkpoint把激活值再压缩一层,我上次跑13B就是这么省下来的。