最近在试着用LoRA微调一个7B的Llama模型,任务挺简单的,就是给客服对话做情感分类。我用的是一张4090(24G显存),batch size设到4就直接OOM了,换成1倒是能跑,但训练慢得离谱,而且loss下降特别不稳定。已经试过gradient checkpointing和混合精度训练,感觉效果有限。看网上有人说可以用DeepSpeed ZeRO或者量化到4bit,但不知道怎么配置才不崩。想问问有经验的老哥,显存不够的情况下,怎么平衡速度和效果?另外,有没有推荐的库或者trick,能让我这张卡撑住微调?先谢过了!
自己用PyTorch微调Llama,显存总爆掉,有什么省钱又实用的技巧吗?
全部回复
共 157 条说实话你这个问题太典型了,24G跑7B LoRA其实完全能玩,但batch size=4确实有点贪。我建议你先检查一下是不是把LoRA的target modules设太多了,一般只改q_proj和v_proj就能省不少显存,而且效果不会差太多。另外你说loss不稳定,我怀疑是学习率太高了,LoRA微调用1e-4到3e-4之间比较稳,你可以试试warmup steps拉长一点,比如总步数的5%。
至于DeepSpeed ZeRO,我个人觉得单卡上意义不大,ZeRO-2的offload会把CPU内存拖进来,速度反而更慢,除非你愿意忍受每步多几倍的时间。4bit量化倒是真香,用bitsandbytes的NF4配置,加上双重量化,显存能压到10G以内,这样batch size=8甚至16都没问题,而且你只是做情感分类这种简单任务,精度损失基本可以忽略。不过我踩过一个坑,4bit下如果模型层数太深,反传的时候梯度会有点抖动,建议你把LoRA的dropout设成0.1,然后跑几个epoch看下验证集准确率,别光看训练loss。
还有个野路子,你可以用PEFT的gradient accumulation配合batch size=1,等效batch size=8,但注意每accumulate几步就手动清一下缓存,不然碎片会越来越多。最后推荐你试试Unsloth这个库,它专门优化了Llama的显存占用和速度,我的经验是同样配置下能省30%显存,训练速度还快20%,而且是pip直接装,配置特别简单。你先把batch size降到2,加上4bit和gradient accumulation,大概率能稳住了。
说实话你这个问题我太有共鸣了,之前我拿3090调7B的时候也是被OOM折磨到怀疑人生。你试的那几个方法我都踩过坑,gradient checkpointing加混合精度其实只算基础操作,真正的瓶颈反而在优化器状态和中间激活值上。我后来换了个思路,直接用bitsandbytes的4bit量化加载模型,配合peft库的LoRA,batch size直接拉回8,loss曲线瞬间就稳了,训练速度比之前1batch还快不少。不过要注意4bit下最好用nf4数据类型,还有记得把梯度的dtype也设成float32,不然会卡在反传那步。另外DeepSpeed ZeRO其实不太适合单卡场景,它主要是多卡省显存,单卡上反而会引入额外通信开销,不如把精力放在PagedAdamW这种优化器上,它能自动把优化器状态换到CPU内存里。你要是实在不想碰量化,还有个野路子是把输入序列截断到256token以内,客服对话情感分类其实用不到太长上下文,显存占用直接砍一半。对了,你loss不稳定可能跟学习率有关,LoRA微调一般建议用cosine衰减,初始lr设到2e-4左右,比默认的线性要稳很多。最后提醒下,记得关掉eval时的梯度计算,有时候验证阶段爆显存就是这玩意儿在偷内存。
4090跑7B LoRA其实不用上DeepSpeed,bitsandbytes的4bit量化加QLoRA是首选,显存能压到10G以内,batch size直接拉到8都没问题。你loss不稳大概率是学习率太高,试试降到1e-4以下,另外用paged_adamw优化器能省不少显存碎片。如果还想提速,可以把序列长度截断到256,客服对话一般用不着长上下文,这样显存占用还能再降一截。对了,gradient checkpointing开着没问题,但别跟量化一起用,有时反而会拖慢速度。
4090跑7B LoRA其实挺尴尬的,24G看着不小但真用起来处处是瓶颈。你试的gradient checkpointing和混合精度属于基础操作,想再压显存可以看看QLoRA,4bit NF4量化配合双卡(哪怕一张垃圾卡做offload)能省不少,不过要注意量化后loss曲线会毛躁一点,建议把学习率调低些,比如2e-4以下,收敛会稳很多。另外你提到batch size=1 loss不稳定,很可能是没开梯度累积,试试累积8步再更新参数,效果等同于batch size=8但显存占用差不多。DeepSpeed ZeRO其实不太适合单卡场景,它的offload开销反而拖慢速度,不如直接用PEFT库的LoRA实现,配置里把lora_r设成16,lora_alpha设32,target_modules选q_proj和v_proj就行。还有个小技巧,把输入序列截断到128或256,客服对话一般不会太长,这能直接砍掉一大块显存。最后实在不行就换Freeze(只训练最后几层)或者用GPTQ量化后的4bit模型,精度损失对情感分类任务基本无感。你试过用bitsandbytes的8bit优化器吗?那个对显存也有帮助,但记得要把optimizer的betas调低一点。
说实话你这情况我太懂了,24G卡跑7B LoRA确实紧巴巴。我建议直接上QLoRA,把基座模型量化到4bit,实测能把batch size提到8甚至16,而且loss曲线比1的时候稳多了。另外你试试把优化器换成AdamW 8bit,PagedOptimizer那个功能能省不少显存峰值,配合gradient checkpointing基本能撑住。还有个小窍门,把序列长度截断到512以内,对情感分类这种任务完全够用,显存又能省一截。
4bit量化加LoRA,batch开1也能稳,速度慢就换ZeRO stage2,别硬刚batch size。
试试QLoRA吧,7B在24G上跑4bit很稳,loss不稳就把学习率降到1e-4以下。
这题我熟,之前用3090调7B也卡得欲仙欲死。你4090其实有条件上4bitQLoRA,bitsandbytes配peft的prepare_model_for_kbit_training,量化后显存占用能砍一半多,batch size拉回4问题不大。另外建议试试ZeRO stage 2,配合offload优化器状态,比gredient checkpointing管用多了。还有个小trick,把序列长度截到256,客服对话一般够用,速度能快不少。
24G跑7B LoRA其实挺宽裕的,你batch size设4爆掉大概率是序列长度或者attention缓存没控制好,试试把max length砍到512,再配合gradient accumulation凑等效batch。4bit量化建议直接上bitsandbytes,加载时设load_in_4bit=True,配合peft的LoRA,显存能压到10G以内,速度反而比bf16快。至于loss不稳,多半是学习率太高,LoRA用1e-4左右,warmup几步看看。另外别迷信DeepSpeed,单卡上ZeRO反而拖慢速度,不如把optimizer换成adamw8bit省显存。
这题我熟,4090跑7B LoRA其实没必要硬刚ZeRO,先试试把batch size固定成1然后梯度累积开8步,效果跟batch 4基本一样,loss稳很多。4bit量化用bitsandbytes加peft的prepare_model_for_kbit_training就行,配置里记得把bnb_4bit_quant_type设成nf4,能省一半显存还不太掉点。另外你loss不稳定可能是学习率太高,LoRA一般1e-4到2e-4就够,别跟全参微调一样上5e-5。还有个野路子,把序列长度截到256,客服对话一般没那么长,显存直接砍半,速度还能翻倍。
24G跑7B LoRA其实挺够用的,关键问题可能出在数据加载和优化器状态上。你试试把batch size固定成1,然后梯度累积设到8,这样等效batch size还是4,但显存峰值会低很多,loss曲线也会稳一些。另外你用的LoRA rank是多少?如果设得太大(比如64以上),中间激活值照样吃显存,降到8或者16试试,任务简单的话效果不会差太多。4bit量化的话推荐bitsandbytes的NF4配置,配合peft库的prepare_model_for_kbit_training,基本能再把显存砍半,但注意量化后学习率要调小一点,不然容易震荡。还有个野路子是冻结embedding和lm_head之外的层,只训attention层,虽然会损失点精度,但速度能快不少。另外检查下你的数据加载器是不是num_workers设太高了,有时候CPU内存爆了也会反过来拖慢显存释放。最后,如果实在嫌麻烦,可以试试Unsloth这个库,它对Llama的LoRA优化做得特别狠,号称能省70%显存,我上次用它跑7B训练时batch size能开到8都不爆。
试试QLoRA加4bit量化,batch size开2照样稳,loss波动大就把学习率降到1e-4。
试过QLoRA没?4bit加双卡量化,24G跑7B绰绰有余,速度还比纯LoRA快一截。
其实你这个问题我踩过一模一样的坑,4090跑7B LoRA,batch size 4不OOM才怪,我最后是batch size 1加梯度累积(设个8)才稳住,loss曲线也好看了很多。量化4bit的话可以试试bitsandbytes的NF4配置,配合PEFT的LoRA,显存能压到12G左右,速度损失其实能接受。另外DeepSpeed ZeRO 2在这张卡上有点鸡肋,不如直接开offload optimizer到CPU,省下的显存全给batch size。对了,你数据量小的话,可以试试把序列长度截到256,很多情感分类任务根本不需要长上下文,效果影响不大。
试试QLoRA配4bit量化,显存直接砍半,把batch提到8没问题。或者换Unsloth,同样效果下省显存还快不少。
说实话你这个配置跑7B LoRA确实有点极限,但24G不至于这么惨。试试peft的4bit量化加载,配合bitsandbytes的nf4配置,显存能压到10G左右,batch size开8都没问题。另外别死磕AdamW,换Adafactor或者8bit优化器,loss稳定很多。我上次就是这么把8B模型塞进4090的,速度还挺惊喜。
4090跑7B LoRA其实不用上ZeRO,直接上QLoRA就行,4bit量化加nf4配置,batch size 8都没问题。你loss不稳大概率是学习率太高,试试1e-4到3e-4区间,再加个warmup。另外paged optimizer会偷显存,记得开一下,能省不少。
我上次用同样配置跑金融情感分类,序列长度裁到256,准确率只掉了不到1%,但训练速度快了快一倍。你那个客服对话要是太长,先看看长度分布,别让padding浪费显存。还有,别用transformers自带的trainer,用peft库的prepare_model_for_kbit_training,它自动处理梯度检查点,省心很多。
24G跑7B LoRA其实挺极限的,但你这个任务简单,试试把LoRA的rank降到8,target modules只选q和v,再配合4bit量化,显存能省出一大截。另外别用AdamW了,换Adafactor或者Lion,优化器状态占的显存能少一半,batch size提到4应该没问题。loss不稳大概率是学习率太高,调到1e-4左右配个warmup,会稳很多。还有个小技巧,把输入序列截断到256,客服对话一般用不到那么长,速度能翻倍。
直接上QLoRA,4bit量化加paged optimzer,24G跑7B绰绰有余,速度和显存都稳。
24G跑7B LoRA其实挺宽裕的,你OOM大概率是seq len太长或者batch accumulation没用起来。试试把max length砍到512,batch size设1然后梯度累积开8步,效果和batch 4差不多,loss也能稳下来。4bit量化的话推荐bitsandbytes加peft的prepare_model_for_kbit_training,配置上load_in_4bit加nf4类型基本不会崩,但记得把tokenizer的pad侧设对。另外DeepSpeed ZeRO2在这场景下其实不如直接开offload省心,真要用就把offload_optimizer打开,stage设2就行。
4bit量化加QLoRA,显存直接砍半,batch size提到8没问题,速度还快。
DeepSpeed ZeRO Stage 2配合offload,4090跑7B很稳,就是配置别照抄默认的。