最近在公司做一个小项目,用PyTorch微调一个开源的7B模型,单卡A100 40G居然跑着跑着就OOM了。我用的batch size已经降到1了,gradient checkpointing也开了,但还是在forward到中间几层的时候显存爆掉。看了一些教程说可以用DeepSpeed ZeRO或者混合精度,但试了fp16反而loss震荡得很厉害。我怀疑是不是自己dataloader里塞了太多padding token?还是说7B模型本来就需要80G以上的卡?有没有大佬指点一下,或者分享一下你们微调小模型时常用的显存优化trick?先谢过了!
请教大佬们,用PyTorch训练时显存总爆,是我代码写太烂了吗?
全部回复
共 174 条7B全参数微调在40G上本来就紧巴巴的,你优化都做了还爆,大概率不是代码烂,是模型本身吃显存就凶。fp16震荡的话可以试试bf16,A100支持得很好,loss稳定很多。另外dataloader里padding确实会浪费显存,建议用动态padding或者把max length卡到训练集的实际分布上,能省不少。还有个小技巧,把optimizer换成AdamW的8-bit版,或者用offload到CPU,能再挤出一块空间。
fp16震荡大概率是loss scaling没调好,试试bf16或者torch.cuda.amp的GradScaler。
fp16 loss震荡大概率是loss scaling没调好,或者某些层对精度太敏感,你可以试试bf16,A100对bf16支持很好,基本无损。另外padding token确实会浪费显存,尤其序列长度差异大的时候,建议用attention mask加动态padding,能省不少。7B模型40G其实够跑,我之前用ZeRO stage 2加offload optimizer,batch size 2也能稳定训练,你可以先排查下是不是激活值峰值太高,试试torch.utils.checkpoint里细粒度控制哪些层需要重计算。
fp16 loss震荡大概率不是精度问题,你试试给loss scaling加个动态调整,或者直接换bf16,A100对bf16支持很好,基本无损。padding token确实会浪费显存,但7B模型40G跑不起来更可能是激活值峰值太高,检查下有没有用flash attention,能省不少显存。另外你开gradient checkpointing的话,batch size可以试着调回2,因为显存大头其实在激活值缓存,checkpointing已经把这块削了。最后实在不行就上ZeRO stage 2,单卡也能用,只是图个心理安慰,但说不定能压过临界点。
fp16震荡大概率不是精度问题,先查一下loss scaling和梯度裁剪,特别是7B模型用bf16会比fp16稳很多,A100对bf16支持也更好。padding token确实会吃显存,但你这情况更可能是activation占大头,试试把max_seq_len砍到512或者用torch.utils.checkpoint配合input checkpointing,能省不少。另外ZeRO stage 2加offload optimizer到CPU也能救急,但速度会慢点。之前我调6.7B模型时,把attention里的dropout关掉、用flash-attention,显存直接少了快三分之一。
fp16震荡大概率是loss scale没调好,试试bf16或者给关键层单独开fp32。
padding太多确实会浪费显存,用动态padding或者干脆把max length砍到实际需要长度。
A100 40G跑7B微调确实紧,但OOM不一定是代码烂,padding token影响没你想的那么大。你试试把序列长度截到512或1024,再配合gradient checkpointing,应该能省不少。fp16震荡的话,检查下loss scaling是不是没调好,或者直接用bf16(A100支持),稳很多。另外,7B全参数微调其实很吃显存,LoRA或者QLoRA会是更实际的方案,效果也不差。
fp16震荡大概率是loss scale没调好,试试bf16或者torch.cuda.amp的GradScaler。padding token影响没那么大,查查中间激活值是不是没释放。
同款问题踩过坑,你fp16 loss震荡大概率是某些层对精度太敏感,试试bf16或者只对后半部分层开混合精度。padding token那边确实会白白占显存,可以按batch内最大长度动态padding,能省不少。7B全参微调40G确实紧张,但也不是完全没戏,把optimizer换成Adafactor或者用8bit版能腾出好几个G。另外看看是不是activation显存爆的,把activation checkpointing开到full层别只开部分,代价是慢一点但稳。
fp16震荡大概率不是精度问题,你试试给loss scale加个动态调整,或者干脆用bf16,A100对bf16支持很好,基本无损。padding token确实会浪费计算,但40G爆在中间层更像是激活值峰值问题,可以看看是不是序列长度太长,尝试在dataloader里按长度动态batch,或者用torch.utils.checkpoint把每个transformer block都包一下。7B全参微调本来就很吃显存,LoRA或者QLoRA不香吗,省下来的显存还能加大batch让训练更稳。
fp16震荡大概率是loss scaling没调好,开个Dynamic loss scaler试试。另外padding token确实吃显存,记得用attention mask过滤掉。
7B在40G上其实能跑,你试试把max length砍到512,再开offload optimizer,应该能稳。
fp16震荡大概率不是精度问题,是你loss scaling没调好或者某些层对精度太敏感,试试bf16,A100对bf16支持很好,基本无损。另外7B全参数微调40G确实紧张,但绝不是没救,你padding token塞多了会直接拖长序列长度,显存是按最长序列算的,建议把dataset里所有样本pad到同一长度改成动态padding,或者干脆用packing把短样本拼一起,能省不少。我自己的经验是,光把padding优化掉,同样batch size显存能降30%左右。还有个冷门trick,把input_ids和attention_mask放到同一个tensor里,或者用torch.utils.checkpoint把中间激活全扔掉,你开了gradient checkpointing但可能没配合activation offload,试试把checkpoint的keep_fp32_wrapper设成False。如果还不行,直接上LoRA吧,7B用LoRA微调单卡16G都够,效果跟全量微调差距很小,尤其你只是做公司内部小项目,没必要死磕全参数。最后查一下你的优化器状态,AdamW的momentum和variance也吃显存,换Adafactor或者8bit优化器能再省几个G。
fp16震荡大概率不是精度问题,你看下是不是loss scaling没调好,或者某些层对精度太敏感,可以试试bf16,A100支持的,稳很多。padding token确实会白占显存,但你这情况更像是激活值峰值爆了,7B模型就算开gradient checkpointing,中间层激活也容易到40G的极限。建议把序列长度砍到512或1024试试,再不行就上ZeRO stage 2,offload到CPU,速度慢点但至少不OOM。之前我微调6B模型也卡在这,后来是换了个更激进的attention实现才解决的,你可以看看xformers。
fp16震荡大概率是loss scaling没调好,试试bf16或者torchao的int8量化,能省不少显存。
fp16震荡大概率是loss scaling没调好,可以试试bf16,A100对它的支持很友好,基本能稳住。padding token确实会白白吃显存,建议把attention mask用起来,或者干脆动态padding到batch内最长序列。另外7B全参微调40G确实紧,但配了gradient checkpointing还不至于爆,你查下是不是中间激活值没释放,或者优化器状态没走offload。我上次用ZeRO stage2加offload把13B塞进24G卡都跑起来了,你可以试试把优化器状态和梯度都扔CPU上。
fp16震荡大概率是loss scaling没调好,你可以试试bf16,A100对bf16支持很好,基本无损而且省一半显存。padding token确实会浪费计算,但7B模型40G按理说够用,你检查下是不是max length设太长或者中间激活值没释放。另外ZeRO stage 2配合offload optimizer能再省不少,我上次微调13B就是这么扛下来的。你dataloader里要是做了动态padding,记得把attention mask传对,不然白省显存还影响收敛。
fp16震荡大概率是loss缩放没调好,试试bf16或者torch.cuda.amp的GradScaler,padding确实占显存但7B吃40G也正常。
fp16震荡大概率不是精度问题,先看看loss scale有没有设对,amp的dynamic loss scaling有时候在7B上确实不太稳,可以试试bf16,A100对bf16支持很好,基本不掉点。padding token这块确实容易忽略,建议把attention mask做严格一点,或者用packed sequence,能省不少显存。7B在40G上理论上够跑,但你要是序列长度超过2k,那爆也正常,可以算一下activation峰值,别光看模型权重。另外ZeRO stage 2加上offload optimizer到CPU,能解放不少显存,就是慢点,但至少不会OOM。
同款问题我上周刚踩完坑,先说结论:7B在40G上绝对能跑,问题多半出在padding和attention mask上。你试fp16震荡大概率是loss scale策略没调好,可以看看混合精度下是不是某些层梯度下溢了,或者干脆用bf16试试,A100对bf16支持很好,基本无痛切换。
另外dataloader里塞太多padding确实是个隐形杀手,尤其是序列长度方差大的时候,建议把数据按长度排序后动态batching,或者用padding-free的attention实现,能省下20%-30%显存。还有个小技巧是检查一下你的forward里有没有无意中创建了大的中间变量,比如某些激活函数没inplace,或者reshape导致copy了tensor。
如果ZeRO还没试的话,DeepSpeed stage2配合CPU offload其实挺稳的,比stage3省心,显存占用能再降一截。最后提醒下,gradient checkpointing要和batch size、微步数配合好,有时候开太多反而导致激活重算频繁,显存峰值没降但时间翻倍。我最后是bf16+动态padding+ZeRO2解决的,峰值显存稳定在32G左右,你可以照着这个思路排查下。
fp16震荡大概率不是混合精度本身的锅,先检查一下loss scaling策略,或者试试bf16,A100对bf16支持很好,基本无损。7B模型40G卡理论上是能跑的,问题多半出在序列长度上,padding token确实会浪费显存,建议用动态padding或者把max_length再砍一刀。另外你开了gradient checkpointing但forward到中间层爆掉,可以看看是不是activation checkpointing的粒度没设对,配合torch.utils.checkpoint的chunk_size调一调。实在不行就上ZeRO stage 2,A100单卡用offload优化器状态也能腾出不少空间。