最近在试着用LoRA微调Llama 3 8B,设备是RTX 4090 24G,数据集大概5万条对话。结果刚跑第一轮batch size设成4就直接OOM了,试了梯度累积和混合精度(bf16),还是撑不住。
我看网上有人说用bitsandbytes量化到4bit能省显存,但我加载模型后生成文本变慢了好多,而且微调完效果有点飘。
想问问大家,除了换硬件,有没有什么实用的显存优化技巧?比如分片加载、offload或者改attention?另外,8B模型用4bit微调会不会损失太多能力?求指个路,有点迷茫。
大佬们,用PyTorch跑Llama 3微调,显存爆了怎么办?
全部回复
共 161 条24G跑8B LoRA其实够用,batch size先调成1试试,配合gradient accumulation慢慢往上加,别一上来就4。bitsandbytes 4bit微调确实会掉点,尤其对话数据多了容易飘,建议先用NF4量化试试,比普通4bit稳一些。分片加载其实不太解决训练显存,offload到CPU或者用Flash Attention倒是能省不少,尤其是长序列场景。至于能力损失,8B模型4bit微调后推理速度慢主要是量化反量化开销,训练时影响小点,效果飘的话可以把学习率调低点再跑几个epoch看看。
我最近也在折腾类似的事,4090跑8B确实紧巴巴的。你试过把batch size降到1再加梯度累积吗?我这么搞过一次,虽然慢点但至少不OOM了。4bit微调我用了感觉效果还行,但生成速度确实肉疼,可能是量化后推理时反量化开销大。另外你可以看看PEFT的LoRA配置,把rank调低点,比如8,显存能省不少。
我踩过坑的是,别光盯着显存,数据加载那边也能省,比如把数据集预处理成tokenized的存盘,读的时候直接load,能省不少临时显存。你有试过用torch.compile吗?我开了之后显存占用没降,但速度上来了,可能间接帮了点忙。至于4bit会不会损失能力,我觉得看任务,简单指令微调影响不大,但要是复杂推理确实会飘。
5万条对话这规模,batch size 4确实有点猛了,我上次3万条用2都差点爆,你试试把seq len限制在1024以内,能省不少。4bit微调8B其实挺常见的,效果飘可能不是量化的问题,是学习率没调好,LoRA的alpha和r也可以再压一压。我倒是好奇你用的什么优化器,AdamW的话换8bit版能再抠出几个G。
24G跑8B LoRA其实挺紧的,但你这配置不是没救。试试把batch size降到1,配合8倍梯度累积,效果跟batch size 4差不多,显存能省一大截。4bit量化确实会让推理变慢,但微调时用QLoRA的话,我体感影响不大,效果飘可能是学习率没调好,建议降到1e-4左右再试试。另外可以把attention改成flash attention,显存能再省个20%左右。至于能力损失,8B模型4bit微调后任务表现一般掉不了太多,但如果你用的是对话数据,可能得注意下数据质量,有时候是数据本身的问题不是量化的问题。
24G跑8B LoRA其实挺极限的,batch size 4爆掉正常,我一般直接设1加梯度累积,效果不比大batch差。4bit微调确实会掉点,尤其对话数据多的时候,建议试试8bit加paged optimizers,显存能省不少而且稳一些。你那个生成变慢可能是量化后推理没开融合算子,换个加载方式或者用vLLM跑会好很多。另外可以看看unsloth的优化版LoRA,显存占用能再压一截,速度还快。
说实话你这配置已经算不错了,24G跑8B LoRA按道理不该这么憋屈,问题多半出在数据集预处理和batch size的匹配上。5万条对话如果没做padding统一或者长度截断,显存会被无效token白白吃光,建议先看看序列长度分布,把超过2048的样本直接砍掉或者用动态padding,效果立竿见影。
4bit量化确实会掉点,尤其LoRA本身可训练参数就少,你再把底座精度压那么狠,微调出来的模型飘是正常现象。我试过用NF4配合双阶段训练,先冻结量化层只训adaptor,再解冻部分层做p-tuning,能稍微缓解,但推理速度确实没办法,bitsandbytes的dequantize开销就在那。
分片加载和offload是能救急,但4090的PCIe带宽会成为新瓶颈,训练速度会掉到蜗牛爬。我更推荐你试试gradient checkpointing加上更激进的梯度累积,比如batch size 1配8步累积,这样等效batch size还是4,但峰值显存能砍一半。另外attention那块可以换成flash-attention-2,PyTorch 2.2以上直接调,能省不少显存而且不损失精度。
最后问个关键问题,你用的LoRA rank和alpha设的多大?很多人默认16或者32,但8B模型用8到12就够,再高纯属浪费显存。如果改完这些还爆,那可能得考虑用QLoRA的官方实现,它内部做了不少显存碎片整理的优化,比自己拼装组件稳定多了。
24G跑8B LoRA按理说应该够,但5万条对话数据确实有点猛,你batch size=4爆掉大概率是序列长度太长,很多对话数据集平均token数动不动就上千,先看看你的max_seq_len设了多少,如果超过2048建议直接截断或者用packing,能省不少。4bit微调效果飘不一定是量化的问题,LoRA rank和alpha调过没,8B模型用QLoRA其实很成熟,但你要注意把4bit的base model冻结,只训adapters,另外生成变慢是因为bitsandbytes的dequantize有开销,推理时可以用vLLM或者把adapter merge回fp16再导出,速度就正常了。显存优化这块,你可以试试torch.compile加FlexAttention,或者干脆用Unsloth那个优化过的LoRA实现,我实测同样设置下能省30%到40%显存,训练还快一截。至于分片和offload,CPU offload在4090这种卡上反而会因为PCIe瓶颈拖慢速度,不建议,除非你数据特别长。还有一个骚操作是把optimizer换成Adafactor,省一大块显存,但收敛速度要自己调。最后,8B模型用4bit微调能力损失真的不大,尤其LoRA本身参数就少,关键是学习率和epoch别乱来,5万条数据跑2个epoch基本够了,多了容易过拟合。你先试试把max_seq_len砍到1024加Unsloth,大概率直接起飞。
说到4bit微调效果飘,我猜你八成是直接用了QLoRA默认配置,但没调target modules。Llama 3的attention和mlp层全量注入LoRA其实很吃显存,试试只选q_proj和v_proj,能把激活内存砍掉快一半。另外你那个5万条对话,如果单条长度超过1k token,建议先按长度分桶再动态padding,不然短样本也会被长样本的padding拖爆显存。
关于offload,我建议别用CPU offload,虽然能跑但慢到怀疑人生。你可以试试把optimizer state用bitsandbytes的8bit版,配合梯度检查点,24G应该能塞下batch size 2加梯度累积16。生成变慢是正常的,4bit的decode阶段要反量化,微调时多用几个epoch让模型适应低精度分布,效果会稳一些。
其实8B用4bit微调损失没想象中大,重点看你的任务——如果是对话能力,反而能带来正则化效果,但要是做数学或代码,那掉点就明显了。你换个思路:先用4bit跑通流程,最后两轮再切回8bit精调,显存峰值只出现在切换瞬间,效果能拉回不少。另外注意torch.compile加flash attention,这俩配合能省30%左右显存。
24G跑8B LoRA应该不至于这么惨吧,你试试把batch size降到1然后梯度累积开大点,5万条数据用8步累积效果差不多,另外检查下是不是序列长度太长,把max_len截到1024能省一大块。4bit微调确实会让效果飘,我试过用QLoRA+adapter稍微好点,但推理速度慢是硬伤,建议还是保持bf16,把attention切成FlashAttention-2能省不少显存,而且速度还能提。你要是实在想上大batch,可以考虑把优化器状态offload到CPU,虽然慢点但能撑住。
试试unsloth吧,LoRA加4bit量化能压到12G以内,速度还比bitsandbytes快不少。
微调8B用4bit损失其实可接受,重点是把lora rank调低点,别让adapter学太飘。
试试unsloth优化+8bit LoRA,显存能砍一半,速度还快,4bit确实容易掉点。
你batch size压到2配合梯度累积,再把序列长度截到512,24G应该能跑。
4bit微调掉点正常,试试QLoRA加paged optimizers,24G跑8B够用,别上4 batch。
4090跑8B用unsloth优化下,显存直接砍半,速度和效果都能保住。
24G跑8B LoRA按理说够用,你batch size 4爆显存大概率是序列长度太长或者数据集里样本长度差异大,试试把max_seq_len砍到1024或者512,配合梯度累积到32步,效果不会差太多。4bit微调确实会掉点,尤其对话任务,建议换成8bit加NF4,显存压力小很多,生成速度也没那么崩。另外你如果只是做LoRA,别用bitsandbytes加载整个模型,直接用PEFT的prepare_model_for_kbit_training,然后把attention改成flash attention,能省不少。你数据集5万条,其实可以先在小样本上跑通流程,再全量跑,别一上来就追求batch size。
4bit微调确实伤能力,试试QLoRA+梯度检查点,24G跑8B把batch压到1问题不大。
5万条对话单卡4090确实极限了,我建议你把batch size压到1,配合gradient checkpointing和paged adamw,再把序列长度截到512试一下。4bit微调掉点主要是量化噪声叠加LoRA秩太低,把lora r提到32或者用nf4+双重量化会稳一些。实在不行就只冻结前几层,或者用Unsloth优化过的内核,能再挤出一半显存。生成变慢可能是bitsandbytes没走4bit推理加速,试试加载时把torch_dtype设成float16。
试试把batch size压到1,配合梯度累积把有效batch撑回来,4090跑8B LoRA其实是够的,关键是把序列长度也砍一砍,5万条对话里肯定有长文本,截断到1024能省一大截。4bit微调掉点确实明显,我上次用NF4跑完BLEU掉了快2个点,生成慢是因为bitsandbytes的反量化开销,建议换GPTQ或者直接用QLoRA的官方实现,能好很多。另外你试试torch.compile加flash attention,我这边显存直接少了3G,速度还快了。
24G跑8B LoRA按理说能挤一挤,但5万条对话确实有点猛,batch size 4直接爆很正常。你试过把batch size降到1,然后靠梯度累积把有效batch撑到16或32吗?这个组合比单纯调大batch更吃显存友好,而且收敛稳定不少。bitsandbytes的4bit确实会拖慢推理,因为反量化有开销,微调效果飘也常见,主要是量化误差在梯度回传时被放大,建议你试试NF4加double quant,或者干脆用8bit,速度和精度平衡会好很多。分片加载和offload其实对单卡帮助有限,更实际的是开torch.compile,能省不少显存,虽然编译时间有点肉疼。至于attention,可以把flash attention打开,省显存的同时还提速,但记得检查你用的transformers版本支不支持。8B模型用4bit微调,如果任务不是特别复杂,能力损失其实可控,但要是对话质量要求高,建议还是回bf16,把序列长度截断到512或768。另外你数据集5万条,可以试试packing策略,把短样本拼一起,能有效提高吞吐。最后别迷信某个单一技巧,组合拳才是王道,先砍序列长度,再上梯度累积,最后再考虑量化。
24G跑8B LoRA按理说够呛但也不是完全没救,batch size 1加梯度累积是底线,先把seq_len砍到512试试,数据截断比硬撑长上下文划算多了。4bit微调确实会掉点,尤其对话任务,建议你用QLoRA的nf4加上lora的rank调到16,效果比基础4bit稳不少。另外可以试试torch.compile加flash attention,能省不少显存,速度还快。我上次用同样配置跑7B,batch size 2加offload到CPU,勉强能过,但训练慢得想砸电脑,你5万条数据得多等几天。
4090 24G跑8B LoRA确实紧巴,但你这个配置不该第一轮就炸,大概率是数据集预处理时padding策略没弄对,把序列长度撑爆了。试试把max_seq_len砍到1024,同时用unsloth那个优化版LoRA,显存能再省一半。4bit微调8B其实可行,但别用bitsandbytes默认的nf4,换成fp4配合lora的target modules只改q/k/v/o,效果会稳定不少,生成变慢是因为4bit反量化有开销,推理时换回fp16权重就好。我建议你先从batch size=1+梯度累积16起步,把显存占用日志打出来看看瓶颈到底在哪,5万条数据其实用QLoRA跑两三个epoch就够,别追求大batch。
24G跑8B LoRA按理说够用,你batch size降到2或者1,配合梯度累积到8,应该能稳住。4bit量化确实会掉点,微调时建议只量化基座,LoRA层保持bf16,效果会稳一些。另外试试torch.compile和flash-attention,显存能再省一截,速度也快。至于能力损失,8B模型4bit微调做任务够用,但生成质量会略降,实测比fp16差几个点,能接受就行。