最近在试着用LoRA微调Llama 3 8B,设备是RTX 4090 24G,数据集大概5万条对话。结果刚跑第一轮batch size设成4就直接OOM了,试了梯度累积和混合精度(bf16),还是撑不住。
我看网上有人说用bitsandbytes量化到4bit能省显存,但我加载模型后生成文本变慢了好多,而且微调完效果有点飘。
想问问大家,除了换硬件,有没有什么实用的显存优化技巧?比如分片加载、offload或者改attention?另外,8B模型用4bit微调会不会损失太多能力?求指个路,有点迷茫。
大佬们,用PyTorch跑Llama 3微调,显存爆了怎么办?
全部回复
共 162 条说实话4090 24G跑8B LoRA确实紧巴,我建议你把batch size压到1,配合8bit的NF4量化试试,微调用QLoRA的话效果比4bit稳定不少,生成慢是因为bitsandbytes的dequantize开销大,正常现象。你那个5万条数据量其实不小,可以试下只冻结前几层,只训练后几层的LoRA,省显存效果明显。另外attention那块别瞎改,容易出诡异bug,优先考虑torch.compile加flash attention,能省个2-3G。4bit微调确实会掉点,但你要是任务不复杂,一般任务够用了。
试试unsloth吧,LoRA+4bit能省一半显存,25G跑8B稳得很,速度还快。
同款4090,我上次跑7B也是硬撑,后来把batch size压到1,配合8bit的bitsandbytes,再把序列长度截到512才勉强跑完。4bit微调确实会飘,尤其对话类数据,感觉量化误差被放大了。如果你一定要用LoRA,可以试试只看量化后的权重做前向,反向传播用原始权重,虽然慢点但稳。另外你5万条数据其实可以分阶段喂,先预训一小部分看loss曲线,别一上来就全量跑,省得白等。
24G跑8B LoRA其实卡在激活值上,batch4确实太顶了,试试gradient checkpointing加上batch1,配合8bit优化器(比如AdamW8bit),5万条数据用step累积效果差不多。4bit微调掉点主要在中文任务上,英文影响小,如果数据集偏口语可以试试把量化改成NF4加双量化,速度损失会小很多。offload到CPU只建议开optimizer和梯度,参数别动,不然慢到怀疑人生。另外看下是不是用了flash attention,没换的话显存能差出3-4G。
4090 24G跑8B LoRA确实紧巴,你这配置其实挺典型的。我试过把batch size压到1,配合8bit的QLoRA,再把序列长度截到1024,能稳定跑完5万条,但速度确实慢一半。4bit微调能力损失在LoRA场景下其实不大,飘的话可能跟学习率或者adapter rank有关,建议把lr调到1e-4以下试试。分片加载跟offload对训练帮助有限,主要省的是内存,显存瓶颈还得靠量化跟梯度检查点扛。你数据集里长对话多吗?我怀疑是attention的显存峰值在作怪,可以看看flash-attention装没装。
你的配置其实够用,问题大概率出在5万条数据没做shuffle和分组打包,长度参差不齐导致padding浪费太多。试试把batch size降到1,配合梯度累积到8,同时用torch.compile和flash attention,应该能稳很多。4bit微调确实会让效果飘,尤其LoRA rank别太低,建议16起步,另外bitsandbytes对推理速度影响大,微调完再转回bf16部署就行。
24G跑8B全参微调确实勉强,但LoRA不应该爆成这样,你试试把batch size降到1或者2,然后配合gradient checkpointing,这个比梯度累积更省显存,我上次用类似配置跑7B模型直接省了快一半。4bit微调确实会掉点,但主要看任务,如果是对话生成可能感觉明显,分类任务就还好,你可以考虑用NF4加double quant,比普通4bit稳一些。另外检查下你数据集的padding策略,5万条对话如果长度不齐,动态padding能省不少空间,别用固定长度硬塞。
试试4bit+梯度累积到16,batch先压到2,效果飘可能是lr没调,降到1e-4看看。
24G跑8B LoRA按理说够用,你batch size 4太大了吧,我一般设1或者2配合梯度累积32步,显存峰值能压到14G左右。4bit微调确实会掉点,但主要影响生成质量,分类或对话任务还能接受,建议量化后加些LoRA rank到32补偿一下。offload到CPU会慢很多,不如试试torch.compile加flash attention,我这套组合下来显存省了差不多30%。另外你数据量5万条不算小,可以检查下有没有padding浪费token,用动态padding能再挤点空间出来。
试试把batch size压到1再加64步梯度累积,LoRA rank降到8,4bit微调其实没那么玄乎,效果飘大概率是学习率没调好。
4bit微调8B损失不大,但你这数据量5万条,不如先拿2万条跑个baseline,看看是不是预处理环节出了问题。
试试把batch size降到1,然后开gradient checkpointing,配合paged optimizers,24G能跑8B的4bit微调。
试试unsloth优化+8bit加载,能省一半显存,4bit效果确实飘,8bit安全感强多了。
试过unsloth没?专为微调优化的,显存能砍一半,速度还快。
4bit微调8B其实够用,效果飘就调下学习率,别用默认的。
24G跑8B LoRA其实挺极限的,我试过把batch size压到1再加64步梯度累积,配合bf16和8bit量化能勉强跑起来,但速度确实感人。你说的4bit变慢我猜是bitsandbytes的LLM.int8()对某些层做了混合精度计算,试试换bnb_4bit_use_double_quant=True加nf4,能快一点。至于效果飘,LoRA rank别太高,16左右就够,另外数据集有5万条的话,考虑只用其中1-2万条做子集先验证流程,省得白烧显卡。分片加载用accelerate的device_map='auto',offload到CPU虽然有瓶颈,但至少能防止OOM中断训练。
4bit微调确实影响稳定性,试试unsloth或QLoRA,显存能压到12G左右,效果稳很多。
我之前也爆过,后来把batch降到1配合梯度累积,再加flash attention,24G勉强能跑。
24G跑8B LoRA其实不用上4bit,你把batch size降到1,配合梯度累积到32步,再加paged AdamW,显存基本能压住。4bit微调确实会掉点,尤其对话数据,效果飘很正常,我之前试过感觉像模型得了失忆症。offload到CPU会拖慢速度,但总比OOM强,你可以试试把optimizer和梯度offload,激活值留在GPU。另外改attention的话,可以试试FlashAttention-2,能省不少显存,但记得先确认你用的PyTorch版本支持。
试试8bit加paged optimzer,bs压到1,堆梯度累积,24G能跑得动,4bit丢精度是真的。
别死磕LoRA了,QLoRA配合unsloth那套,省显存还快,效果飘大概率是学习率没调对。
24G跑8B LoRA按理说够用,你batch size直接砍到1或者2,配合梯度累积到8,先把训练稳定性稳住再说。4bit微调确实会掉点,但你要是用QLoRA加适配器,效果飘多半是学习率没调好,试试1e-4以下。offload到CPU我试过,慢得怀疑人生,不如把序列长度截断到512,数据集里多余轮次直接丢掉。另外可以把attention换成flash-attn,省显存效果立竿见影,就是安装有点折腾。
24G跑8B的LoRA按理说够用,你试试把batch size降到1,然后用梯度累积到32步,效果和batch 4差不多,显存能省下一大截。4bit微调确实会让模型能力打折,特别是对话这类需要细腻语义的任务,我建议你试试8bit加QLoRA,速度和效果平衡会好很多。另外把序列长度截到1024,数据集里太长的样本直接丢掉,也能缓解很多压力。你检查下是不是把attention的算力浪费在padding上了,用flash attention能再挤出一部分显存。
这问题太典型了,4090跑8B LoRA按理说24G是够的,你batch size 4直接爆大概率是序列长度没限制住,或者优化器状态没开分片。我建议你先试试把max_seq_len砍到1024,然后配合gradient_checkpointing,batch size调到1,梯度累积开8步,这样显存能压到15G以内。4bit微调确实会掉点,但主要影响下游任务的稳定性,如果你用QLoRA记得把lora的dropout调低点,学习率也降一档,效果会稳很多。另外你提到的offload,我试过把优化器状态offload到CPU,速度慢个30%但能腾出5G显存,实在不行再上这招。
8B用4bit微调其实损失没想象中那么大,关键看你数据集和任务类型,如果是对话生成这种通用场景,量化后的模型能力衰减基本在可接受范围内。你生成变慢可能是bitsandbytes的4bit推理没走融合算子,试试加载时加个torch.compile或者用vLLM来跑,速度能提回来不少。显存优化的话,除了楼上说的,我强烈建议你检查下dataloader是不是num_workers开太多,有时候这也会占额外显存,尤其Windows下特别明显。还有个小技巧,把LoRA的target modules改成