最近在公司做一个小项目,用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对它的支持很友好,基本不掉精度。padding token确实会白白占用显存,建议动态padding或者干脆把attention mask用起来,能省不少。7B全参微调40G确实紧,但也不是完全没戏,LoRA或者QLoRA值得一试,效果未必差。你checkpointing是包了整个model还是只包了中间层?有时候手动分段开能省更多。
7B全参数微调在40G上确实很极限,但OOM多半不是模型本身的问题。你试试看把padding全去掉,用动态batch按最长序列padding,能省不少显存;另外fp16震荡的话,检查下是不是loss scale没调好,或者某些层对精度太敏感,可以试试bf16。还有个小技巧,把optimizer换成Adafactor,能省一半优化器显存,很多人忽略这个。实在不行就上LoRA吧,效果其实不差,省下的显存还能加大batch。
fp16震荡大概率不是精度问题,你试试给loss scaling加个动态调整,或者干脆用bf16,A100对bf16支持很好,能省一半显存还不容易炸。padding那边确实是个坑,把dataloader里attention mask做好,别让模型算无效位置,省下的显存比你想的多。7B在40G上单卡微调其实够用,关键看你怎么切模型,可以试试把optimizer状态offload到CPU,或者用ZeRO stage2,别一上来就stage3。另外你查下是不是activation显存峰值卡在中间层的某个大tensor上,手动把那几层换成torch.utils.checkpoint试试。
fp16震荡大概率是loss scaling没调好,试试bf16,A100支持得很稳。
fp16震荡大概率是loss scale没调好,试试bf16或者torchao的int8量化,7B用40G完全够。
fp16震荡大概率不是精度问题,是你learning rate没跟着调,混合精度下lr通常要降到原来的1/3甚至1/5,另外loss scaling策略也看看,动态loss scaler有时候会疯狂放大梯度。padding token确实会浪费显存,但7B模型跑40G爆掉更可能是activation memory爆炸,你试试在forward里把不需要的中间变量del掉,或者用torch.utils.checkpoint把每个transformer block都包起来,别只开全局开关。另外检查下是不是dataloader的num_workers太多,每个worker会复制一份模型权重到显存,这个坑我踩过。如果还不行,直接上ZeRO stage 2,配合offload optimizer到CPU,A100 40G跑7B微调应该够,除非你序列长度拉到4k以上。最后建议你监控一下每层显存占用,用torch.cuda.memory._record_memory_history能看到具体哪一层爆的,别盲猜。
A100 40G跑7B微调确实会紧,但OOM大概率不是模型本身的问题。你试试把padding token彻底去掉,用attention mask把非padding部分单独处理,显存能省不少。fp16震荡的话,可以换bf16试试,A100对bf16支持很好,稳定性比fp16强很多。另外检查一下是不是优化器状态占了大头,用AdamW的话可以把eps调大一点,或者直接上8bit优化器。
fp16震荡大概率是loss scaling没调好,试试bf16,A100支持得很稳。
fp16震荡大概率是loss缩放没调好,试试bf16或者给padding mask加上,能省不少显存。
fp16 loss震荡这个我太有同感了,之前调stable diffusion的时候也遇到过,后来发现是loss scaling没调好,你可以试试动态loss scaler,或者干脆用bf16,A100对bf16支持很友好,精度损失小很多。另外你说怀疑padding token,这个方向其实挺对的,7B模型哪怕一个pad token也会走过全部transformer层,你可以试试把dataloader里按长度动态batch,或者用flash attention,它自带padding mask优化,显存能省不少。不过说实话,单卡A100 40G微调7B裸模型本来就紧巴巴的,我猜你多半还用了lora之外的全参数微调?如果只是做下游任务,强烈建议试试QLoRA,4bit量化加lora,显存能压到20G以内,速度还快。至于DeepSpeed ZeRO,单卡其实意义不大,那是多卡才划算,你不如先看看显存到底花哪了,用torch.cuda.memory_summary()打印一下,大概率是激活值占大头,这时候开activation offload或者把中间层输出存到CPU上再取回,都能解决。最后再补一句,别迷信gradient checkpointing,它省显存但是翻倍计算时间,有时候你batch开到4配合checkpointing反而比batch=1不checkpointing更稳,这个得自己试。
说实话你这配置和设置不该OOM的,我拿A100跑7B微调时batch size=1加gradient checkpointing也就吃20G左右,你大概率是序列长度或者padding没处理好,试试把max length限制到1024,然后dataloader里用collate_fn动态padding到batch内最长而不是全局最大。
fp16震荡的话别直接硬上,可以试试bf16,A100支持得很好,loss稳定很多;另外DeepSpeed ZeRO Stage 2加offload optimizer到CPU能省不少显存,就是慢点,但总比爆掉强。
对了你检查过模型里是不是有没冻结的embedding层?有时候代码里不小心把某些大tensor留在计算图里也会导致显存峰值异常,建议用torch.cuda.max_memory_allocated()看下峰值到底在哪一步炸的。
说实话你这配置跑7B微调确实有点紧,但40G A100理论上不是完全没戏,问题大概率出在padding和序列长度上。你试过把max_length从默认的2048砍到1024或者更短吗?很多开源数据集的padding token占比高得吓人,就算batch size是1,实际算力也浪费在无效token上,显存自然就爆了。另外fp16震荡的话,可以试试bf16,A100对bf16支持很好,稳定性比fp16强不少,loss基本不会飞。ZeRO Stage 2或者3配合offload也是个思路,但7B模型用ZeRO可能有点杀鸡用牛刀,先试试把attention的显存优化打开,比如xformers或者flash attention,能省不少。我之前微调6.7B模型,开flash attention加bf16,峰值显存能压到24G左右,你参考下。还有个容易被忽略的点,看看是不是优化器状态没分片,AdamW的momentum和variance在fp32下占显存很大,用bitsandbytes的8位优化器能再省一截。建议你先用torch.profiler看下每一层的显存分配,定位到底是哪一部分爆的,别急着全盘改代码。
说实话7B在40G上单卡微调确实得抠得很细,这个锅不完全在你。padding token影响没那么大,但建议把max_seq_len压到训练所需的最短长度,省下的显存很可观。fp16震荡大概率是学习率和warmup没配好,试试bf16,A100支持得很好,loss会稳很多。另外ZeRO stage 2外加offload optimizer到CPU,基本能再省出10G左右,你这配置跑起来应该没问题。
fp16震荡大概率不是精度问题,你看下是不是loss scaling没设置好,或者某些层本身就不适合半精度。7B在40G上其实能跑,关键得看序列长度和attention的计算量,你把max length砍到1024再试试。另外dataloader里padding确实会浪费显存,用collate_fn动态padding到batch内最长样本能省不少。还有个小技巧,如果只是微调,把不需要梯度的参数显式requires_grad_(False),省下的显存可能比你想象的多。
说实话你这配置和设置已经很到位了,7B模型在40G卡上微调本来就属于极限操作,不全是代码问题。A100 40G开gradient checkpointing后,纯模型权重加优化器状态大概就要吃掉25G左右,forward中间激活值稍微一波动就爆很正常。你提到dataloader里padding多,这确实是个隐藏杀手,建议试试把attention mask弄严格点,或者用动态batch按长度分组,能省不少显存。fp16震荡的话,可以换bf16试试,A100对bf16支持很好,loss稳定性比fp16强很多,很多开源项目现在默认bf16。另外ZeRO stage 2配合offload optimizer到CPU也能缓解,但要注意速度会慢一些。我自己的经验是,把序列长度截断到2048,再加上flash attention,基本能压进40G。别急着上80G,先看看是不是tokenizer把大量padding塞进去了,那个有时候比模型本身还吃显存。
fp16震荡大概率是loss scaling没调好,试试bf16或者给关键层单独开fp32。padding真别塞太多,动态padding能省不少显存。
说实话7B模型在40G上微调确实紧巴,但绝对不是完全没戏,你fp16震荡大概率是loss scale没调好,建议试试bf16,A100对bf16的支持很稳,能直接省一半显存还不用折腾动态缩放。另外padding token这个点你抓得挺准的,很多新手都栽在这,dataloader里把attention mask处理好,用动态padding或者直接把长序列截断到1024以内,显存能肉眼可见地降下来。
我自己的经验是,光开gradient checkpointing不够,还得配合optimizer的状态切分,比如AdamW的momentum和variance其实占了大头,用DeepSpeed ZeRO-2或者哪怕手动把优化器状态放到CPU上,都能再挤出不少空间。你提到ZeRO试了但没细说,是不是只开了stage 1?stage 2加上offload效果会明显很多,但要注意CPU通信会拖慢速度,小数据集上其实无所谓。
还有个容易被忽略的点,就是input ids的维度,有些tokenizer会把特殊token算进去导致实际序列比你想的长,你打印一下每步的max length看看是不是有异常长的样本在拖后腿。最后实在不行就换LoRA或者QLoRA,4bit量化加低秩适配器,7B模型在24G卡上都能跑,40G绰绰有余,效果也不比全量微调差多少,公司项目够用了。
fp16震荡大概率不是精度问题,先查一下loss scaling是不是没开对,或者试试bf16,A100对bf16支持很好,基本无感。padding token确实会浪费显存,7B模型哪怕1个batch塞进512长度的padding,中间激活也很恐怖,建议用attention mask加动态padding。另外40G跑7B微调其实够用,你试试把序列长度砍到1024以内,再加个activation offload,应该能稳。
fp16掉loss不稳大概率不是精度问题,你先查下dataloader里padding有没有mask到位,我之前就是没把attention mask传对,显存直接翻倍。7B在40G上完全能跑,你试试把input长度截断到512,再开个offload optimizer,比硬扛ZeRO省事。另外如果只是微调,冻结前几层embedding和layernorm,能省不少显存,效果基本不掉。
fp16震荡大概率不是精度问题,先查查loss scaling是不是没开对,或者某些层对精度敏感,试试bf16会稳很多。7B在40G上其实能跑,你dataloader里padding确实是个大头,建议把attention mask和动态padding用起来,能省不少显存。另外可以看看是不是激活值峰值在中间几层爆了,配合activation offload或者手动把forward拆开,分段清中间变量。
我之前也遇到过类似情况,最后发现是position embedding的缓存没清,每步都在涨。你检查下是不是把整个序列都pad到最大长度了,如果公司数据长短不一,用collate_fn动态拼batch能省一半以上。DeepSpeed ZeRO2比ZeRO3省通信,对单卡也有帮助,可以先试试stage2。
别急着怪自己代码烂,这坑我踩过。7B在40G上不开梯度检查点肯定爆,开了还爆大概率是序列长度太长,或者中间张量没及时释放。你试试在每层forward后手动del掉不用的中间结果,加上torch.cuda.empty_cache()。另外fp16震荡就换bf16,A100对bf16支持很好,基本无损。