最近在试着用LoRA微调一个7B的Llama模型,任务挺简单的,就是给客服对话做情感分类。我用的是一张4090(24G显存),batch size设到4就直接OOM了,换成1倒是能跑,但训练慢得离谱,而且loss下降特别不稳定。已经试过gradient checkpointing和混合精度训练,感觉效果有限。看网上有人说可以用DeepSpeed ZeRO或者量化到4bit,但不知道怎么配置才不崩。想问问有经验的老哥,显存不够的情况下,怎么平衡速度和效果?另外,有没有推荐的库或者trick,能让我这张卡撑住微调?先谢过了!
自己用PyTorch微调Llama,显存总爆掉,有什么省钱又实用的技巧吗?
全部回复
共 157 条你这个问题我太懂了,之前用3090调7B的时候也是疯狂OOM。建议直接上QLoRA,把4bit量化加上,显存占用能砍掉一大半,你那个batch size提到8都没问题。配置上记得用bitsandbytes的nf4类型,然后peft库里的LoRA参数把r设成8,alpha设16,基本就稳了。另外可以试试把序列长度限制在512以内,客服对话一般够用了,这样速度会快很多,loss稳定性和全精度比差不了太多。
试试4bit QLoRA + paged optimizer,7B在24G上batch size能到8,速度还比纯fp16快。
把序列长度截到256,情感分类用不着那么长,显存直接省一半。
4090跑7B LoRA其实瓶颈不在显存总量,而在激活值峰值,试试把batch size拆成梯度累积,配合ZeRO stage 2,能稳很多。4bit量化建议用bitsandbytes的NF4,配合peft的target_modules设置,比默认配置省一半显存,但loss震荡的话可以调低学习率到1e-4。另外检查下是不是把embedding和lm_head也冻结了,这两个不冻显存占用会高不少。
试下QLoRA+4bit,你这任务7B用不到全参,batch1加梯度累积照样稳,速度还快不少。
用bitsandbytes搞4bit量化,再配合paged optimizer,24G跑7B绰绰有余,loss不稳就把学习率调低点。
24G跑7B LoRA还爆显存,大概率是序列长度和batch size没搭配好,试试把max length砍到512,再把batch size调到2配合梯度累积,效果会好很多。4bit量化建议直接用bitsandbytes的NF4格式,配合peft库的LoRA配置,基本能省一半显存,而且loss波动会小很多。DeepSpeed ZeRO2在单卡上其实提升不大,不如把精力放在优化数据加载和调整学习率上,另外可以试试把优化器换成Adafactor,省显存效果很明显。
24G跑7B LoRA其实挺宽裕的,问题大概率出在数据加载和优化器状态上。你可以试试把optimizer换成AdamW 8bit,来自bitsandbytes,能直接省下好几G,另外用gradient accumulation把batch size凑到16左右,loss会稳很多。至于4bit量化,QLoRA那套配置其实很成熟,直接套用HuggingFace的示例脚本就行,记得把trust_remote_code打开。还有个容易被忽略的点,检查下是不是把验证集也塞进Dataloader了,之前我就犯过这错导致显存莫名暴涨。速度慢的话,开torch.compile能提个20%左右,代价是编译时间,但值得一试。
24G跑7B LoRA按理说够用了,你试试把batch size固定成1然后梯度累积开个8步,效果基本等同batch 4但显存压力小很多。还有那个loss不稳定,大概率是学习率太高,LoRA的lr调到1e-4以下会稳不少。4bit量化用bitsandbytes配peft库挺省心的,我上次8G显存都能跑13B,不过你既然有24G其实不太需要。真想上ZeRO的话,stage 2加offload optimizer到CPU就行,记得把zero_optimization的stage设成2,offload_optimizer的device改成cpu,其他默认。
24G跑7B LoRA还爆显存,大概率不是显存不够,是优化器状态和中间激活值没打理好。你可以试试把LoRA的target modules换成只调query和value,别碰全部线性层,参数量直接砍半,显存能省出一大截。关于batch size,1的时候loss抖大概率是学习率太高,调到1e-5以下再配个warmup,稳定很多。DeepSpeed ZeRO Stage 2搭配offload optimizer到CPU,确实能救急,但注意要开pin memory,不然数据搬运能把训练拖成PPT。4bit量化的话,推荐bitsandbytes的NF4配置,记得加载时设好bnb_4bit_use_double_quant和compute_dtype,一旦跑起来比FP16省一半显存。还有个冷门trick,把输入序列截断到256或者512,情感分类根本不需要长上下文,这比改任何训练配置都直接。最后,梯度累积设成8,等效batch size就是8,配合混合精度,速度不会比原来慢太多。你要是试完还是卡,可以考虑换PEFT库里的最新版,它对Llama的显存优化比你自己手搓的LoRA脚本好不少。
4090跑7B LoRA其实没必要硬上ZeRO,直接bitsandbytes的4bit量化加LoRA,显存能压到10G左右,batch size开到8都没问题。你loss不稳大概率是学习率太高,LoRA一般建议1e-4起步,再配合warmup steps会好很多。另外试试PEFT库的prepare_model_for_kbit_training,它会自动帮你处理好梯度检查点和量化相关的坑,比自己手动调省心多了。
试试QLoRA配4bit,显存直接砍半,batch开8没问题,loss不稳就把学习率调到1e-4以下。
4090跑7B LoRA其实不用硬上ZeRO,试试peft的target_modules全量替换成q_proj和v_proj,再配合bitsandbytes的4bit量化,batch size直接拉回8,loss曲线比你现在稳多了。另外记得把gradient_accumulation_steps调成4,效果等同大batch但显存压力小很多,我这么跑过单卡微调,速度大概能快三分之一。你那个loss不稳定八成是lr太高,降到1e-4配合warmup试试?
4bit量化加LoRA肯定够用,再把batch设成2配合梯度累积试试,速度损失没那么大。
4090跑7B其实不用硬上DeepSpeed,试下QLoRA加4bit的NF4量化,显存直接砍半,batch size能提到8左右,速度比你想的快。另外loss不稳大概率是学习率太高,调到1e-4以下,再加个warmup steps,会稳很多。你用的是peft库吗?那个带gradient checkpointing的配置,比你自己手写省事。
试下QLoRA配4bit,24G跑7B够用,batch开2加梯度累积,稳定得很。
你这情况我太熟了,24G跑7B LoRA按理说不该这么憋屈,batch size 4爆掉大概率是序列长度和attention缓存没控制好。试试把max length砍到512,再加个4bit的QLoRA,bitsandbytes配置一下基本能稳在batch size 8,速度反而会上去。loss不稳的话,建议把学习率调到1e-4以下,再用warmup跑几百步看看,别一上来就全量更新。另外PEFT库的prepare_model_for_kbit_training别忘了加,能省不少显存,还不用折腾DeepSpeed那些复杂配置。
试试QLoRA加4bit,24G跑7B绰绰有余,速度比纯LoRA快不少,loss不稳就调低学习率。
直接用bitsandbytes配peft,显存能压到10G以内,batch size开到8没问题。
说实话你这情况我太懂了,24G跑7B LoRA本来应该挺宽裕的,batch size 4爆掉大概率是没开gradient accumulation或者attention实现没优化。你试试把batch size固定到1,然后梯度累积步数设成8,效果等同于batch size 8,显存占用几乎不变,loss也会稳很多。另外你提到混合精度,建议别用fp16,直接用bf16,4090对bf16支持很好,数值稳定性比fp16强一大截,loss乱跳的问题可能就解决了。至于DeepSpeed ZeRO,单卡上ZeRO-2意义不大,ZeRO-3反而会拖慢速度,真不如去试QLoRA,4bit的NF4量化配合double quantization,显存能压到10G以内,你甚至可以开更大的batch。还有个小trick,把序列长度截断到256或者128,客服对话一般不会太长,这样显存直接省一半,速度还快。最后强烈推荐bitsandbytes的8bit优化器,配合paged AdamW,几乎不占额外显存,比官方AdamW省好几G。先按这个组合跑一轮看看,不行再上unsloth,那个库专门优化过微调流程,训练速度能再翻倍。
试试QLoRA配4bit,显存直接砍半,batch开8没问题,速度还比你现在快。
单卡4090玩7B就老老实实上bitsandbytes,别折腾DeepSpeed,配置麻烦收益还小。
你这配置跑7B LoRA确实极限,但24G不该这么惨。试试把batch size固定成1,然后用gradient accumulation把有效batch堆到16或32,loss会稳很多。另外强烈建议上bitsandbytes的4bit量化,配合peft的LoRA,显存直接砍半,速度反而可能更快。ZeRO Stage 3对单卡没啥用,别折腾了,重点是把attention的显存占用降下来,比如用torch.compile或者flash-attention,能省不少。最后,如果任务真的只是情感分类,其实可以换个更小的模型,比如Mistral 7B或Phi-3,效果差不多但省心多了。
试试QLoRA配4bit量化,24G跑7B够用,batch开2再加梯度累积,稳得很。