最近开始尝试用Llama 3.1 8B做领域微调,跟着教程写了LoRA,batch size设到1,gradient checkpointing也开了,结果3090(24G)还是OOM。我的输入长度大概2k tokens,是不是跟序列长度有关?看到有人说用DeepSpeed ZeRO-3或者Flash Attention能省显存,但配置起来有点复杂,不太确定是哪里出了问题。另外,torch.compile会有帮助吗?现在用的是PyTorch 2.1,cuda 12.1。求有经验的大佬指点下常见坑,谢谢!
刚转大模型方向,用PyTorch跑LLM微调,显存总爆掉怎么优化?
全部回复
共 8 条你这配置3090跑8B LoRA按理说不会直接爆,2k长度确实有点吃紧,但更可能是attention计算时中间激活值太大。Flash Attention能省不少显存,而且现在PyTorch 2.2以上直接集成,升级一下就自带支持了,不用额外折腾。ZeRO-3主要是把优化器状态分到多卡或CPU上,单卡其实效果有限,不如先试试gradient checkpointing配合Flash Attention,batch size保持1但调低micro batch size到4或者8,看看能不能跑通。torch.compile对显存优化帮助不大,主要提速,建议先把显存问题解决了再折腾这个。
8B模型就算LoRA+gradient checkpointing,2K长度在24G上确实很极限,建议先把序列长度砍到1K或者512看能不能跑通。DeepSpeed ZeRO-3和Flash Attention确实能救,但ZeRO-3配LoRA有时会有参数同步问题,可以试试ZeRO-2或者用bitsandbytes的4bit量化,简单粗暴。torch.compile对LLM微调提升有限,不用急着折腾,先把显存降下来再说。
24G跑8B LoRA微调确实挺极限的,你提到的几个方向都是正解。序列长度2k token对于Llama 3.1来说已经是显存大头了,Flash Attention能有效降低attention层的显存占用,建议优先配起来,实测能省10%-20%。DeepSpeed ZeRO-3主要是切分优化器状态和梯度,对单卡来说收益其实没多卡那么明显,而且配置起来确实容易踩坑,新手可以先试试ZeRO-2或者干脆不用。torch.compile在微调场景下对显存优化帮助不大,主要是加速计算,反而可能因为编译过程增加临时显存占用。另外一个小技巧:检查一下你的LoRA rank是不是设太高了,8或16通常就够用,ranking偏大也会让adapter权重占不少空间。还有,可以尝试把输入截断到1.5k左右看能不能跑通,很多领域任务其实不需要全部2k context。最后确认一下你用的是不是最新版bitsandbytes,4bit量化配合NF4能把模型压到6-7G,配合LoRA基本是2408的标配玩法了。
24G跑8B LoRA微调确实有点极限,但2k长度下3090爆显存应该不只是序列长度的问题,建议先检查下是否忘了关gradient checkpointing的某些子模块,比如注意一下enable_reentrant=False这个参数,有时候默认设置反而会增加显存占用。Flash Attention对长序列提升很明显,而且现在transformers库直接集成好了,from_pretrained时加attn_implementation="flash_attention_2"就行,不用自己配,能省出4-5G。DeepSpeed ZeRO-3的话,如果没用多卡,单卡意义不大,反而可能因为offload拖慢速度,不如试试ZeRO-2或者直接用bitsandbytes的4bit量化QLoRA,我实测24G跑8B 4bit量化加gradient checkpointing能塞下4k长度的batch size 2。torch.compile在训练场景收益不稳定,尤其对长序列可能编译时间太长,建议先搞定显存再考虑。另外可以看看是不是dataloader的pin_memory=True和num_workers开太高导致缓存占显存,改成pin_memory=False偶尔会有奇效。
建议试试bitsandbytes 4bit量化,能直接把显存压到12G以下,配合Flash Attention效果更好。
24G跑8B LoRA确实有点极限,2k长度加上gradient checkpointing还爆的话,大概率是输入padding或者attention计算本身吃掉了大量临时显存。试试torch.compile?它在PyTorch 2.1上对Llama这类模型效果挺明显的,我自己的经验是能省10%-15%显存,而且编译一次后面跑起来很快。Flash Attention也值得搞一下,虽然配置确实有点绕,但装上后长序列的显存占用直接砍半,你搜一下xformers或者flash-attn的官方安装指南,跟着走一遍其实没那么复杂。DeepSpeed ZeRO-3倒不一定非得用,对单卡场景来说配置成本太高,有时候反而因为offload引入额外开销。另外检查一下你的tokenizer是不是把所有样本都pad到2k了,如果实际长度差异大,用动态padding或者packing技巧能省不少。还有个容易被忽略的点:优化器状态占显存,试试AdamW的8-bit版本或者干脆用SGD+合理的warmup,有时候能省出几G来。
我最近也踩过类似的坑,8B模型2k长度确实很容易爆显存。Flash Attention是必装的,能省20%左右显存,而且配置不复杂,pypi直接装就行。ZeRO-3的话,如果单卡用效果不明显,反而是offload到CPU能解燃眉之急。torch.compile建议先别开,有时反而会多占显存,等调通流程再说。另外可以检查下是不是输入padding太长,把pad到固定长度改成动态batch能省不少。
Flash Attention基本是必开的,能省不少显存,torch.compile对训练收益不大可以先不管。