最近在公司做一个小项目,用PyTorch微调一个开源的7B模型,单卡A100 40G居然跑着跑着就OOM了。我用的batch size已经降到1了,gradient checkpointing也开了,但还是在forward到中间几层的时候显存爆掉。看了一些教程说可以用DeepSpeed ZeRO或者混合精度,但试了fp16反而loss震荡得很厉害。我怀疑是不是自己dataloader里塞了太多padding token?还是说7B模型本来就需要80G以上的卡?有没有大佬指点一下,或者分享一下你们微调小模型时常用的显存优化trick?先谢过了!
请教大佬们,用PyTorch训练时显存总爆,是我代码写太烂了吗?
全部回复
共 9 条这问题我太熟了,7B模型在A100 40G上跑微调,batch size=1还爆显存,大概率不是卡的问题,而是数据形状和模型架构的细节没抠到位。
首先你提到padding token太多,这个确实是个常见坑。LLM为了对齐长度,很多人在dataloader里直接往长文本上补padding,但注意力机制会把这些padding也算进去计算,虽然结果会被mask掉,但它们依然占着显存做矩阵运算。你试一下把tokenizer的padding改成max_length或者直接动态batch,按最长的样本截断,不要统一补到某个固定长度,能省不少。另外检查下attention_mask是不是传对了,有些库的实现默认不传mask,模型会傻乎乎地算全序列。
再说fp16震荡这事,7B模型直接用原生的automatic mixed precision确实容易炸loss,尤其是微调阶段,梯度变化大,fp16的动态范围不够。你可以试一下bf16,如果卡支持的话,A100对bf16是原生支持的,稳定性比fp16好很多,而且显存用量跟fp16一样。如果只能用fp16,那得调一下loss scaling的策略,或者干脆先跑一个warmup阶段让模型适应一下。
还有一点容易被忽略:检查下你用的peft库是不是默认把全参梯度都保留了。像LoRA虽然只训练部分参数,但如果不手动设置requires_grad=False,PyTorch还是会给所有参数分配梯度的buffer,那样显存根本省不下来。建议用torch.no_grad()或者peft的disable_adapter模式确认下。
最后说句实在话,7B模型在40G上做全量微调确实勉强,但做LoRA或Q-LoRA的话是够的。你要不是必须用全参,试试4bit量化加LoRA,bitsandbytes库支持得很好,batch size能拉到4-8,训练速度反而更快。
7B模型在A100 40G上单卡微调确实吃紧,尤其你开了gradient checkpointing还OOM,问题大概率出在序列长度上——检查下dataloader里padding后实际最大长度,很多开源代码默认塞到2048甚至4096,光attention的中间激活就能吃掉十几G。fp16震荡的话,试试bf16或者给loss加个动态缩放,另外ZeRO-2配合activation offloading也能再省几个G,但记得把offload device设为cpu。
fp16 loss震荡大概率是loss scale策略没调好,试试bf16,A100原生支持,比fp16稳很多。7B模型40G其实能跑,但主要瓶颈在activations,你开了gradient ch
eckpointing还爆,可能是dataloader里padding太多导致序列过长,检查下每个batch的实际token数,用动态padding或batch sampler按长度分组能省不少显存。
同感,这问题太真实了,7B模型在40G卡上跑确实得精打细算。我先说一个可能被忽略的点:你提到dataloader里padding token太多,这个其实影响挺大的,尤其是如果序列长度参差不齐,你用了collate时统一padding到最长,那显存里会塞满无效的padding。建议试试PyTorch的pack_padded_sequence或者干脆用dynamic padding,每次只padd到当前batch的最大长度,能省不少。
另外fp16震荡的问题,我猜你可能没做gradient scaling?或者loss本身太大导致下溢?可以试试bf16,A100支持bf16,它动态范围比fp16大很多,震荡会小很多,而且很多7B模型微调社区已经默认用bf16了。
至于gradient checkpointing,你确认是开在了每个transformer层上吗?有些实现只在特定模块上开,中间层爆了说明可能没覆盖全。还有一个trick:把模型参数和优化器状态offload到CPU,用DeepSpeed ZeRO-3结合CPU offload,虽然会慢一点,但40G跑7B应该稳的。我上次调一个6.7B模型,batch size 1,开启ZeRO-3 + cpu_offload + bf16,显存峰值大概在32G左右,留了一些余量。
对了,你用的什么微调方法?如果是全参数微调,那确实对显存要求高,可以考虑LoRA或者QLoRA,4bit量化后7B模型在24G卡上都能跑,而且效果在不少任务上不输全参。可以试试bitsandbytes的4bit配置,配合gradient checkpointing,显存能压到20G以内。
fp16震荡的话试试bf16,另外看看tokenizer是不是没设max_length,padding太多确实会吃显存。
fp16 loss震荡可以试试bf16,很多7B模型对fp16敏感,另外检查下padding长度,建议动态batch。
fp16震荡可以试试bf16,另外检查下tokenizer的max_length是不是设得太大了。
fp16 loss震荡可以试试bf16,效果稳得多。另外检查下attention的padding mask是不是写对了。
fp16 loss震荡可以试试bf16,另外检查下tokenizer是不是把padding截断设好了。