最近在试着用LoRA微调Llama 3 8B,设备是RTX 4090 24G,数据集大概5万条对话。结果刚跑第一轮batch size设成4就直接OOM了,试了梯度累积和混合精度(bf16),还是撑不住。
我看网上有人说用bitsandbytes量化到4bit能省显存,但我加载模型后生成文本变慢了好多,而且微调完效果有点飘。
想问问大家,除了换硬件,有没有什么实用的显存优化技巧?比如分片加载、offload或者改attention?另外,8B模型用4bit微调会不会损失太多能力?求指个路,有点迷茫。
大佬们,用PyTorch跑Llama 3微调,显存爆了怎么办?
全部回复
共 162 条试试unsloth吧,同样LoRA能省一半显存,速度还快,4bit微调8B效果其实够用。
24G跑8B还爆显存,大概率是序列长度和batch的锅,你这5万条对话里长上下文肯定不少。建议先把max_seq_len砍到2048甚至1024,同时batch size降到1配合梯度累积到32步,显存占用能直接少一半。4bit微调其实没那么玄乎,关键是别用NF4直接用FP4,再加LoRA的r调低到8,效果飘大概率是学习率没跟着调。
分片加载和offload都是治标不治本,真正吃显存的是激活值,试试gradient checkpointing,虽然慢点但能省出3-4G。另外你可以用torch.compile把attention换成FlashAttention,显存能再省一截,但注意别和bitsandbytes的4bit混用,容易出奇怪bug。我上次用QLoRA微调7B模型,4bit下效果比全精度掉了不到2个点,关键看你的下游任务,如果是对话生成这种对细节敏感的,可能要多调几轮。
24G跑8B全量微调确实紧张,但LoRA+4bit不至于这么惨,你试试把batch size降到1,梯度累积设个8或16,这样显存峰值能压下来不少。另外别用bitsandbytes的4bit做训练,推理还行,微调时量化误差会被梯度放大,效果飘很正常,建议用NF4或者干脆8bit加QLoRA。分片加载和offload到CPU挺有用,但注意offload会让速度慢,得平衡一下。想问问你数据集清洗过吗,5万条如果质量参差,也可能导致loss波动大,不如先砍到1万条跑通流程再说。
24G跑8B LoRA其实卡在激活值上,你batch size 4加上5万条数据的长尾效应,OOM太正常了。我试过把seq len限制到2048,配合gradient checkpointing,batch size能塞到8,但速度慢得怀疑人生。4bit微调不是不行,关键是你要用QLoRA那套,把NF4和双重量化都开了,同时学习率调低到1e-4以下,不然确实飘。生成的慢是因为bitsandbytes的4bit反量化有额外开销,你可以考虑把embedding和lm_head留在fp16,只量化attention和mlp,推理快不少。至于offload,我建议优先用CPU offload而不是disk,因为SSD随机读写太慢,但offload后显存确实能省一半,代价是训练时间翻倍。改attention的话,试试torch的scaled_dot_product_attention,配合flash attention v2,能把中间激活值压下去不少,而且基本不掉点。最后关于能力损失,我实测4bit微调后的下游任务分数比8bit低大概3到5个点,但如果你只是做指令跟随,体感差别不大,关键看你的任务类型。
4bit微调确实飘,试试QLoRA加paged optimizer,24G跑8B能稳不少,但速度慢就忍忍吧。
试过把batch size压到2加梯度累积吗?4090跑8B LoRA这配置其实够用,再不行就换序列长度短的子集。
我之前也踩过这个坑,24G跑8B LoRA确实紧巴巴的。你试试把batch size压到1,配合梯度累积到8,然后把序列长度截断到1024,能省不少。4bit微调效果飘大概率是量化后损失了精度,建议用QLoRA的nf4配置,比普通4bit稳一些。实在不行就上Unsloth,优化过的内核能省一半显存,速度还更快。
4090跑8B全参本来就不现实,LoRA已经很极限了。你试试把LoRA的rank降到8,target modules只选q_proj和v_proj,别贪多。显存还爆的话,可以开torch.utils.checkpoint(就是activation checkpointing),用时间换空间。至于4bit微调,我用下来感觉对话类任务还行,但推理速度确实拉胯,建议训练完再转回bf16。
你这配置跑5万条确实硬核,我2万条都得折腾半天。试试把optimizer换成AdamW的8bit版,能省2-3G显存。另外,别用bitsandbytes的4bit,试试GPTQ量化,加载后推理快很多,微调效果也稳。改attention的话,flash-attention 2能省不少,不过得看CUDA版本兼容不。我目前是8bit+梯度累积+batch=2,
24G跑8B LoRA按理说够用,你batch size 4直接OOM大概率是序列长度太长或者数据集里样本长度方差太大,试试把最大长度砍到1024甚至512,同时用torch.utils.checkpoint,这玩意儿能省一半左右显存,代价就是慢点。4bit微调效果飘很正常,QLoRA那套其实对超参很敏感,学习率得调低到1e-4左右,而且要用paged_adamw优化器,不然训练不稳定。你说的生成变慢,是因为bitsandbytes的4bit反量化开销就在那,推理时用8bit会快不少,微调完再转回去就行。分片加载没必要,24G单卡又不是多卡,offload到CPU更慢,不如把attention换成flash-attn,显存和速度都能改善。另外5万条对话对8B来说不算多,你试试只微调前10层或者用rsLoRA把rank加大到64,有时候效果反而更稳。最后,4bit微调能力损失肯定有,但主要影响的是知识保留,对话能力和指令遵循还行,要是任务偏推理还是建议至少8bit。
24G跑8B还爆显存确实难受,不过你试试把batch size压到1,配合8bit的bitsandbytes加梯度累积,我之前7B模型这么干能省出不少空间。4bit微调掉点问题我遇到过,感觉是学习率要调低点,比如2e-4改成1e-4,稳定性会好很多。另外你可以看看Unsloth那个库,专门优化LoRA显存,效果比纯原版强不少。生成变慢可能是量化后推理没走优化核,试下加载时把torch_dtype和device_map都指定好,别让bitsandbytes自动调度。
24G跑8B LoRA按理说够用,你batch size 4爆掉可能是序列长度和attention缓存吃太多,试试把max_seq_len砍到1024或2048,再开gradient checkpointing,显存能省一半。4bit微调确实会掉点,尤其对话数据容易飘,建议用NF4加双重量化,或者直接上QLoRA的paged optimizer,效果稳一些。如果你数据集里长文本多,可以试试flash attention,不光省显存还提速,但得确认你CUDA版本支持。另外别全指望offload到CPU,慢得想哭,优先把batch size压到1加梯度累积,再配合deepspeed stage 3试试。
你的batch size 4在24G上确实太激进了,我之前用4090跑7B LoRA都是batch size 1加梯度累积,先把内存占用压到12G以内再说。4bit微调效果飘很正常,QLoRA对超参和数据集很敏感,建议把学习率调低一半多试几个seed,或者直接用NF4加double quant,比普通4bit稳不少。另外可以试试torch.compile加flash attention,能省下不少显存还提速,offload到CPU虽然慢点但至少能跑起来,先验证流程再优化速度吧。
24G跑8B LoRA其实没那么玄乎,关键在batch size和序列长度,你试试把batch size降到1,然后开gradient checkpointing,这俩搭配能省不少。4bit微调确实会让效果飘,尤其对话数据容易崩,我建议你换NF4加double quant,比普通4bit稳一点。另外你说生成变慢,那是正常的,量化后推理本来就有开销,但微调阶段影响不大,可以分开处理。
24G跑8B LoRA还爆显存,大概率是数据加载和优化器状态在作妖,试试把batch size压到1再加梯度累积,同时开gradient checkpointing,能省不少。4bit微调确实会掉点,尤其对话任务上效果飘很正常,建议用NF4加double quant,生成慢就开torch.compile或者vLLM跑推理。另外你5万条数据其实不算多,可以先用QLoRA跑个2-3个epoch看看趋势,别一上来就追求完美效果。分片加载和offload是最后手段,能不开就不开,速度损失太明显。
4bit微调Llama3确实容易飘,建议试试QLoRA加paged optimizer,显存能压下来效果也稳点。
4bit微调8B确实容易飘,试试QLoRA加paged optimizer,能稳不少。
offload到CPU挺实用,把优化器状态扔过去,显存能省一大截。
24G跑8B LoRA按理说是够的,你batch size 4 OOM大概率是序列长度太长或者数据没做padding截断,先检查下是不是有一批超长样本把显存峰值拉爆了。4bit微调确实会掉点,尤其对话任务对生成质量敏感,我试过NF4+双阶段微调(先冻底层后解冻顶层)能缓解一些,但速度慢是硬伤。分片加载加CPU offload是保底方案,不过offload层数设太多会让训练变得极慢,建议只offload优化器状态。另外你试试torch.compile加flash attention,单卡能省出2-3G,配合gradient checkpointing应该能把batch size提到8。关于4bit损失能力的问题,个人经验是如果微调数据质量够高,效果飘可能不是量化导致,而是学习率没调好,LoRA rank设低点(比如8)用2e-4到3e-4试试。还有个偏门思路是把数据集按长度排序,每个batch内长度接近,这样padding开销小很多,我上次这么做显存占用直接降了15%。最后实在不行就换QLoRA+PEFT的官方示例配置,人家给的默认参数都是调过的,比自己瞎试稳。
4090跑8B LoRA确实紧巴,试试把batch压到1再加8步梯度累积,4bit微调用QLoRA的话效果其实还行。
分片加载加CPU offload能救急,但速度损失你得忍,或者干脆砍数据量到2万条先跑通。
4090 24G跑8B全量微调本来就紧,LoRA加4bit量化是正解,但你说的效果飘大概率是量化参数和LoRA rank没调好,试试把quant_config里ftfy设成True,同时rank降到8,alpha用16,能稳不少。显存还爆的话就把batch size压到1,梯度累积加到32,别硬扛。至于4bit会不会伤能力,看任务,简单指令跟随影响不大,复杂推理确实会掉点,但比OOM强。另外你试试unsloth这个库,专门优化过llama微调显存,同配置能多塞一倍batch。
说实话24G跑8B LoRA按理说没那么容易爆,你batch size 4确实有点猛了,我一般4090上开1加梯度累积8步,效果跟4差不多,但峰值显存能降一半还多。bf16加LoRA本身没问题,关键是你看下是不是把attention的seq len拉太长了,5万条对话如果单条超2k token,那计算图内存会指数涨,建议先截断或者用Flash Attention,这个能省不少。bitsandbytes 4bit那个我试过,加载慢是因为它要反量化做权重更新,而且QLoRA对学习率特别敏感,飘很正常,你得把lr降到1e-4以下再试试。真要省显存,别用offload到CPU,那玩意慢到怀疑人生,不如把优化器状态用Adafactor或者8bit Adam,能省好几G。另外你可以考虑把模型分片到两张卡上,但4090不支持NVLink,通信开销大,效果一般。至于4bit会不会伤能力,我体感是推理任务掉点明显,但微调完做对话生成差别不大,前提是你数据集质量够高。最后建议你开个wandb盯一下每层的显存占用,很多时候是某个embedding层或者norm层在作妖,定位到了就好办。
4bit+LoRA够用,但记得冻结所有原参数只训adapter,效果飘就调低学习率到1e-4试试。
试过unsloth没?LoRA加4bit能省一半显存,效果比bitsandbytes稳不少。