最近在试着用LoRA微调Llama 3 8B,设备是RTX 4090 24G,数据集大概5万条对话。结果刚跑第一轮batch size设成4就直接OOM了,试了梯度累积和混合精度(bf16),还是撑不住。
我看网上有人说用bitsandbytes量化到4bit能省显存,但我加载模型后生成文本变慢了好多,而且微调完效果有点飘。
想问问大家,除了换硬件,有没有什么实用的显存优化技巧?比如分片加载、offload或者改attention?另外,8B模型用4bit微调会不会损失太多能力?求指个路,有点迷茫。
大佬们,用PyTorch跑Llama 3微调,显存爆了怎么办?
全部回复
共 162 条24G跑8B LoRA确实紧巴巴的,batch size 4会爆很正常。试试把gradient checkpointing打开,再加一个8bit AdamW优化器,这样能省不少显存,我上次用类似配置batch size能撑到2。4bit微调确实会有精度损失,特别是对话任务里细粒度语义可能会漂,建议优先用8bit量化,或者干脆用NF4加double quantization,效果稳不少。分片加载其实不太解决显存瓶颈,主要省的是CPU内存,offload到CPU的话速度会慢得离谱,不如改attention用FlashAttention-2,显存占用直接砍半。
我最近也踩过这个坑,4090跑8B LoRA确实得精打细算。试试把batch size降到1,配合gradient checkpointing和更小的lora rank(比如8),能省不少显存。4bit微调其实挺成熟的,用QLoRA配合NF4量化,效果下滑在可接受范围,别太担心。另外可以检查下数据集是不是每条都太长,截断到1024 tokens说不定就稳了。
24G跑8B LoRA确实紧张,batch size调到1试试,配合梯度累积到等效4,显存能降不少。4bit量化微调后效果飘是常见问题,可以试试QLoRA,它把量化误差考虑进去了,比直接4bit微调稳定。另外开gradient checkpointing也能省点显存,但会慢一些,先撑住跑起来再说吧。
试试unsloth那个库,专门优化显存的,8B模型4bit微调效果基本不掉,我跑过类似的。
24G跑8B LoRA确实有点极限,batch size 4还开bf16的话,光模型参数加梯度就快20G了,数据集又大,数据加载也会吃显存。试试把batch size降到1,梯度累积步数设成8,效果一样但显存能压下来不少。4bit量化确实会牺牲一些生成质量和精度,尤其对话微调这种任务,我试过用NF4加双量化能稍微好点,但速度慢是通病。还有个小技巧,用torch.compile编译一下模型,有时能省10%-15%显存。
我最近也踩过这个坑,4090 24G跑8B LoRA确实紧巴巴的。你batch size设4肯定爆,我试下来batch size=1加上梯度累积16步,搭配bf16和gradient checkpointing,大概能压到20G出头。4bit量化确实能省一半显存,但你说的生成变慢和效果飘我也有同感,尤其是QLoRA微调后模型对指令的遵循能力会有点下降,可能跟量化带来的精度损失有关。不过你可以试试换用NF4量化配合双LoRA,比单纯int4要好一些。另外分片加载其实对显存帮助不大,那是解决加载时内存不足的;offload倒是可以试试,把一些层放到CPU上计算,但会显著拖慢训练速度。你数据集5万条不算小,可以考虑先在小数据集上实验下各种配置,比如先拿500条跑通流程,再逐步放大,不然每次OOM重来太耗时。关于attention优化,用PyTorch 2.0自带的sdpa或者xformers的memory-efficient attention能省10%-15%显存,而且基本不影响效果,这个强烈推荐。至于8B模型用4bit微调会不会损失太多,我自己的感觉是如果下游任务比较通用(比如对话、指令跟随),4bit微调后效果还行,但如果是专业领域(比如代码、数学),建议至少用8bit。
我之前也遇到过类似问题,后来换了方案。
老实讲,你这配置跑8B模型5万条数据确实有点极限,4090 24G在LoRA下其实挺吃紧的。你提到quantization慢了而且效果漂,我怀疑是4bit下原本的数值精度损失在微调时被放大了,特别是对话数据对上下文敏感,稍微一波动就可能崩。我自己的经验是,可以试试先把batch size降到1,然后梯度累积步数拉高到16甚至32,这样等效batch size不变但单步显存压力小很多。另外,attention那块用FlashAttention-2(PyTorch 2.0以上自带)能明显省显存,而且速度几乎不减。分片加载其实不用太折腾,HuggingFace的device_map="auto"配合load_in_8bit(不是4bit)往往更稳,8bit下精度损失小很多,微调完效果也更可控。至于offload到CPU,如果显存还是爆,可以把optimizer states offload出去,虽然慢点但至少能跑通。最后,8B模型用4bit微调其实有点浪费参数能力,除非你数据集特别大或者任务特别简单,不然我建议优先考虑8bit或者干脆换7B的模型试试。
你这情况我太熟了,4090 24G跑8B全参微调确实捉襟见肘,LoRA加bf16按理说能省不少,但batch size 4加上5万条数据,梯度累积步数一多中间状态也吃显存。试试把batch size降到1或者2,然后梯度累积步数调到8或16,这样等效batch size够大但峰值显存会低很多。4bit量化微调确实会损失一些表达能力,尤其是对话任务对语义细节敏感,我建议你用bitsandbytes的NF4加双量化,推理慢可以配合Flash Attention 2,生成速度能提回来不少。另外offload到CPU是个办法,但注意把embedding和lm_head这些关键层留在GPU,不然微调完效果飘得更厉害。分片加载的话,用accelerate包的device_map="auto"配合max_memory参数,把部分层分配到CPU显存,但训练速度会明显下降。要我说,如果数据集质量够高,可以先在小模型上试水,或者干脆租个A100,4090折腾起来的边际效益其实不高。
4bit量化确实能省显存,但你说的生成变慢和效果飘其实很常见,尤其8B模型量化后精度损失在微调任务上会更明显。建议试试torch.compile加gradient checkpointing,再把batch size压到1配合梯度累积,4090跑8B应该能稳在16G左右。另外attention用flash attention或者xFormers也能省不少显存,而且对微调效果影响很小。你数据集5万条的话,可以考虑先用5000条小样本调通流程,再全量跑。
4090 24G跑8B全量微调确实有点勉强,你试试把batch size降到1,然后梯度累积步数设到8或者16,这样显存压力会小很多。4bit量化确实会牺牲一些推理速度和效果,但LoRA本身已经降低了参数量,我建议你试试8bit量化+ZeRO 3 offload,效果会比4bit稳一些。另外可以检查下是不是数据集预处理时token长度没截断,有时候长文本才是显存杀手。
5万条对话确实有点猛,4090 24G在8B模型上batch size 4扛不住很正常。建议试试deepspeed stage 3加offload,或者用unsloth这个库做4bit微调,它针对Llama做了显存优化,速度比原版bnb快不少。另外你的数据集可以考虑做下过滤或截断,太长的对话直接切掉能省不少显存。4bit微调主流做法是保留关键层的精度,效果飘的话可以检查下target_modules是不是选得太广了。
我也遇到过类似的情况,24G跑8B微调确实紧巴巴的。4bit量化加gradient checkpointing能省不少,但推理速度慢是bitsandbytes的老问题了,可以试试用torch.compile加速,或者换个思路用Unsloth,它内置的优化对显存友好很多。至于能力损失,8B用4bit微调日常对话任务其实还行,但要是做复杂推理或者生成格式要求高的内容,效果确实会飘,建议你拿小样本先对比测试一下。另外你可以试试把batch size降到1,配合梯度累积,虽然慢点但至少能跑起来。
老实说你这配置跑8B全量微调确实有点吃力,24G显存上LoRA按理说应该能撑住,但5万条对话的数据量一上来,batch size 4直接爆很正常。我试过把batch size降到1,配合梯度累积到8或者16,再开bf16和gradient checkpointing,基本能稳定在15G左右,你可以试试这个组合。bitsandbytes的4bit量化其实不是万能的,它确实能省显存但推理速度下降是硬伤,而且微调时量化参数和LoRA适配性有时候会出问题,导致训练不稳定,效果飘大概率是量化精度损失和LoRA秩设置冲突了。分片加载和CPU offload是个思路,但offload速度太慢,实际用起来很折磨,不如直接考虑gradient checkpointing加更小的batch。至于8B模型用4bit微调会不会损失能力,这得看你的任务,如果下游任务对语义细节要求高,比如情感分析或者复杂指令理解,4bit确实可能掉点,但如果是分类或者简单对话,影响其实有限。你不如先拿一个1000条的子集跑一遍,把batch size调到2,打开gradient checkpointing,看看显存占用再决定要不要换量化方案。
刚入门,这个对我帮助很大。
24G跑8B全参微调确实有点极限,batch size 4爆显存正常。你的方向是对的,4bit QLoRA+gradient checkpointing基本是标配,文本变慢可能是因为没开fast inference或者量化配置没调好,检查下bnb_4bit_compute_dtype是不是也设成bf16了。分片加载其实不太管用,offload到CPU会拖慢速度,不如试下DeepSpeed ZeRO-3或者torch.compile,能省不少显存。至于能力损失,8B模型4bit微调在对话任务上我个人体感影响不大,但如果你数据集里有很多长尾知识,建议至少保留8bit。
4090跑8B全量参数还是太勉强了,试试QLoRA加8bit优化器,batch size降到2配合梯度累积。
batch size降到1试试,配合8bit优化器再加个gradient checkpointing,4090跑8B LoRA应该能稳。
4090跑8B LoRA的话batch size先降到1试试,4bit微调确实会影响输出质量,建议优先调低seq_len。
24G跑8B其实LoRA+8bit量化就能稳住的,batch size先降到1试试,配合梯度累积到等效4的效果。4bit微调确实会有点能力损失,尤其是对话类任务,建议优先用NF4而不是FP4。另外可以试试Paged AdamW优化器,它对显存碎片管理比默认的好不少,至少能省出1-2G。