最近在公司做一个小项目,用PyTorch微调一个开源的7B模型,单卡A100 40G居然跑着跑着就OOM了。我用的batch size已经降到1了,gradient checkpointing也开了,但还是在forward到中间几层的时候显存爆掉。看了一些教程说可以用DeepSpeed ZeRO或者混合精度,但试了fp16反而loss震荡得很厉害。我怀疑是不是自己dataloader里塞了太多padding token?还是说7B模型本来就需要80G以上的卡?有没有大佬指点一下,或者分享一下你们微调小模型时常用的显存优化trick?先谢过了!
请教大佬们,用PyTorch训练时显存总爆,是我代码写太烂了吗?
全部回复
共 174 条40G跑7B微调确实紧,但爆在中间层大概率不是模型大小问题,而是激活值峰值太高。你可以先试试torch.utils.checkpoint配合gradient accumulation,把有效batch size撑起来,同时用torch.profiler看下到底是哪层峰值爆的。另外padding token太多确实会浪费显存,建议把dataset里按长度排序后动态padding到batch内最大长度,能省不少。fp16震荡的话,可以试试bf16,A100支持得挺好,loss稳定很多。别急着上DeepSpeed,先把这些基础项调完再说。
fp16震荡大概率不是精度问题,你试试给loss加个warmup或者换个优化器,AdamW配cosine schedule会稳很多。7B模型就算用A100 40G,如果序列长度超过2048,单卡确实很吃力,建议先查一下实际显存占用是不是被激活值吃掉的。padding token影响不大,但如果你用动态padding+attention mask,能省不少算力。另外ZeRO stage 2配合gradient checkpointing,通常能把峰值压下来一半,你可以先不开混合精度试试。我自己的经验是,把模型切分到多卡比硬调单卡省心多了。
fp16震荡大概率是loss scaling没调好,padding确实浪费显存,试试动态padding加bf16。
7B全参微调40G确实紧,建议先看下有没有显存碎片,换torch.compile或者offload优化器状态试试。
fp16震荡大概率是loss scale没调好,试试bf16,A100支持得很稳。
padding token确实吃显存,但你这情况更像activation爆了,建议把gradient checkpointing开到每层都checkpoint,别省。
fp16震荡大概率不是精度问题,先查一下loss scale和梯度裁剪,尤其是7B这种规模,bf16通常比fp16稳得多。另外padding token确实会白吃显存,建议把attention mask用起来,或者干脆动态padding到batch内最长序列。A100 40G跑7B全参微调本来就紧,LoRA或者QLoRA能省一大半显存,效果也不差。你试试把gradient checkpointing和offload优化器一起开,我这边之前也遇到过类似情况,最后是换成LoRA才跑通的。
fp16震荡大概率是loss scaling没调好,试试bf16,A100支持而且稳得多。
说实话7B模型在40G卡上微调是可行的,你这个配置不该OOM,问题大概率出在padding上。我遇到过类似情况,dataloader里如果序列长度参差不齐,padding token一旦多了,计算量会虚高好几倍,显存自然就炸了。建议你先统计一下实际序列长度分布,用动态padding或者把max length设得保守点,这一步省下来的显存可能比开什么checkpointing都管用。
另外fp16震荡厉害不一定是不该用混合精度,很可能是学习率没跟着调,或者loss scaling策略不对。我一般会先用bf16试试,A100对bf16支持很好,数值稳定性比fp16强不少,如果模型本身是bf16预训练的,微调时直接沿用就好。你还可以看看是不是优化器状态占了大头,AdamW的动量项在7B规模下很吃显存,换8bit优化器或者用Adafactor能立竿见影。
再有个细节,gradient checkpointing开了以后要确认它真的作用在每一层上,有时候你只是开了全局开关但模型内部某些模块没实现对应接口,等于没开。最后如果实在不行,可以试试把模型切成两半用CPU offload,虽然慢点但至少能跑通,先把效果验证了再说。
fp16震荡大概率是loss scaling没调好,另外检查下attention mask,padding多确实吃显存,试试动态padding或者packing。
7B全参微调40G确实紧张,先确认下是不是序列长度太长,把max_len砍到1024再试试。
fp16震荡大概率是loss scale没调好,试试bf16,A100支持得很稳。
fp16 loss震荡大概率不是精度问题,先检查下loss scaling是不是没开对,或者试试bf16,A100上比fp16稳很多。padding token确实会白白吃显存,用attention mask或者把序列截断到实际长度能省不少。7B模型40G其实够跑,但要看你的序列长度和额外显存开销,建议用torch.cuda.max_memory_allocated()看看峰值到底在哪层。另外ZeRO stage 2配合offload optimizer能省一大半优化器状态,比单纯开gradient checkpointing效果好很多。
说实话7B模型在40G上微调确实有点紧张,但绝对不是不可能,你这情况大概率不是代码烂,而是几个坑叠一起了。padding token这个点你怀疑得很对,长序列里大量pad会被计算进去,白吃显存,建议用attention mask把pad区域彻底屏蔽,或者干脆动态batch按长度分组,能省不少。另外fp16震荡的话,试试bf16,A100对bf16支持很好,loss稳定性比fp16强一个档次,很多开源项目默认就是bf16。还有你开了gradient checkpointing但显存还是爆,可以查一下是不是optimizer state占了太大头,7B模型光AdamW的state就要好几个G,用DeepSpeed ZeRO Stage2或者直接上8bit optimizer能把这块砍掉一大半。最后提醒一下,如果中间几层爆,可能是activation峰值太高,可以考虑把输入序列截断到2048试试,很多任务其实不需要那么长的上下文。我之前微调13B也就用40G勉强够,关键还是看你怎么分配显存,而不是一味降batch,希望这些对你有用。
fp16震荡大概率是loss scaling没调好,试试bf16,A100支持得很稳。
fp16震荡大概率是loss缩放没调好,试试bf16或者给pad设mask,A100跑7B不该爆的。
fp16震荡大概率不是精度问题,是你loss scaling没调好,或者学习率太大,可以先试试bf16,A100对bf16支持很好,基本无损。padding token确实会浪费显存,但你这都OOM在中间层了,更像是activation memory的问题,建议用torch.utils.checkpoint把每个transformer block都包一下,别只开全局开关。另外7B全参微调40G确实很极限,哪怕batch size=1,中间层的hidden state也吃不少,可以看看是不是序列长度太长,把max length砍到1024或2048试试。我之前微调6.7B用类似配置,差不多是35G左右,你参考下。
fp16震荡大概率是loss scale没调好,试试bf16或者torchao的int8量化,能省不少显存。
fp16震荡大概率不是精度问题,先查一下loss scaling有没有设置好,或者试试bf16,A100对bf16支持更好,能省不少显存。padding token确实会白占显存,建议在collate_fn里按batch内最长序列动态padding,别用全局max length。另外7B全参微调40G确实紧,如果只是任务适配,可以试试LoRA或者冻结大部分层只训最后几层,效果不一定差。还有个小技巧,把optimizer换成Adafactor,能省一半优化器状态显存。
这情况我也踩过坑,7B全参微调在40G上基本是极限操作,就算勉强跑起来也容易不稳定。你试试把input_ids里pad的部分在attention mask里彻底屏蔽,有些实现会忽略但实际还是参与计算。另外fp16震荡的话,看看是不是某些层梯度爆炸,可以加个grad clip,或者改用bfloat16试试,A100上稳很多。实在不行就上ZeRO stage 2,比混合精度省心。
fp16震荡大概率不是精度问题,你先看看loss曲线是不是前期就炸,如果是的话把learning rate降一个量级试试,我遇到过类似情况。padding token确实会浪费显存,但7B模型40G跑不起来主要还是activation占大头,你可以试试torch.utils.checkpoint配合input chunking,把sequence切成几段过。另外ZeRO stage 2其实比fp16更稳,单卡也能用,就是速度会慢点。你模型是用的decoder-only架构吧?那attention的中间变量很容易爆,建议开一下flash-attention,能省不少。
说实话7B模型在40G上微调是完全可行的,我甚至用3090 24G跑过,所以肯定不是模型大小的锅。你fp16震荡大概率是loss scale没调好或者某些层对精度太敏感,可以试试bf16,A100对bf16支持很好,而且动态范围大,基本不会像fp16那样溢出。另外padding token确实是个隐藏杀手,你把attention mask和labels都处理一下,确保pad位置不参与loss计算,能省不少显存,我一般还会在collate_fn里按序列长度排序然后动态padding,这样每个batch的冗余计算会少很多。你开了gradient checkpointing但还爆,可能是在forward中间层时激活值峰值太高,可以配合activation offload或者把模型切成几段手动控制显存释放。还有个小众技巧,optimizer换成Adafactor或者8bit adam,能省下大概2-4G的优化器状态显存。最后建议你用torch.cuda.max_memory_allocated()打点看看具体哪一步峰值最高,别瞎猜,定位到具体层再针对性优化。
7B全参数微调在40G上确实很极限,但你说fp16震荡明显,先检查下是不是loss scale策略没调好,或者某些层对精度特别敏感。我之前遇到过类似情况,把padding token attention mask显式处理一下,再配合gradient accumulation模拟大batch,能稳定不少。另外你试试ZeRO stage 2加offload optimizer,有时比单纯开gradient checkpointing省得多。
fp16震荡大概率是loss scale没调好,试试bf16,A100支持得很好,能省一半显存。
padding token多了确实浪费显存,建议动态padding或者按长度分桶,能省不少。