最近在试着用LoRA微调Llama 3 8B,设备是RTX 4090 24G,数据集大概5万条对话。结果刚跑第一轮batch size设成4就直接OOM了,试了梯度累积和混合精度(bf16),还是撑不住。
我看网上有人说用bitsandbytes量化到4bit能省显存,但我加载模型后生成文本变慢了好多,而且微调完效果有点飘。
想问问大家,除了换硬件,有没有什么实用的显存优化技巧?比如分片加载、offload或者改attention?另外,8B模型用4bit微调会不会损失太多能力?求指个路,有点迷茫。
大佬们,用PyTorch跑Llama 3微调,显存爆了怎么办?
全部回复
共 161 条4090 24G跑8B微调确实有点极限,我试过把batch size降到1、配合gradient checkpointing和4bit量化,大概能压到14G左右,但训练速度会慢一些。你提到的offload也可以试试,把优化器状态挪到CPU上能省不少显存,就是来回传输数据会拖慢速度。至于4bit微调后的效果,我实测下来感觉任务简单的话影响不大,但复杂对话任务确实会有点“丢细节”,如果对效果要求高还是建议用8bit或者直接上3090 24G双卡。
4090 24G跑8B模型确实吃紧,你试试把batch size降到1,配合gradient checkpointing和4bit量化,我这样跑过5万条数据勉强能撑住。4bit微调能力损失其实没那么夸张,尤其在LoRA这种低秩适配下,主要影响的是收敛速度而不是最终效果,gen速度慢可能是bitsandbytes没开fast inference。另外attention上可以试下FlashAttention-2,能省不少显存还提速。
batch size降到1试试,4090跑8B加4bit量化其实够用,效果飘可能是学习率没调对。
说实话24G跑8B模型用LoRA确实挺极限的,你batch size设4不爆才怪。我实际试下来,batch size设1、配合梯度累积到16步,再加bf16和gradient checkpointing,4090勉强能跑起来,但训练速度慢得让人崩溃。你说的bitsandbytes 4bit量化我试过类似情况,推理速度下降是因为它用的是NF4格式,推理时反量化有额外开销,但微调时如果结合LoRA的target modules选得巧(比如只调q_proj和v_proj),效果其实能稳住,不会飘太厉害。不过有个坑:4bit下LoRA的rank别设太高,8到16就够,太大反而容易过拟合还费显存。分片加载和offload到CPU其实不太推荐,因为通信延迟会让训练时间翻倍,除非你用ZeRO-3加NVLink。至于8B模型4bit微调会不会损失能力,得看任务——如果是对话生成这种需要细粒度语义的,确实会有点“塑料感”,但如果是分类或摘要这种粗粒度任务,基本不影响。你不如先拿1万条数据小规模试跑,把batch size压到1,gradient accumulation设到8,看看显存峰值再决定要不要换量化。
这情况我太熟了,4090 24G跑8B模型确实容易卡在显存瓶颈上,尤其是对话数据集一长,序列长度稍微上去点就炸了。你试的4bit量化方向是对的,但bitsandbytes那个速度慢其实跟量化本身的推理优化有关,微调时候可以用QLoRA——它本质是4bit加载+LoRA,训练时梯度只更新低秩矩阵,显存能压到12-14G左右,而且推理速度比纯4bit好不少。不过你说的效果飘,我猜可能是量化后精度损失导致loss震荡,建议把bnb的4bit配置里的nf4打开,再用double quantization,能稳住一点。关于分片加载,其实你这种单卡场景不太需要,offload到CPU倒是可以试试,但会拖慢训练速度,个人觉得不如把batch size降到1然后用梯度累积,配合gradient checkpointing,24G理论上能撑住8B的LoRA微调。至于8B用4bit会不会损失太多能力,实测下来做指令微调影响不大,但如果你数据集里有很多长文本或复杂推理任务,还是建议用8bit或直接bf16+LoRA,牺牲一点batch size换效果稳定。最后检查下你数据加载时有没有做tokenizer的max length截断,有时候一条对话几千token直接撑爆显存,设个512或1024的阈值能省不少。
4090 24G跑8B LoRA确实有点极限,5万条数据的话batch size 4太高了,降到1或者2配合梯度累积试试看,我一般用梯度累积步数8-16能稳住。4bit量化对推理速度影响挺明显的,微调时建议用NF4而不是FP4,效果会稳一点,我用QLoRA调过7B模型感觉能力损失在可接受范围。另外可以试试Flash Attention 2,能省不少显存,而且对训练速度也有提升。分片加载和CPU offload在Hugging Face的Trainer里开DeepSpeed或FSDP就行,不过配置略麻烦,可以搜下现成的配置文件改改。
24G跑8B全参本来就不现实,LoRA加bf16加梯度累积其实已经很极限了。4bit量化确实会让推理变慢,可以试试NF4配合双量化,微调时把bnb_4bit_use_double_quant打开,效果飘可能是学习率没跟着量化一起调。另外把attention换成Flash Attention 2,显存能再省个1-2G,batch size设成2慢慢累积也行,5万条数据不差那点速度。
你这配置跑8B模型确实有点极限,24G显存上LoRA按理说batch size 1或2应该能跑,设成4直接爆很正常。我试过用gradient checkpointing配合bf16,batch size开2,序列长度别超过2048,基本能稳在22G左右。量化到4bit确实会降推理速度,而且微调时梯度更新的精度也会受影响,尤其对话数据多的时候,效果飘可能跟量化后参数敏感度变化有关。你可以试试QLoRA那种做法,把4bit量化后的模型冻结,只训练低秩适配器,这样显存和速度能平衡点。至于分片加载,其实PyTorch的FSDP或者DeepSpeed的ZeRO-3都能把优化器状态offload到CPU,代价是训练时间翻倍。另外attention这块,用FlashAttention-2能省不少显存,尤其长序列场景下效果明显。8B模型用4bit微调,如果任务比较通用,能力损失其实可控,但要是领域很专精,建议先用8bit或者干脆不量化,牺牲点batch size保精度。
试试用Unsloth加载模型,自带4bit优化和梯度检查点,24G跑8B LoRA batch size能拉到8。
24G跑8B模型全量微调确实紧巴巴的,你这情况其实挺典型的。LoRA本身已经省了不少,但5万条数据batch size 4还是太激进了,我一般会先试batch size 1配合梯度累积,比如累积8步等效batch size 8,这样显存压力会小很多。量化到4bit确实能省显存,但生成变慢和效果飘大概率是因为bitsandbytes的4bit对权重精度压缩太狠,微调时梯度更新容易失真,推荐试试QLoRA的NF4量化,它专门为微调设计了归一化分布,效果会稳一些。另外你提到的attention优化可以考虑用FlashAttention-2,它能减少显存占用并加速计算,PyTorch 2.2以上直接支持。分片加载和offload也能用,但offload到CPU会拖慢训练速度,更适合推理场景。关于8B模型用4bit微调会不会损失能力——这取决于下游任务,如果对输出流畅度要求高的话建议至少保留8bit,但如果你只是做特定领域适配,4bit微调后效果其实和全精度差距不大,尤其是LoRA这种参数高效方法。可以先用小数据集验证一下量化后的收敛曲线,再决定要不要上全量。
试试把batch size降到1,配合gradient checkpointing和4bit量化,24G勉强能跑,8B用4bit微调效果影响不大。
我最近也拿4090试过类似配置,batch size设成2加8步梯度累积能跑起来,4确实容易爆。4bit微调我实测在对话任务上loss会高一些,但用qlora加paged optimizer可以把显存压到14G左右,生成变慢可能是量化配置没调好。可以试试先分片加载模型,然后开gradient checkpointing,attention用xformers的memory efficient模式,这三板斧下来基本能跑通。
试试Unsloth框架,显存能省一半,4bit微调8B效果其实还行,别太担心。
这情况太真实了,24G显存跑8B模型确实容易卡在边界上。你试的bf16加梯度累积方向是对的,但batch size=4对于5万条数据来说还是有点贪心,建议先压到1试试,配合gradient_accumulation_steps=8或者16,这样等效batch size能上去但显存峰值会降不少。bitsandbytes 4bit量化确实会牺牲推理速度,尤其是8B模型在4090上跑4bit,反量化开销很明显,微调后效果飘可能是因为量化导致的梯度噪声,可以试试用QLoRA那个方案,把量化后的低秩矩阵保留在bf16精度下更新,效果会稳很多。分片加载和offload到CPU其实不太推荐,因为数据搬运开销大,反而会让训练更慢。attention方面可以试一下FlashAttention-2,对长序列场景显存节省很明显,而且基本无损精度。至于8B用4bit会不会损失能力,实测下来大多数任务能保住90%以上的效果,但如果你的数据集很细粒度或者需要高精度输出,建议还是用8bit或者保留bf16。
4090 24G跑8B全参确实勉强,4bit量化加梯度累积是个路子,但生成变慢可能是bitsandbytes的4bit推理优化没开好,试试加载时加bnb_4bit_use_double_quant和bnb_4bit_compute_dtype=torch.bfloat16。分片加载加offload到CPU也能救急,不过数据量5万条,建议把batch size降到1再配合梯度累积,显存压力会小很多。4bit微调能力损失对大部分任务影响不大,但如果你追求高精度,换成QLoRA加NF4量化试试,效果更稳。
试试deepspeed zero3加offload,8B用4bit微调影响不大,关键是数据集质量要够硬。
建议试试gradient checkpointing,batch size先降到1,用8bit adam加4bit模型,我这样跑13B都能稳在24G。
我也遇到过类似的情况,4090 24G跑8B模型确实有点极限,尤其你的数据集5万条不算小。batch size设4直接OOM很正常,我试过把batch size降到1,然后梯度累积设到16或者32,这样等效batch size没变但显存压力会小很多。另外你提到4bit微调后效果飘,这个我也有体会,主要是量化会引入噪声,特别是LoRA的rank如果设得低(比如8或16),表达能力本身就受限,加上4bit精度损失,训练时学习率得调小一点,比如1e-4甚至5e-5,而且warmup step可以适当拉长,让模型慢慢适应量化后的参数分布。不过如果你对生成速度有要求,4bit确实会变慢,因为反量化过程在推理时是额外开销。还有个思路是尝试用gradient checkpointing,就是把中间激活存成临时文件,虽然慢一点但能省不少显存,配合bf16和batch size 2应该能跑起来。至于分片加载或者offload,我试过把优化器状态offload到CPU,但会拖慢训练速度,而且显存节省有限,感觉不如直接调低batch size加梯度累积来得划算。8B用4bit微调会不会损失能力?我觉得如果任务比较通用,像对话或指令跟随,其实影响不大,但如果是领域知识密集的任务,比如医疗或法律问答,建议至少用8bit,因为低精度下参数更新的有效位数少,容易丢失细节。
你这情况太真实了,4090 24G跑8B模型确实容易卡在显存瓶颈上。4bit量化确实能省不少显存,但推理变慢和效果飘的问题我也有同感,尤其是微调时梯度更新不够精细,模型能力会有一定折扣——不过如果只是做特定任务微调,损失其实能接受,关键看你的数据集质量。
我试过把batch size降到1,配合梯度累积到16步,再加bf16,勉强能跑起来,但速度慢得让人崩溃。分片加载和offload到CPU其实挺有用的,比如用device_map="auto"或者设置offload_folder,能释放一些显存,但代价是训练时间会翻倍,得看你等不等得起。
改attention的话,可以用torch的sdpa或者xformers的memory efficient attention,实测能省10%-15%显存,而且效果基本不变,这个我强烈建议试试。另外,检查一下你的DataLoader是不是把整个数据集都预加载了,有时候用流式读取或者map-style dataset按需加载也能缓解。
至于能力损失,我个人的经验是:如果只是对话微调,4bit量化后效果飘可能是学习率或者LoRA rank没调好,跟量化本身关系不大。你试试把rank从8降到4,或者用rsLoRA,说不定能稳住。
总之别急着换硬件,先优化这些细节,实在不行再考虑租个A100或者用多卡分布式训练。
老实说,24G跑8B模型确实有点极限,我当时用3090也踩过类似的坑。batch size设到4直接爆很正常,建议先降到1试试,配合梯度累积到8或者16,效果其实差不多。混合精度bf16开了是对的,但别忘了把torch.compile打开,能压掉不少显存碎片化的问题。
关于量化4bit,你说的生成变慢和效果飘我也有同感,尤其是微调时如果用了NF4格式,反向传播的精度损失会被LoRA放大。我后来换了个思路:模型本体用8bit加载(bitsandbytes的fp8),只把LoRA的adapter保留在fp16,这样显存占用比4bit高不了太多,但生成速度和效果稳很多。另外可以试试把attention换成FlashAttention-2,对长序列的显存节省特别明显。
至于分片加载和offload,说实话在单卡上意义不大,因为CPU-GPU来回搬数据反而会拖慢训练。如果你数据集有5万条,不如先挑个几千条代表性的跑一轮验证效果,确认方向没问题再全量上。最后8B用4bit微调,我觉得任务依赖性强,如果是对话风格调整影响不大,但要是做指令跟随或者推理能力提升,损失就挺明显的。建议你先用8bit跑个小batch试试,不行再降量化。