最近在试着用LoRA微调一个7B的Llama模型,任务挺简单的,就是给客服对话做情感分类。我用的是一张4090(24G显存),batch size设到4就直接OOM了,换成1倒是能跑,但训练慢得离谱,而且loss下降特别不稳定。已经试过gradient checkpointing和混合精度训练,感觉效果有限。看网上有人说可以用DeepSpeed ZeRO或者量化到4bit,但不知道怎么配置才不崩。想问问有经验的老哥,显存不够的情况下,怎么平衡速度和效果?另外,有没有推荐的库或者trick,能让我这张卡撑住微调?先谢过了!
自己用PyTorch微调Llama,显存总爆掉,有什么省钱又实用的技巧吗?
全部回复
共 157 条老实说你这配置跑7B LoRA确实有点紧,但24G不该这么惨。试试把batch size压到1,同时开gradient accumulation到8,效果等同batch size 8,loss会稳很多。另外强烈建议上QLoRA,用bitsandbytes的4bit NF4量化,显存直接砍半,我记得装个peft库改几行配置就行。还有个野路子:把tokenizer的padding策略改成左侧填充,能减少计算浪费。最后checkpoint别存全量,只存adapter权重,省下的空间能多跑几个epoch。
你这情况我太熟了,24G跑7B LoRA其实不该这么憋屈。试试把batch size固定成1,然后用梯度累积到8或16,效果跟大batch差不多,loss也能稳住。另外强烈建议上bitsandbytes的4bit量化加载模型,再配个peft库的LoRA,显存能直接省掉一大半,速度反而可能更快。DeepSpeed ZeRO Stage 2可以开,但单卡上提升有限,优先把offload optimizer state打开就行。最后别忘用torch.compile,虽然启动慢点,但训练迭代能快不少。
4bit量化加QLoRA,7B在24G上batch能拉到8,试试peft的prepare_model_for_kbit_training。
你这情况跟我上周简直一模一样,4090跑7B LoRA,batch size=4必炸。后来我干脆把LoRA的r值降到8,同时只冻结embedding和lm_head之外的层,显存直接省了快3个G,batch size能开到2了。4bit量化其实没那么玄乎,用bitsandbytes配peft的prepare_model_for_kbit_training就行,记得加载时设load_in_4bit=True,但loss不稳定大概率是学习率太高,试着调到1e-4以下,另外把序列长度截到256,情感分类真用不着512。
24G跑7B LoRA其实完全够用,你试试把batch size固定到1,但梯度累积步数调到8,效果跟batch size 4差不多,loss也会稳很多。另外4bit量化用bitsandbytes加peft的prepare_model_for_kbit_training,配置上注意把llama的attention里那个softmax换成内存友好的实现,能省不少显存。DeepSpeed ZeRO2在这卡上反而容易出兼容问题,不如直接开flash attention,显存占用能降三分之一。我自己之前用同样的配置做多轮对话微调,峰值也就15G左右,你查下是不是tokenizer把长序列padding了,那个特别吃显存。
24G跑7B LoRA其实挺宽裕的,你OOM大概率是seq len太长或者attention没开flash attention。试试把max length砍到512,batch size用梯度累积到8,loss曲线能稳很多。4bit别直接上,先用bitsandbytes的NF4配合peft的prepare_model_for_kbit_training,记得关掉gradient checkpointing,不然速度会雪崩。DeepSpeed ZeRO2对单卡没啥用,ZeRO3反而会慢,不如把精力放在优化数据加载和paged optimizer上。另外检查下是不是把embedding和lm_head也设成可训练了,那两个参数量巨大,冻住能省不少显存。
24G跑7B LoRA其实不该这么憋屈,你batch size降到1还爆的话,八成是LoRA的target modules选太宽了,试试只微调q_proj和v_proj,参数量能砍掉一半多。混合精度你用的是bf16还是fp16?Llama对bf16更友好,loss不稳可能跟这个有关,另外建议把学习率降到1e-4以下,用cosine schedule带warmup,能明显缓解震荡。DeepSpeed ZeRO Offload确实能压显存,但配置坑很多,我建议你先试试PEFT库自带的4bit量化,用bitsandbytes的NF4格式,配合双卡(哪怕另一张是垃圾卡)做offload,比单卡硬扛稳得多。还有个野路子,把序列长度裁剪到128,客服对话一般没那么长,注意力显存直接指数下降。优化器换Adafactor,它比AdamW省显存,效果损失很小,配合gradient accumulation到8步,等效batch size能到8,速度反而比硬跑大batch快。最后别迷信batch size,7B模型用LoRA,batch size 8和16的收敛差异真的不大,你不如把时间花在调lr和warmup步数上。
24G跑7B LoRA其实挺宽裕的,问题多半出在你把batch size怼太高,加上没开gradient accumulation。建议batch size固定1,梯度累积设8步,这样等效batch size还是4,但显存压力小很多。另外你试过bitsandbytes的4bit量化加载模型吗?配合peft库的LoRA,基本能把激活内存砍掉一大半,loss不稳大概率是学习率太高,降到1e-4试试。DeepSpeed ZeRO2在单卡上没用,ZeRO3配合offload倒是能救,但配置麻烦还拖慢速度,不如先试试NF4量化加paged optimizer,我这么做之后4090上sequence length 512能跑到batch size 8。还有个野路子,把优化器换成Adafactor,省显存效果立竿见影,就是收敛得调一调。
我之前也卡在这上面,24G跑7B全参微调纯属折磨,你换LoRA方向是对的,但batch size4爆掉大概率是序列长度和attention缓存吃满了,试试把max length砍到512,客服对话一般用不了那么长。gradient checkpointing和混合精度你开了,但可以再检查下是不是忘了关gradient accumulation的同步开销,batch size1加accumulation steps8,效果和batch size4差不多,但loss不稳可以试试调高学习率warmup比例,或者用AdamW的epsilon改大点。4bit量化我推荐bitsandbytes的NF4,配合peft库的LoRA,直接把base model冻成4bit,只训adapter,显存能压到10G以内,速度反而比fp16快,因为内存带宽压力小,但注意量化后loss曲线会有点抖,别慌,训完eval看指标就行。DeepSpeed ZeRO Stage2配单卡意义不大,ZeRO主要省的是多卡通信冗余,单卡你直接开offload到CPU也行,但速度会掉到怀疑人生,不划算。还有个冷门trick,用torch.compile加reduce-overhead,能省点显存还提速,不过要PyTorch2.0以上,兼容性得试。最后推荐你看看Unsloth这个库,专门优化微调显存,同样是LoRA,它能比原生peft再省30%左右,而且不用改太多代码,我试过是真的稳。
24G跑7B LoRA其实挺宽裕的,你batch size 4就爆大概率是seq length太长或者优化器状态没省。建议先试试8bit AdamW,光这一项就能省下大半显存,再配合gradient accumulation把有效batch size堆到16或32,loss会稳很多。QLoRA的4bit NF4量化是另一个路子,但记得用unsloth或者peft的latest版本,配置上主要注意把target modules选对,别全量化,不然精度崩得快。我个人经验是,情感分类这种任务根本不需要全量微调,冻结全部attention层只调MLP或者干脆只调最后几层,效果差别不大但显存能再降一半。另外你loss不稳定可能不是batch size的锅,检查下学习率,LoRA通常得降到1e-4以下,配合warmup和线性衰减会好很多。DeepSpeed ZeRO在单卡上意义不大,除非你想开offload到CPU,但那个速度会慢到怀疑人生,不推荐。最后可以考虑用bitsandbytes的page_optimizer,配合梯度检查点,24G能塞下16的batch,训练速度比你现在batch 1快三倍不止。
你这需求用QLoRA基本就是标准答案了,4bit量化加NF4格式能把7B压到6G左右,batch size开8都稳。重点是把peft库的target_modules选成q_proj和v_proj,再加个8bit的AdamW优化器,显存能再省一截。另外loss不稳大概率是学习率太高,建议调到1e-4以下,配合warmup steps跑个几百步看看曲线。DeepSpeed ZeRO其实单卡没必要折腾,反而容易出配置问题,不如直接上unsloth,它自带flash attention和智能批处理,亲测速度能快30%。
说到24G跑7B微调,其实瓶颈不在模型本身,而在优化器状态和中间激活值。你batch size 4就爆,大概率是激活值没省到位——gradient checkpointing开了之后,建议把batch size调回4甚至8,因为显存占用大头从激活值变成了参数梯度,这时候你会发现同样显存能塞下更大batch。loss不稳定跟batch size太小直接相关,试试gradient accumulation,每步累积4个mini-batch再更新,等效batch size能到16,收敛会稳很多。
至于ZeRO和4bit,ZeRO对单卡其实帮助有限,它主要是省显存换通信,单卡上不如直接开offload optimizer到CPU,代价是慢30%左右但能腾出近半显存。4bit量化我推荐bitsandbytes的NF4,配QLoRA,几乎不损失精度,7B模型加LoRA权重总共才占6G多,你甚至可以试试8的batch size。还有个冷门trick:把输入序列截断到最长128个token,客服对话通常没那么长,显存立刻降一半,对分类任务影响微乎其微。
库的话,强烈建议transformers加peft,把LoRA的r设成8到16,target modules只选query和value,别全改。另外检查一下是不是把模型本身也设成了训练模式,冻结所有非LoRA参数能省不少优化器内存。如果还嫌慢,用torch.compile加max-autotune模式,虽然编译久点,但训练token吞吐能翻倍。最后,别迷信DeepSpeed,单卡场景下它反而容易引入配置问题,先试QLoRA加梯度累积,大概率能解决。
试试QLoRA配bitsandbytes,4bit量化后24G跑7B轻松,batch可以开到8。
试试QLoRA加bitsandbytes的4bit量化,再配合gradient accumulation,4090跑7B完全没问题。
LoRA微调7B在24G上确实挺吃紧的,batch size开到4基本不现实。你可以试试QLoRA,4bit量化加载底座模型,配上bitsandbytes和peft,显存能省一大半,我拿4090跑过类似配置,bs=2加grad accumulation基本稳。loss不稳可能跟学习率和warmup有关,LoRA一般1e-4到2e-4比较合适,warmup给个几十步。DeepSpeed ZeRO-2也能帮上忙,但单卡收益不如量化直接,配置还容易踩坑。
4090跑7B LoRA其实挺够用的,batch size上不去多半是序列长度和optimizer state吃显存。试试bitsandbytes的4bit量化加载base model,再配合paged_adamw_8bit优化器,显存能省一大截,batch size开到8都没问题。loss不稳的话可以调低学习率到1e-4左右,warmup多给点步数。DeepSpeed ZeRO-2对单卡提升不大,别折腾了,直接用peft加trl的SFTTrainer最省心。
试试QLoRA加bitsandbytes的4bit量化,我4090上batch size能拉到8还不炸,loss也稳多了。