最近在试着用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 条我也遇到过类似的问题,batch size设2其实够用,关键是梯度累积别设太大,4步确实容易让loss震荡,建议先试试2步累积,再配合学习率稍微调低一点。另外rank=32对8B模型确实偏高了,尤其你的数据量才一万条,降到16甚至8效果可能更稳,显存也能省不少。长文本平均500 tokens不算特别夸张,但你可以检查下是不是有超过2048的样本,裁一下能显著降低显存压力。
rank降到16试试,长文本多的话batch size设1加8步累积,效果比硬撑大batch稳得多。
rank降到16试试,长文本多的话batch size设2加8步累积,loss抖就调低学习率到1e-5。
4090双卡跑8B模型,batch size=4溢出挺正常的,毕竟长文本很吃显存。我建议你试试把LoRA rank降到8或16,一般客服场景8就够用了,32确实有点费资源。另外gradient accumulation步数别设太高,2步就差不多了,4步loss抖动大可能是学习率没跟着调低。数据集平均500 tokens的话,可以先用padding截断到512,能省不少显存。
试试把batch size调到1,累积步数设8,rank降到16,长文本可以截断到384 tokens,loss抖动会好很多。
同两张4090在跑,batch size=2加8步累积其实效果更好,你试试把梯度累积步数提到8,学习率降到1e-4,loss抖动大概率是lr太高了。另外rank=32对于8B模型确实偏大,降到16甚至8,显存能省出一大截,收敛也稳得多。长文本500 tokens其实还好,主要问题还是batch和lr没配合好。
rank降到16,累积步数调2试试,长文本多的话预处理切成512token以内能省不少显存。
老实说你这情况我太熟了,之前微调Llama2的时候也卡在显存和速度的夹缝里。我个人觉得LoRA rank设成32对于8B模型来说确实偏高了,尤其你数据集只有一万多条,rank 8到16其实就够用,rank太高反而增加激活显存占用,还会让优化器更难收敛。另外batch size=2加4步累积实际有效batch是8,如果loss抖得厉害,可以试试先固定有效batch为8,然后把梯度累积步数降到2,同时把batch size提到4——但你需要检查数据加载时有没有长尾样本,平均500 tokens不算太夸张,但万一有单条超过2048的,两张4090也得崩。我建议你先把max length限制在1024或768,然后LoRA rank降到16,有效batch size设成8(比如batch size=2,累积4步),优化器用AdamW加权重衰减,学习率调低到2e-4左右,这样一般能稳住loss。如果你还嫌慢,可以试试Q-LoRA,把4-bit量化打开,显存能再省一半,或者用deepspeed的ZeRO-3。另外,你数据集里有没有重复问答对?如果样本分布不均,batch size小的时候梯度噪声大,loss抖动反而正常,可以考虑用动态padding或者按长度分组来减少无效计算。
你这个问题我太有同感了,4090两张跑8B模型batch size确实是噩梦。我试下来rank设16就够用了,32对客服任务有点浪费,反而容易让loss不稳。另外500 tokens不算特别长,但建议把数据集里超过512的样本截断一下,能显著省显存。梯度累积4步没问题,但学习率要适当调低一点,比如从2e-4降到1e-4,抖动会好很多。先试试rank降到16加梯度累积2步,应该能平衡速度和稳定性。
我双卡4090跑类似任务时batch size设的2,梯度累积搞到8步,loss调低学习率到1e-4就稳了,你可以试试。LoRA rank 32对8B模型确实偏高,降到16甚至8能省不少显存,效果差别其实不大。另外长文本多的话,可以把max length卡到384,超出部分截断或分块,不然注意力计算太吃显存。
你这配置跑Llama3-8B确实吃紧,我建议先把LoRA rank降到16试试,8都行,32对8B来说太奢侈了,显存压力主要在参数量上。batch size设2然后用梯度累积8步,这样等效batch size是16,收敛会比4步稳很多,你试试看loss抖动会不会好点。另外平均500 tokens确实偏长,可以检查下数据里有没有特别长的尾巴,适当截断到384或256,能省下不少显存给batch size,epoch时间也能压下来。
我最近也在调类似的配置,两张4090的话batch size 2加梯度累积4步其实挺常见的,但loss抖可能是学习率没跟着调,累积步数增加后实际batch变大,学习率得适当降一点,比如从2e-4降到1e-4试试。另外LoRA rank 32对8B模型确实有点高,尤其你的数据量才一万多,降到16或者8说不定更稳,还能省显存。长文本500 tokens倒还好,但你可以检查下有没有特别长的序列被padding到统一长度,有时候是那个把显存吃爆了。
两张4090跑8B模型确实吃紧,batch size=2加4步累积其实等效batch size=8,loss抖可能不是累积的问题,而是学习率没跟着调低。建议把rank降到8或16试试,LoRA参数量减半能省不少显存,长文本500 token其实还好,但可以试试梯度裁剪加混合精度,我上次微调7B时batch size=2加8步累积效果挺稳的。
-
8B模型在4090上batch size=4确实容易爆,我一般设2再加4步累积,但你的loss抖可能跟学习率有关,试试从1e-4降到3e-5,同时把warmup步数拉长到总步数的10%,收敛会稳很多。
-
长文本占显存很厉害,平均500 tokens的话建议先截断到384或256,如果数据里关键信息大多在前半段,效果损失不大。LoRA rank=32对8B模型偏高,降到16试下,显存能省不少。
-
我猜你用的是AdamW吧?试试把betas调成(0.9, 0.98),weight decay设0.1,配合线性衰减,配合batch size=2+累积4步,我跑类似任务一个epoch能压到6小时左右。
-
另外检查下是不是padding策略的问题,用dynamic padding而不是把所有序列pad到最长,能省30%左右显存。loss抖也可能是梯度裁剪没开,设max_grad_norm=1.0会稳很多。
我也是双4090跑LoRA,你这配置batch size设4崩太正常了,Llama3-8B本身显存占用就高。建议你rank降到8试试,长文本多的话梯度累积步数设2-4就行了,关键是把梯度检查点打开,我这batch size能稳到8。另外loss抖不一定全是batch的问题,检查下学习率是不是太高了,我调到2e-4配合warmup就好很多。
rank设16试试,长文本多的话batch size 2加8步累积,我这样调4090单卡跑得挺稳。
两卡4090跑8B模型batch size设4确实有点极限了,尤其你平均token还不短。我建议先把rank降到16试试,大部分场景下8-16就够了,32对LoRA来说有点浪费参数。梯度累积4步抖动很可能是因为学习率没跟着调,累积步数翻倍的话学习率建议降个20%-30%试试。另外检查下是不是用了flash attention,没开的话显存能差不少,长文本多的话可以试试把max length设成512,有些超出部分直接截断影响不大。
rank降到16试试,长文本多的话梯度累积步数翻倍到8,loss抖动可能是学习率没跟着调低。
rank降到16试试,长文本多的话batch size设2配合8步累积,loss抖动就调低学习率。
显存这块我跟你配置差不多,两张4090跑8B模型,batch size=4确实容易崩,尤其你平均500 tokens算长文本了,我一般单卡设batch=1,两张卡开数据并行等效batch=2,然后用gradient accumulation堆到8步以上,这样实际batch就是16,loss收敛比直接设小batch稳定很多。你那个累积4步抖,很可能是因为累积步数太少,等效batch还是太小,梯度噪声大,建议试试累积8步甚至16步,学习率相应调低一点,比如从2e-4降到1e-4。另外rank=32对8B模型其实不算高,但如果你数据集只有一万条,rank降到16甚至8反而可能更好,防止过拟合也能省显存——我试过8B+Lora rank=8,效果跟32几乎没差别。长文本的话,建议把max_length截到384或512,超出部分直接砍掉或分段处理,不然attention计算量会暴涨。还有个小技巧:用deepspeed stage2或flash attention能显著降低显存占用,我开了之后batch size直接翻倍。你loss抖也可能是学习率或者warmup步数没配好,试试先跑50步看趋势,别急着看整个epoch。