最近在试着用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 4再加5万条数据,爆显存太正常了。4bit量化掉精度是必然的,尤其对话任务对细节敏感,效果飘不奇怪,我建议你试试把LoRA的rank降到8或16,同时用gradient checkpointing,这俩组合能省不少。另外你说加载后生成变慢,可能是bitsandbytes的4bit推理没走量化加速,试试torch.compile或者换GPTQ量化格式,速度会好很多。分片加载和offload到CPU其实不太适合微调,因为频繁传输会拖慢训练,不如牺牲一点batch size换稳定。
4090跑8B本来就得精打细算,试试8bit加梯度检查点,batch降到2先跑通再说。
4bit微调8B损失确实明显,建议试试8bit加paged optimizers,显存能省不少。
24G跑8B LoRA按理说是够的,你batch size 4爆显存大概率是序列长度太长或者优化器状态没省。试试把max_seq_len砍到1024,再用paged_adamw优化器,我这边8B全参微调都能压到16G以内。4bit量化掉精度是肯定的,尤其对话任务容易飘,建议优先考虑gradient checkpointing加梯度累积,实在不行就换QLoRA的double quant。你数据集5万条真不小,先用1万条跑通流程再上全量,别一上来就硬刚。
24G跑8B LoRA其实挺极限的,但没到完全没救。batch size 4对4090来说确实太大了,你可以试试batch size 1加梯度累积,配合bf16,显存能压到15G左右。4bit量化的问题在于它把embedding和lm_head也量化了,这两块对精度影响最大,建议用QLoRA的NF4格式,同时把这两个模块保留为fp16,效果会稳很多。至于生成变慢,那是bitsandbytes的4bit反量化开销,正常现象,但微调阶段速度慢点无所谓,推理时再换回原模型就行。另外你可以开gradient checkpointing,能再省30%显存,代价是训练时间变长,但总比OOM强。分片加载和offload到CPU不太建议,除非你内存特别大,不然数据搬运的IO瓶颈会让你怀疑人生。关于能力损失,4bit微调其实可以看作一种正则化,数据集够大的话泛化反而可能更好,但5万条对话对8B来说不算大,建议先用1万条子集试跑通,确认loss下降再全量上。改attention的话,torch的sdpa已经够高效了,没必要折腾flash-attn,除非你显存实在挤不出那几百MB。最后说句实在的,实在不行就租一张A100或两卡3090跑deepspeed stage3,体验完全不一样。
试试unsloth的LoRA,显存能省一半,4bit微调8B其实够用了,效果飘大概率是学习率没调好。
试试unsloth的4bit训练,显存能砍一半,速度还快,效果比bitsandbytes稳不少。
你这个配置跑8B LoRA其实挺极限的,5万条数据量也不小,我建议先把seq_len砍到512或768试下,数据集里长样本比例不高的话影响不大。4bit微调确实会掉点,尤其对话任务上飘得明显,不如试试把LoRA rank降到8或者16,配合gradient checkpointing和更小的batch size,虽然慢点但稳定。另外可以看看PEFT的prepare_model_for_kbit_training是不是把requires_grad都设对了,有时候冻结参数没处理好也会白占显存。
5万条对话全量塞进去确实太猛了,batch size先砍到1或者2,配合梯度累积到16步,显存压力会小很多。4bit微调我用过,效果飘大概率是学习率没调好,建议降到1e-4以下试试。另外可以开flash attention和torch.compile,能省不少显存,速度还快。至于能力损失,8B模型4bit做LoRA其实还好,主要看任务难度,简单指令微调问题不大。
5万条对话用LoRA其实可以试试把batch size压到1,配合8步梯度累积,效果和batch 4差别不大,而且显存占用能降一半。4bit微调掉点主要看任务,如果是对话生成确实容易飘,但把LoRA的rank提到16或者32,再加点dropout,稳定性会好很多。另外你试试用unsloth那个库,它对Llama的attention做了优化,24G跑8B bf16 LoRA完全够用,我上次2万条数据都没爆过。
batch size别硬刚4了,直接设1,梯度累积开大点,显存立马松快。4bit微调的问题不在量化本身,而是你加载后推理变慢,那是因为bitsandbytes的4bit反量化开销大,建议微调用4bit,保存后转回8bit或bf16重新加载推理,效果和速度都能兼顾。另外可以试试QLoRA加paged optimizer,那个对显存碎片处理特别好,我实测能多塞30%的数据。
说实话24G跑8B LoRA不该爆,你检查下是不是把整个模型都放GPU了,试试device_map='auto'加max_memory限制,把部分层offload到CPU。4bit微调能力损失其实没那么夸张,关键是你数据集太大,5万条对LoRA来说噪声多,不如随机抽个1
说实话24G跑8B LoRA按理说应该能挤一挤的,你batch size 4确实有点猛,我一般先压到1或者2,配合gradient accumulation把有效batch撑起来,显存占用能降一大截。bf16加梯度检查点基本是标配了,你试试把checkpointing开开,虽然慢点但真的能省好几个G。4bit微调我试过,推理慢是bitsandbytes的dequantize开销,微调完效果飘大概率是学习率没调好,4bit的QLoRA本来就敏感,建议lr降到2e-4以下,warmup拉长点。另外你5万条对话其实不算多,考虑只用其中一部分做子集先跑通流程,确认稳定了再全量上。offload到CPU也是个路子,但注意把offload的层数控制好,不然频繁搬运反而卡死。改attention的话,换成flash attention能省不少内存,而且速度还快,PyTorch里直接调就行。最后说下能力损失,8B用4bit微调其实指令跟随能力掉得不多,主要影响长文本生成的一致性,你要是任务对逻辑要求高,建议还是8bit加梯度检查点稳一点。
4bit微调确实容易飘,建议试下QLoRA加paged optimizer,batch降到2配合梯度累积,稳很多。
4090跑8B LoRA就别硬刚batch 4了,降到1加梯度累积试试,显存占用能砍一半还稳。4bit微调确实容易飘,建议换成8bit加paged_optimizer,效果和速度平衡点。
你这情况我之前也踩过,4090跑8B LoRA其实不用上4bit,先把batch size降到1,配合梯度累积到8或者16,显存就稳了。4bit微调确实会掉点,尤其对话数据多的时候,生成质量飘是正常的,建议试试8bit加nf4,或者用QLoRA的double quant,效果会好不少。另外可以开torch.compile和flash attention,能省不少显存,速度还快。你那个5万条数据,其实可以先抽几千条试跑通流程,确认没问题再全量上,省得反复OOM心态爆炸。
24G跑8B LoRA按理说能挤一挤,你batch size 4爆掉大概率是序列长度或者数据集里长样本拖累的,可以先检查下有没有超过2048 token的对话。4bit微调效果飘太正常了,尤其你用NF4的时候,LoRA的秩和alpha得跟着调,我试过秩16加alpha 32能稳一点,但推理慢确实无解,bitsandbytes的4bit反量化开销就在那。分片加载你试试accelerate的device_map='auto'加CPU offload,把部分层扔到内存,虽然慢点但能跑通,注意offload后梯度更新会变慢,最好配合梯度累积到32步。改attention的话,试试torch的SDPA或者flash-attn,能省不少显存,但你要是已经用了bf16可能提升有限。另外5万条数据其实可以砍到1万条先跑通流程,看loss曲线正常了再全量上,别一上来就硬刚。至于能力损失,我体感4bit微调在指令跟随上还行,但复杂推理和生成连贯性会掉,如果你任务偏生成还是建议8bit加gradient checkpointing,牺牲点速度换质量。
说实话4090 24G跑8B LoRA确实有点极限,但5万条数据真没必要硬啃全量,你可以试试把max_seq_len砍到512或者更短,对话数据本身冗余度就高,这招通常能省30%以上显存。4bit微调效果飘大概率是学习率没调对,建议降到2e-4以下,另外把target_modules只选q_proj和v_proj,别全上。分片加载加offload到CPU能救急,但速度会慢到怀疑人生,我上次试过基本就是边跑边等。你那个batch size 4改成2再加8步梯度累积,效果差不多但显存能压一半,实在不行就换QLoRA的nf4格式,比普通4bit稳定不少。
5万条对话用LoRA跑8B,24G确实有点极限了,batch size 4不爆才怪。我建议你试试把batch size降到1,配合梯度累积到8或16,这样显存压力会小很多,效果基本不受影响。4bit量化掉精度是真的,尤其对话任务里容易飘,不如先用8bit,或者干脆用NF4加上double quant,能稍微稳一点。另外你可以看看PEFT的prepare_model_for_kbit_training,那个参数没调对的话,LoRA训练时显存分配会很不合理。还有个偏方,把数据集先做一遍去重和截断,超过1024 token的样本直接过滤,能省不少显存,毕竟对话里长尾样本占比很高。
试试unsloth吧,同样LoRA能省一半显存,速度还快,4bit微调8B效果其实够用。
你这情况我太熟了,4090跑8B LoRA,batch size 4确实有点贪,降到2配合梯度累积就稳了。4bit微调效果飘大概率是NF4的量化噪声太大,试试8bit加Double Quantization,速度和精度能平衡不少。另外可以开torch.compile加flex_attention,省显存的同时还能让显存分配更均匀,数据集5万条的话建议先抽5000条跑通流程再全量上。
4090跑8B LoRA,batch size=4确实有点激进,我一般设2加梯度累积到8,效果没差多少但显存稳得很。4bit量化掉精度是必然的,尤其对话数据多的时候,生成飘可能跟这个有关,建议试试8bit加LoRA,速度慢点但能力保得住。另外你可以开torch.compile,配合flash attention v2,能省不少显存,我试过能多塞20%的batch。5万条数据其实不小,先跑个几千条子集看看loss曲线,确认没问题再全量上,别一上来就硬刚。