最近在公司做一个小项目,用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支持且稳得多。
fp16震荡大概率不是精度问题,是你learning rate没跟着调,一般要降到fp32的1/3左右。另外7B模型光是参数就占14G,加上激活值和梯度,40G确实紧张,但绝对没到必须上80G的程度。你试试把max_seq_len砍到512,padding那边用attention mask遮住就别真塞那么多token,还有optimizer换成AdamW的8bit版,能省下不少显存。顺便检查下是不是dataloader的num_workers开太多导致内存碎片,这玩意儿有时候比显存更坑。
fp16震荡大概率不是精度问题,你试试给loss scaling设个动态范围,或者干脆用bf16,A100上稳得很。7B全参微调40G确实紧,但batch=1+梯度累积不至于爆在中间层,查下是不是padding没mask掉,导致attention矩阵算了一堆无效位置。另外你开gradient checkpointing的时候,记得把输入也切成slice,不然激活值照样存整份。我上次调一个小模型也是卡这,最后把序列长度砍到512才跑通。
fp16震荡大概率不是精度问题,是loss scaling没调好,或者学习率需要跟着降一降,你可以试试bf16,A100对bf16支持很好,基本无损。另外padding token这个方向确实值得查,7B模型哪怕一个batch塞太多无意义token,激活值也会非常可观,建议把attention mask和动态padding用起来。40G跑7B微调其实够用,我见过有人用LoRA加gradient checkpointing把峰值压到20G以内,要不你先换个思路,别全量微调,只训adaptor试试?
fp16震荡大概率是loss scale没调好,试试bf16,A100支持且稳得多。
fp16震荡大概率不是精度问题,是loss scaling没调好,可以先试试torch.cuda.amp的GradScaler,顺便把动态loss scaling开着,能解决大部分震荡。另外padding token确实是隐形杀手,7B模型就算pad到512长度,无效计算也白吃显存,建议用attention mask配合动态padding,或者直接按batch内最大长度截断。A100 40G跑7B全参微调本来就紧,你要是用了AdamW,光优化器状态就占快2倍模型显存,换8-bit Adam或者只用LoRA能省一大截。最后check一下是不是中间激活值没释放,pytorch的autograd有时会把中间tensor留着,手动del一下或者用torch.no_grad包住不需要梯度的前向部分试试。
A100 40G跑7B微调确实紧,但OOM不全是代码的锅,padding token影响没那么大,fp16震荡大概率是loss scaling没调好。你可以看看是不是激活值峰值太高,试试在forward里手动插torch.cuda.empty_cache(),或者把序列长度截断到2k以内。另外ZeRO stage 2配合offload optimizer能省不少,但记得关掉gradient checkpointing避免冲突。我之前用类似配置跑13B都扛得住,你检查下是不是哪里把梯度也塞进显存了。
fp16震荡大概率是loss scale没调好,试试bf16或者torch.cuda.amp的GradScaler,另外padding记得用attention mask屏蔽掉。
fp16 loss震荡的话,大概率是loss scaling没调好,或者模型里有敏感层,可以先试试bf16,A100对bf16支持很好,基本无损。另外padding token确实会白白占显存,建议把attention mask用起来,或者干脆动态padding到batch内最长序列。7B全参数微调40G确实紧,但也不至于完全跑不动,你可以看看是不是激活值缓存没释放,试试max_memory参数给torch.cuda设置个上限,让一部分参数留在CPU。
fp16震荡大概率不是精度问题,你看看是不是loss scaling没调好,或者某些层对精度特别敏感,可以试试bf16,A100支持得很好。padding token塞太多确实会浪费显存,但你这情况更可能是7B模型本身在40G上做全参数微调就够呛,建议先跑个纯forward看基线占用。我之前用LoRA把7B压到24G内很稳,你可以考虑下参数高效微调,效果也不差。另外检查下dataloader的num_workers和pin_memory,有时候这些细枝末节反而会拖后腿。
fp16震荡大概率不是精度问题,先查一下loss scale是不是没调好,或者某些层对精度太敏感,可以试试bf16,A100支持得很好。另外7B全参微调40G确实很极限,就算能塞进去也不稳定,建议先看看是不是padding token导致的无效计算,用attention mask把padding部分屏蔽掉能省不少显存。
fp16 loss震荡大概率不是因为精度,是你学习率没跟着调,试试把lr降到原来的十分之一,或者用bf16,A100对bf16支持更好。padding token确实会浪费显存,但7B模型40G卡理论上能跑,重点检查一下attention mask和position id有没有正确传。另外你开gradient checkpointing的话,记得把输入尺寸压到最长序列的80%左右,很多dataloader默认按batch里最长样本padding,这也会莫名其妙多吃不少显存。如果还不行,就看看是不是优化器状态占了太多,AdamW的momentum在7B上能吃掉好几个G。
7B全参微调40G确实悬,试试LoRA只调适配器,显存直接砍一大半。
7B模型单卡40G微调,就算bs=1加checkpointing,optimizer states和activations加起来也很容易炸,尤其你如果没冻结底层或者用了全参数微调。fp16 loss震荡可以考虑换成bf16,A100原生支持,数值稳定得多,基本不用调loss scale。dataloader的padding确实要看看,但一般不是OOM主因,建议先print一下每个step的显存占用,定位到底是哪一层爆的。实在不行就上LoRA或者ZeRO-2,7B全参微调本来就挺吃卡的。