最近在试着用LoRA微调Llama 3 8B,设备是RTX 4090 24G,数据集大概5万条对话。结果刚跑第一轮batch size设成4就直接OOM了,试了梯度累积和混合精度(bf16),还是撑不住。
我看网上有人说用bitsandbytes量化到4bit能省显存,但我加载模型后生成文本变慢了好多,而且微调完效果有点飘。
想问问大家,除了换硬件,有没有什么实用的显存优化技巧?比如分片加载、offload或者改attention?另外,8B模型用4bit微调会不会损失太多能力?求指个路,有点迷茫。
大佬们,用PyTorch跑Llama 3微调,显存爆了怎么办?
全部回复
共 162 条4090跑8B LoRA其实卡在数据集太大,5万条对话单轮epoch就够呛。你可以试试把batch size压到1,配合8bit的bitsandbytes,然后开gradient_checkpointing,显存能省下一大半,速度损失还能接受。至于4bit微调,说实话对8B这种大模型影响不小,尤其是对话任务,效果飘很常见,建议先用QLoRA的NF4格式,再调低学习率多跑几轮,比纯4bit稳。另外分片加载和offload对单卡意义不大,不如把注意力换成flash-attention-2,能再挤点显存。你数据集这么大,不如先抽个5000条跑通流程,确认效果再上全量,省得反复试错。
4090 24G跑8B LoRA确实紧巴,batch size 4不爆才怪,我一般直接设1配16梯度累积。4bit微调效果飘大概率是NF4加双重量化导致的,试试8bit加paged_optimizer,或者干脆用qlora的embedding和lm_head单独调精度。另外可以开flash attention 2,显存能省15%左右,配合zero-inference的分片加载先把推理跑通再微调。最后别迷信offload,除非你CPU内存大,不然换数据交换的延迟够你喝一壶的。
24G跑8B LoRA其实挺极限的,但5万条对话真没必要一次全塞进去。你可以试试把数据集切成长度更均匀的块,配合动态padding,能省不少显存。另外batch size 4爆了不代表batch size 1也爆,先降到1把流程跑通,后面再用梯度累积找平衡,别一上来就挑战极限。
bitsandbytes 4bit那个生成变慢我太懂了,因为量化后kernel效率低,尤其你还在微调,反向传播时反量化开销更大。如果非要用,建议把NF4和双重量化都打开,再把attention里算力最重的几个层保留为8bit或者bf16,混合精度策略比全量4bit稳得多。
关于能力损失,4bit微调8B确实会有点飘,尤其是对话任务这种需要细粒度语义的。我之前用QLoRA跑7B模型,效果比全精度差一截,后来改成把embedding层和lm_head留在高精度,只量化中间层,效果就回来了。你可以试试这个思路,反正这两层占显存也不大。
分片加载和offload其实治标不治本,除非你同时开CPU offload,但那样训练时间会翻倍。更推荐你看看torch.compile和flash attention,PyTorch 2.0之后这两招能压掉不少显存和速度开销,我用了之后batch size能从1提到3。
还有个冷门技巧,把数据集里特别长的样本过滤掉,或者用梯度检查点加显存,虽然慢点但能换容量。我觉得你先跑通小batch,再用nvidia-smi盯着看哪一步峰值最高,比盲目调参靠谱。最后问下,你用的是peft库的LoRA吗?target_modules设置成哪些层了?这个影响也很大。
说实话你把batch size降到1加上gradient checkpointing应该就能跑,24G跑8B LoRA没那么极限。4bit微调确实会掉点,但关键看你的任务领域,如果是通用对话影响不大,垂直领域可能会飘得厉害。offload到CPU会拖慢训练速度,不如试试把序列长度截断到512,很多数据集没那么吃长上下文。另外你5万条数据其实可以先小batch跑通流程,确认效果再上全量,别一上来就挑战极限配置。
5万条对话这体量,直接上LoRA确实勉强,我建议你先用8bit跑,bf16+gradient checkpointing把batch压到1,然后开梯度累积到8,显存肯定够。4bit生成变慢是因为反量化开销,微调时其实没太大影响,但效果飘的话可以试试只量化attention层,其他层保持bf16。还有,你检查下是不是transformer版本太新,有些实现会默认缓存KV导致爆显存,关掉use_cache试试。
我倒是觉得你这问题不全在显存,5万条数据LoRA完全够用,但batch4太贪了。先试试sequence packing,把短对话拼起来填满上下文,能省不少显存。至于4bit微调,别全量化,用QLoRA那种NF4格式再加double quant,效果比普通4bit稳很多
我最近也在搞类似的事,4090跑8B LoRA确实极限,batch size降到1加上梯度累积到8步基本能稳住,但你还得把序列长度截到1024以内。4bit量化掉精度是必然的,尤其对话数据多的时候,效果飘大概率是学习率没跟着调低,建议试试0.0001左右。offload到CPU只在反向传播时开能救一点显存,但速度会慢到怀疑人生,不如先砍数据集到2万条看看曲线是否还正常。
24G跑8B LoRA其实不该这么狼狈,你试试把batch size降到1,配合8倍梯度累积,效果等效但显存压力小很多。bitsandbytes 4bit确实会掉点,尤其对话类任务,建议改用NF4加double quant,速度损失会好一点。另外记得把gradient checkpointing打开,这玩意能省一半激活内存。分片加载和offload是最后的招,能不动就别动,IO瓶颈比显存更头疼。
24G跑8B LoRA其实挺极限的,batch size 4爆正常,我试过把batch压到1再加64步梯度累积,配合bf16和gradient checkpointing能勉强稳住。4bit微调确实会掉点,尤其对话数据多的时候,建议你试试QLoRA加NF4,别用FP4,效果能稳一点。至于生成变慢,大概率是bitsandbytes的dequantize开销,你可以看看能不能把adapter和base model分开存,推理时只加载adapter。分片加载和offload到CPU也是办法,但会拖慢训练速度,不如先减小序列长度或者过滤下数据集里超长的样本。
试试8bit加paged_adamw,batch降到2,序列长度砍到512,24G能跑起来,效果比4bit稳。
4bit微调掉点确实明显,尤其对话任务,建议先看下tokenizer和padding有没有对齐。
4bit微调掉点正常,试试QLoRA加paged optimizers,batch再砍半,能稳不少。
24G跑8B还开4的batch确实极限了,我一般先砍到2再加梯度累积,等效batch靠累积拉回来,显存能稳不少。4bit微调掉点其实没想象中严重,尤其LoRA只训一小部分参数,但你说的生成变慢可能是bitsandbytes的4bit推理没走优化内核,试试加载后转回bf16权重做推理。另外可以开torch.compile加flash attention,显存和速度都有改善,5万条数据不用一次塞满,用流式加载配合缓存也很关键。offload到CPU是最后手段,但会拖慢训练,建议先调batch和梯度累积试试。
24G跑8B LoRA按理说够的,你batch size 4爆显存大概率是序列长度太长或者优化器状态吃得多,试试把seq len砍到1024再加gradient checkpointing,能省一半。4bit微调确实会影响下游任务稳定性,尤其对话生成容易飘,建议用NF4加double quant,或者干脆用8bit加LoRA,速度损失能接受。offload到CPU也是个办法,但会慢很多,适合偶尔调参用。另外你5万条数据量不算大,先小batch跑几百条看看loss曲线,别急着上全量。
24G跑8B LoRA按理说够了,你batch size 4爆显存大概率是序列长度或数据集里长样本太多,试试把max_seq_len砍到1024甚至512,再配合gradient checkpointing,能省不少。4bit微调确实会有精度损失,但如果你用QLoRA的NF4格式,效果其实比普通4bit稳定,生成变慢可能是没开flash-attention或bitsandbytes没走对路径。我自己的经验是,LoRA rank设16左右,target modules全加在attention层上,5万条数据跑起来完全没问题,你检查下是不是把embedding也加了trainable。要是还爆,就把优化器换成AdamW 8bit,显存占用能再降1-2G。
这题我熟,之前用4090跑7B也卡得死去活来。你试试把batch size降到1,配合gradient checkpointing,加上8bit优化器(比如bnb的AdamW8bit),显存能省下一大截。4bit微调确实会掉点,尤其是在对话任务上,但如果你数据集质量够高,实测影响能接受。另外可以试下unsloth框架,它对Llama做了专门的显存优化,同样配置下比我手动改的省了快30%。
24G跑8B LoRA其实挺极限的,你batch size直接砍到1或者2,配合8bit的AdamW优化器能省不少。4bit微调确实会掉点,但你要是用QLoRA的NF4格式再加回一些LoRA rank,效果会稳很多,别用普通的4bit量化。生成变慢可能是因为bitsandbytes的4bit在解码时没做kernel融合,试试最新的版本或者换GPTQ。另外5万条数据可以先用更小子集试跑通,确认loss在降再全量上,不然排查问题成本太高。offload到CPU太慢不推荐,你真要省显存不如把输入序列截断到512,对话数据一般够用了。
试试unsloth吧,省显存效果明显,4bit微调8B真不太行,我试过掉点挺狠的。
4090跑5万条LoRA确实紧巴,把batch压到1加梯度累积,再把seq len截到512试试。
4bit微调掉点正常,试试QLoRA加paged optimizers,4080都能跑8B。
4090跑8B LoRA还得靠Unsloth,这个库对显存优化做得特别狠,同样设置能比原生省一半左右,加载时间也快。4bit微调其实没你想的那么玄乎,QLoRA那套理论就是干这个的,只要把学习率调低点、用上paged_adamw,效果基本能贴近16bit。你试试把gradient_checkpointing也开开,batch size压到2,加上梯度累积凑成4,应该就能跑动了。生成变慢是bitsandbytes反量化的正常代价,但微调阶段慢一点无所谓,推理时再换回原模型就行。
试试unsloth那个优化版LoRA,4090跑8B能塞下batch 8,速度还快不少。4bit微调影响不大,效果飘可能是学习率没调好。
试试unsloth吧,显存占用能砍一半,训练速度还快,4bit微调8B效果其实够用。
之前我也爆显存,后来把batch size压到1加梯度累积,再用paged_adamw优化器,总算跑起来了。
我之前也遇到过类似的情况,24G跑8B LoRA确实紧巴巴的。4bit微调效果飘大概率是量化造成的精度损失,建议试试把quantization和LoRA的target modules分开处理,或者用NF4加double quant,能稳不少。另外你可以看看torch.compile,配合gradient checkpointing能省出一大块显存,虽然编译慢点但训练稳。分片加载和offload到CPU其实不太适合你这个场景,反而容易拖慢速度。