最近在试着用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其实挺极限的。建议把batch size降到1,配合梯度累积到32步,效果差不多但显存压力小很多。4bit量化确实会掉点,尤其对话数据多的时候,可以试试把量化后的模型冻结embedding层,只训练attention和mlp,效果会稳一些。另外开gradient checkpointing能省不少,虽然慢点但至少不OOM。
这问题我上个月刚踩过坑,4090跑8B LoRA确实得精打细算。你可以试试把batch size压到1,配合8bit的QLoRA,然后开gradient checkpointing,这样显存能省出一大截。4bit微调能力损失其实没那么夸张,主要看你的数据集跟任务领域,如果和基座分布差太远效果才会飘。另外生成变慢大概率是bitsandbytes没走对路径或者没开fast quant,换个参数或者升级下版本能缓解不少。
同样的显存预算下,把attention改成flash-attention能省差不多2-3G,加上把优化器换Adafactor,比AdamW吃显存少很多。你5万条数据其实可以先抽样5000条试跑通,确认不OOM再全量上,省得浪费时间调参。4bit微调8B我试过,只要用QLoRA的原始配置,效果和bf16差距在5%以内,但速度确实会慢,这个只能忍了。
如果你数据集里有大量长对话,试试把最大序列长度从默认2048砍到1024,很多样本根本用不到那么长,显存直接减半。另外可以把模型切到NF4量化,比普通4bit稳一些,配合paged_adamw优化器能防峰值波动。我个人感觉8B用4bit微调,
24G跑8B LoRA,batch size 4确实有点紧张,但也不是完全没救。你试过把batch size降到1,然后梯度累积步数拉高吗?我上周用类似配置跑7B,bs=1加8步累积,显存峰值能压到18G左右,训练速度其实没慢太多,主要瓶颈在数据加载上。bitsandbytes的4bit我也踩过坑,推理慢是因为反量化开销,你试试加载时把bnb_4bit_compute_dtype设成float16,别用默认的float32,速度能回来一些。至于效果飘,大概率是量化后的数值精度影响到了LoRA适配器,建议把target_modules里那些关键线形层单独挑出来用8bit,其余保持4bit,混合精度微调我试过比全4bit稳很多。另外你可以开一下gradient_checkpointing,这个能省30%左右显存,代价是训练时间多20%,但总比OOM强。还有个偏方,把数据集用packing策略打成固定长度序列,减少padding浪费,5万条对话大概能省出两成的有效batch空间。关于8B用4bit会不会损失能力,说实话,如果是做指令跟随或者简单闲聊,区别不大,但你要是微调它做推理或数学题,效果跳水挺明显的,建议这种场景直接上QLoRA加冻结embedding层,或者考虑换7B的Mistral,性价比高不少。
24G跑8B LoRA其实卡在序列长度和batch的乘积上,你试试把max_seq_len砍到512,配合gradient_checkpointing应该能塞下batch4。bitsandbytes的4bit确实会牺牲速度,但微调效果飘可能是学习率没调好,我一般用4bit+LoRA会把lr降到1e-4以下。另外可以看看PEFT的target_modules是不是只选了q_proj和v_proj,少选几个模块也能省不少显存。
24G跑8B全参本来就紧,但LoRA+bf16按理说能挤进去,你batch size压到1了吗?梯度累积16步效果其实差不多。4bit微调确实会掉点,尤其对话数据,建议试试把quantization config里的double quant关掉,或者用NF4别用FP4,能稳一些。另外你数据集5万条太大,先抽1万跑通流程,确认loss在降再上全量,不然调参都费劲。
24G跑8B LoRA其实挺极限的,但你这配置不该第一轮就炸,batch size 4确实激进,先降到1或者2试试,配合梯度累积到8或者16,效果差不多但显存压力小很多。4bit量化掉精度是肯定的,尤其对话数据多的时候,模型容易变得“飘”,我猜你用的是QLoRA那套吧,建议把NF4的blocksize调大一点,或者试试8bit加double quant,速度会好一些。另外你说加载后生成变慢,大概率是bitsandbytes的4bit在反量化时开销大,可以考虑把embedding和lm_head留在fp16,这俩参数占比不小但很吃精度。至于offload,CPU offload对LoRA来说不太划算,通信瓶颈比显存更烦人,除非你batch size实在调不动了。还有个思路是换attention,比如用flash_attention或者xformers的memory efficient kernel,能省不少激活内存,但记得先确认你的transformers版本支持llama3。最后想问你用的什么框架,peft加transformers还是别的?如果是纯原生写,建议上accelerate的zero2或者zero3,能自动切分优化器状态,5万条数据其实也够你跑几个epoch找规律了。
我之前也遇到过类似情况,4090跑8B LoRA确实紧巴。你可以试试把batch size降到1,配合梯度累积到32步,效果上基本等价但显存压力小很多。另外,4bit量化微调确实会有精度损失,特别是对话任务上,建议改用NF4加上双量化,再用QLoRA的paged optimizers,能缓解一点。至于offload,如果数据加载不是瓶颈,把优化器状态offload到CPU比模型offload更划算。最后,如果你数据集质量高,其实可以考虑用8bit加LoRA rank设低点,效果比4bit稳。
24G跑8B LoRA确实紧张,我试过把batch size压到1加梯度累积,再把seq len截到512,勉强能跑但速度感人。4bit微调掉点其实没想象中严重,关键是记得在训练时把lora模块也转成bf16,生成慢多半是反量化开销,你可以试试QLoRA的double quant。另外检查下是不是数据集里有些长样本把显存峰值拉爆了,按长度分桶能省不少。
24G跑8B LoRA按理说够用,你batch size 4爆显存可能是数据长度没截断或者梯度检查点没开,试试把seq_len砍到1024再加gradient_checkpointing,能省不少。4bit微调确实会飘,尤其LoRA rank低的时候,建议量化到8bit或者NF4配合paged_optimizer,效果会稳一些。另外你5万条数据其实可以先小batch跑通再往上加,别一上来就贪大。
4090 24G跑8B LoRA本来就紧,5万条对话这数据量也不小,batch size 4属实有点贪了。我建议你先试试把batch size降到1,配合梯度累积到8或16,显存压力能小一大截,速度反而可能更快。
4bit微调确实会飘,尤其LoRA本身精度就敏感,我试过用NF4加double quant,效果比普通4bit稳一些,但推理速度慢是通病,基本只能牺牲速度换显存。要不你试试8bit加CPU offload?虽然慢点,但精度损失小很多。
另外你提到改attention,其实可以试试FlashAttention-2,PyTorch 2.x直接调用就行,能省不少显存,而且不用改模型结构。我跑7B模型时开这个加梯度检查点,24G能塞下batch size 2,效果比量化靠谱。
最后想问下,你用的peft库版本是最新的吗?之前有过显存泄漏的bug,更新到0.10以上会好很多。实在不行就分成两段微调,先冻结前半部分,再解冻后半部分,也能绕过去。
24G跑8B LoRA其实卡在门槛上,batch size 1加上gradient checkpointing基本能稳,你试试把seq长度限制在1024,数据预处理时做packing能省不少显存。4bit微调确实会掉点,你可以先拿一小批数据对比下全精度和4bit的loss曲线,效果没差太多再继续。offload到CPU太慢不推荐,倒是可以把optimizer状态用AdamW的8bit版,省出来的显存够你加batch size了。
24G跑8B LoRA按理说够用,你batch size 4确实猛了,降到1+梯度累积16步,效果一样但显存直接砍半。4bit微调掉点很正常,尤其对话数据,我试过把量化改成NF4+双重量化,生成慢但能忍,关键学习率要调低到1e-4左右。另外你试试torch.compile+flex_attention,能省不少显存,效果比offload稳。
4090跑8B LoRA其实不用太慌,我试过把batch压到1,配合gradient checkpointing,再把序列长度截到512,5万条数据照样能跑,就是慢点。4bit微调确实会飘,尤其对话任务,建议你用NF4加double quant,然后LoRA的rank别拉太高,8到16就行,效果能稳住。生成变慢大概率是bitsandbytes的4bit推理没走优化内核,试试加载后转回bf16再生成,或者直接用unsloth的版本,省显存还快。分片加载和offload对单卡意义不大,主要靠batch和序列长度硬扛,实在不行就换QLoRA加paged optimizer,能再省一截。
24G跑8B LoRA其实不用上4bit,你把batch size压到1,配合梯度累积到8,再把序列长度截到1024,大概率能稳。4bit微调确实会掉点,尤其对话数据多了之后效果飘很正常,我试过8bit加QLoRA会好一些。另外可以试试torch.compile加flash attention,能省不少显存,但注意别和梯度检查点一起开,容易冲突。
5万条数据不算小,你不如先拿几千条小batch跑通流程,确认loss下降趋势再全量上。offload到CPU的话,训练速度会慢得让你怀疑人生,不太推荐。真要省显存,改attention比换量化实在,比如用滑动窗口或者稀疏注意力,但得看你数据集里长对话占比高不高。
24G跑8B LoRA其实够用,你试试把batch size降到1,配合8倍梯度累积,再用paged_adamw优化器,显存能压下来不少。4bit微调确实会掉点,但如果你只训少量轮次,用QLoRA加lora_alpha调大点,效果飘可能是学习率没配好,试试1e-4加warmup。分片加载没必要,offload到CPU会慢得你想砸电脑,改attention的话不如直接换FlashAttention-2,省显存还提速。我自己的经验是,5万条数据用4bit训一轮够了,多了反而过拟合。
说实话24G跑8B LoRA按说不该这么惨,你试试把batch size降到1,配合梯度累积到32步,效果和batch 4差不多。4bit微调确实会损失点精度,但用QLoRA的NF4格式比普通bitsandbytes稳很多,生成变慢可能是因为没开torch.compile或者KV cache没优化。另外可以开flash attention2,显存能省不少,训练速度还快。效果飘的话,建议把学习率降到1e-4以下,LoRA rank别超过16,先跑几百条数据看看loss曲线再上全量。
4bit微调8B确实容易飘,建议试试QLoRA加paged_optimizer,能省不少显存。
offload到CPU会慢到怀疑人生,不如把batch降到1,用梯度累积撑住。
试试把batch size降到1配8bit,加paged_adamw优化器,能稳不少。4bit微调8B损失还行,主要看数据质量,别太慌。
试试unsloth的4bit微调,显存能压到12G以内,效果比bitsandbytes稳多了。
24G跑8B LoRA按理说应该够的,你batch size 4直接OOM有点奇怪,会不会是序列长度撑爆了?5万条对话里如果有些长上下文,padding到统一长度会浪费大量显存。建议先检查下数据预处理,试试动态padding加打包,能省不少。另外你可以把LoRA的target modules只放在attention的q和v上,别全加上,r值也降到8试试,效果不一定差多少,但显存能降一截。
至于4bit微调,说实话我觉得有点得不偿失。QLoRA虽然能跑,但生成速度慢是因为反量化开销,而且你感觉效果飘,大概率是量化后梯度的精度损失在LoRA低秩更新时被放大了。我自己的经验是,8bit加bf16混合精度其实是个甜点,显存占用比4bit高不了太多,但稳定性好很多。如果实在要省显存,试试torch.utils.checkpoint(激活重计算),这个对LoRA特别友好,配合gradient accumulation,batch size 4应该能压到20G以内。
还有个偏门点的方法,用DeepSpeed的ZeRO-3加offload到CPU,但你要注意offload参数会拖慢训练速度,别指望它提速。另外你把优化器换成AdamW的8bit版(bitsandbytes里那个),也能省几个G。最后关于能力损失,8B用4bit微调,如果任务简单比如分类或者短问答,问题不大,但对话生成这种语义敏感的,建议至少用8bit,别为了跑通牺牲太多质量。先试动态padding和激活重计算,这两个最容易见效。