最近在试着用LoRA微调Llama3-8B来做我们公司的客服问答模型,数据集大概一万多条。我显卡是两张RTX 4090,但每次设batch size=4就显存溢出,降到2又跑得慢得要命,一个epoch要十几个小时。我看别人说用gradient accumulation可以等效大batch,但设成4步累积之后,loss曲线一直抖,收敛效果很差。想问问有经验的大佬,这种规模的模型和数据集,batch size和累积步数到底怎么搭配比较合理?还有,是不是我LoRA的rank设太高了(我设的32)?或者数据集里长文本太多(平均500 tokens)导致的?求指教,我调了好几天了,心态有点崩。
用LoRA微调Llama3-8B做客服,batch size设多大才不崩?
全部回复
共 154 条建议rank降到8试试,长文本是显存杀手,gradient accumulation配4步没问题但lr得跟着调小。
你这情况我太懂了,双卡4090跑8B其实挺尴尬的,单卡4batch确实顶不住。建议试试batch size=2加上gradient accumulation=8,等效16的batch,但注意lr要相应调低一点,不然loss抖很正常。另外LoRA rank=32对8B来说确实偏高,砍到16甚至8试试,显存占用和收敛速度都会有明显改善。还有你那500tokens的平均长度不算特别夸张,但建议把超过1024的样本截断一下,padding到统一长度,别让显存浪费在无效token上。
显存溢出大概率不是rank的问题,8的rank对8B模型足够用了,32纯粹是浪费。你试试把batch size降到1,累积步数设8,等效batch size还是8,但每步的显存压力小很多。loss抖的话,检查下学习率,LoRA微调一般1e-4到2e-4就行,太大容易震荡。另外平均500 tokens确实偏长,可以把超过512的样本截断或者做滑窗处理,不然attention计算太吃显存。我跑过类似规模的客服数据,单卡4090 batch=1累积8步,一个epoch大概4小时,你可以参考下。
显存这块,2张4090跑8B模型,单卡batch size=2其实挺正常的,问题大概率出在LoRA rank=32上,试试降到16甚至8,显存和速度都会好很多,效果一般差不了太多。累积步数设4没问题,但loss抖多半是学习率没配合调低,你试试把lr从2e-4降到1e-4左右,同时把warmup步数加长点,让累积的等效batch别太激进。另外平均500 tokens确实偏长,建议把超过512的样本截断或者按长度分桶,padding不要搞到统一长度,能省不少显存。我自己微调过类似规模的客服数据,感觉你现在的瓶颈不是batch太小,而是超参没对齐,先调rank和lr,比纠结batch size更值得花时间。
调成1+8累积试试,loss抖大概率是lr没跟着降,rank16够用了。
4090两张跑这数据量别急,先砍到8累积+cosine调度看看。
4090两张跑8B居然还爆显存?先看看是不是序列长度没截断,500tokens确实偏长了。
4090双卡跑8B其实可以试试把LoRA rank降到16,效果一般不会差太多,显存能省下一大截。你那个loss抖动大概率不是累积步数的问题,而是学习率太高,试试调到1e-4以下,配合warmup几步会稳很多。另外平均500 tokens确实偏长,可以把超过768的样本截断或者按长度分桶,不然padding浪费显存很亏。我上次跑类似数据量是batch size=2、累积8步,效果比4步累积稳,你可以交叉验证下。
4090双卡跑8B,batch2+累积4是标配,你loss抖大概率是lr没跟着调,试试降到2e-4。
rank32确实偏高,客服场景16就够,省下的显存还能把batch提上去。
我之前也遇到过类似情况,两张4090跑8B其实挺吃紧的。你试试把LoRA rank降到16,同时把序列长度截到400以内,显存压力会小很多。至于梯度累积,4步有点激进,改成2步配合batch size=2,loss会稳一些,一个epoch大概能快三分之一。另外检查下是不是用了flash attention,这个对长文本省显存效果很明显。
- 4090双卡跑8B,batch size=4溢出其实挺常规的,你试试把LoRA换成4bit量化加载,显存直接省一半。
- 累积步数4真不算多,loss抖大概率不是累积的问题,你先把学习率降到1e-4以下看看,我上次就是被这坑的。
- 平均500 tokens确实偏长,可以把数据截断到384,效果不会差太多,训练速度能快不少。
- rank=32对于客服任务有点浪费,16就够用了,你那个显存压力有一半是rank和长文本一起堆出来的。
- 我当初调的时候也卡了三天,最后发现把数据集里重复模板去掉,epoch时间直接砍半,你可以检查下数据里有没有大量相似问答。
之前跑类似项目也踩过这个坑,8B模型在4090上单卡batch size=4本来就接近极限了,两张卡用梯度累积没毛病,但loss抖大概率是学习率和warmup没配合好,建议把LR降到2e-4左右试试。LoRA rank=32对这个数据量确实偏高,降到16能省不少显存,而且效果基本不缩水。另外平均500 tokens不算特别长,但如果你有padding到固定长度的话,把max length从2048砍到1024也许能直接塞下更大的batch。
显存溢出大概率不是rank的问题,先试试把gradient checkpointing开了,8B模型开这个能省一半显存,batch size就能回到4。累积步数设4没问题,但loss抖可能是学习率没跟着batch size调,等效batch翻倍后学习率也得相应提一点,不然收敛会不稳。另外你平均500 tokens确实偏长,packing的时候把长样本单独抽出来处理,别硬塞进一个batch里,我上次这么搞完显存直接降了30%。
4090双卡跑8B,batch size=4爆显存大概率不是rank的问题,32对LoRA来说真不算高,你试试把序列长度截到512或者用梯度检查点,能省不少显存。累积步数设4但loss抖,很可能是学习率没跟着调,等效batch变大之后lr得相应提一点,不然收敛就是会不稳。另外你这数据量1万多条,其实可以试试batch size=1加累积8步,效果不一定比大batch差。我上次调类似配置,把max length砍到400,显存直接少用3G,训练速度还快了20%。
我最近也在调类似的配置,不过用的单卡A100,batch size同样卡在4上。你这情况大概率不是rank的锅,32对LoRA来说挺常规的,问题可能出在梯度累积和优化器状态的交互上——累积步数设4但loss抖,我猜是你学习率没跟着调,累积步数翻倍相当于batch翻倍,学习率得相应往上拉一点,不然收敛会特别挣扎。另外500 tokens确实偏长,LoRA对序列长度挺敏感,建议先试试把输入截断到384或者256,很多客服场景其实关键信息都在前面,效果未必差。还有个实操技巧:把数据集里特别长的样本单独过滤掉,或者用动态padding按batch内最大长度算,这样显存利用率会高不少,我之前这么搞直接省了30%的占用。最后实在不行,试试8-bit优化器或者换AdamW的offload版本,能再挤点空间出来。你先试试截断和调学习率,大概率能稳住loss。
说实话你这配置跑不动大概率不是batch size的锅,LoRA rank 32在8B上确实偏高了,降到16试试,显存能省不少。另外平均500 tokens的长文本建议把max_seq_len砍到512,不然attention计算量直接翻倍。梯度累积步数设4没问题,但loss抖大概率是学习率没跟着调,累积步数翻倍学习率也得相应放大,你试试3e-4起步。我先前用单卡4090微调7B,rank16+bs2+累积8步,一个epoch也就四小时,你可以参考下这个组合。
4090上LoRA跑8B本来就不太能直接塞大batch,两张卡的话每张batch size=2再开gradient accumulation=8,效果比你现在稳得多,关键是把学习率降到1e-4左右,别用默认值。你rank=32确实偏高了,8-16就够用,尤其客服这种垂直领域任务,太高反而容易过拟合。长文本这块建议把超过512 token的样本截断或做滑窗切分,不然显存和loss抖动都是它闹的。还有个细节,累积步数不是越大越好,超过8步loss就容易飘,你试试4步但把lr调低,应该能跑通。
4090双卡跑8B还爆显存,大概率不是batch size的锅,你试试把LoRA的target modules从全部线性层改成只冻住q_proj和v_proj,显存能省将近一半。rank=32对于一万条数据确实偏高了,尤其你平均500 token,参数量直接翻倍,建议先降到8看看loss曲线,我上次用rank=16效果跟32差不多但训练快了一倍。另外gradient accumulation设4步没毛病,但loss抖不代表收敛差,你确认下是不是learning rate没跟着调,等效batch翻倍的话LR得相应调大,不然优化器步长太小会来回震荡。还有个小技巧,把序列长度从512砍到384,哪怕截断一点长文本,对显存和速度的提升都巨大,客服数据其实很少有非要500 token才能答的。你要是急用,可以先试试DeepSpeed ZeRO stage 2加offload optimizer,双卡4090跑8B完全够,batch size直接上8都没问题。最后问下,你用的transformers版本是多少,新版对Llama的attention实现优化过,老版本会多占不少显存。
这个思路不错,收藏了。
你这配置跑8B的LoRA其实挺尴尬的,两张4090互联带宽不够,batch size稍微大点就爆显存很正常。我建议你先别盯着rank=32,那个对8B来说确实偏高了,降到16试试,效果损失不大但显存压力能小不少。至于gradient accumulation,loss抖不一定全是它的锅,你试试把学习率调低点,比如从2e-4降到1e-4,同时把warmup steps拉长,我遇到过类似情况这么搞就稳了。另外平均500 tokens确实偏长,你可以试试把超过512的样本截断或者做一下长度分布统计,很多客服问答其实关键信息就集中在开头和结尾。还有个骚操作是先用4bit量化加载模型再LoRA,显存能省一半,速度反而可能更快。你那个一个epoch十几个小时也太夸张了,我怀疑是不是数据加载或者并行策略没优化,检查下DataLoader的num_workers和pin_memory,还有模型并行是不是用的device_map="auto"。最后,一万多条数据用8B其实有点大材小用,如果效果一直不理想,可以先用7B或者更小的模型跑通流程再换回来。
说实话你这配置跑8B不该这么痛苦,问题大概率不在batch size。4090 24G显存两张,LoRA rank 32确实偏高了,砍到8或者16试试,显存能省不少。另外loss抖不一定是累积步数的事,你检查下学习率,LoRA微调一般1e-4到2e-4就够,太高了累积梯度也会震荡。长文本倒是会影响显存,但500 tokens算正常范围,可以试试把max length裁到384,很多客服问答用不到那么长。我自己的经验是batch size 2加累积8步,效果比batch 4加累积4步稳定,你反着调一下说不定就顺了。