最近在试着用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试试,长文本多的话把max_length砍到384,batch size设2加8步累积,我这么调稳得很。
rank 32对8B模型确实偏高了,试试16或8,同时把梯度累积步数降到2,应该能稳住loss。
batch size 2加上gradient accumulation设成8步试试,等效batch size就是16,我试过这个组合在4090上挺稳的,loss抖动也小很多。另外rank=32对于LoRA来说确实偏高了,降到16或8,显存压力能降不少,而且效果差别不大。长文本500 tokens其实还好,但如果截断策略没处理好,padding太多也会浪费显存,可以检查下tokenizer的max_length设置。
rank设16试试,长文本多的话batch size别硬撑,累积步数改成2步加梯度裁剪更稳。
显存瓶颈其实挺常见的,两张4090跑8B模型batch size设4确实容易爆,建议你先降到2,然后把gradient accumulation设到8,这样等效batch就是16,收敛会稳很多。LoRA rank 32对于客服任务可能偏高了,可以试试降到16或者8,显存压力会小不少。另外你提到长文本多,平均500 tokens的话,可以检查一下数据里有没有特别长的尾巴,适当截断到384或者512能省不少显存。
两张4090跑8B模型batch size 4炸显存太正常了,毕竟长文本500 tokens确实吃显存。建议先试试rank降到16,LoRA参数量少一半,显存压力会小很多。梯度累积4步loss抖大概率是学习率没配合降,试着把lr调到1e-4到5e-5之间,或者先用cosine schedule稳住。另外可以检查下数据集里是不是有特别长的样本,跑之前按长度排序或者截断到384 tokens,能省不少显存。你目前一个epoch十几小时确实离谱,调完这些应该能压到五六小时。
rank设16试试,长文本多的话batch size可以降到1然后梯度累积8步,收敛会稳很多。
rank降到8试试,长文本多的话batch size设1累积8步,loss能稳很多。
batch size设2加8步累积,rank降到16,长文本截断到384,我这样跑稳得很。
4090双卡跑8B模型,batch size设2其实挺正常的,我自己的经验是单卡4090跑7B模型batch size只能塞下1,你这双卡能跑2已经不错了。显存溢出不一定全是batch size的锅,长文本平均500 tokens确实很吃显存,你可以试试把max_length砍到384或者256,很多客服问答其实用不了那么长上下文,效果不会有太大影响。
关于gradient accumulation导致loss抖动,我怀疑不是累积步数本身的问题,而是你learning rate没跟着调。累积4步等效batch size=8,但lr如果还保持在原来batch size=2的水平,相当于有效更新步数变少了,梯度更新幅度相对变大,自然容易抖。建议把lr降到原来的1/2或者1/3试试,或者用余弦退火调度器慢慢降。
LoRA rank设32对8B模型来说确实偏高了,尤其你那才一万多条数据,rank太高容易过拟合不说,还白白增加显存占用。我自己试过8到16之间效果就很好了,32带来的提升微乎其微,但训练时间会明显变长。你可以先降成8跑个几百步看loss趋势,如果收敛稳定再慢慢加。
另外提醒一点,你双卡4090记得用DeepSpeed ZeRO stage 2或者FSDP,光靠原生PyTorch的DataParallel显存利用率很差。我试过同样配置,用上ZeRO之后batch size能直接从2提到4还不崩。
rank 32对8B模型确实偏高了,试试16,batch size用1加8步累积,loss会稳很多。
4090双卡的话batch size=2加gradient accumulation=8试试,我跑类似规模的数据集这么设显存稳得住,loss也平滑。你rank设32对于8B模型确实偏高了,降到16甚至8效果差不多但省不少显存。长文本500 tokens不算特别夸张,但建议检查下有没有超长尾巴,把max length设成512或者用动态padding能缓解。收敛抖可能跟学习率也有关,LoRA一般用1e-4左右,你可以调低点看看。
我也是两张4090跑类似任务,batch size设2加gradient accumulation 8步试过,loss确实抖但调低学习率到1e-4就好很多了。你rank 32对8B模型确实偏高了,尤其数据量才一万多,降到16甚至8试试,显存压力会小不少。另外平均500 tokens不算长,但建议把超过1024的样本截断或过滤掉,能省不少显存。实在不行可以先用deepspeed stage2开offload,虽然慢点但至少不崩。
同样4090双卡,之前也踩过这个坑。你设rank32确实偏高了,LoRA在7B以上模型里rank8到16就够用,显存能省出一大截。另外平均500 tokens确实容易爆,建议把超长文本截断到384或512,配合gradient accumulation时把learning rate适当调低到1e-4左右,loss抖动会好很多。
4090双卡跑8B用bs=4溢出有点反常,你确认下是不是忘了开gradient checkpointing?那玩意能省一半显存。另外rank=32对8B确实偏高,试试8或16,loss震荡可能跟这个有关,长文本倒不是主因。
累积步数4没问题,但建议配合lr调度器用warmup,前几百步把学习率从0慢慢拉起来,能压住前期震荡。我跑类似任务一般bs=2+累积8,等效16,效果比bs=4+累积4稳。
你这数据量其实不大,可以试试直接全量微调部分层,或者把max_seq_len砍到256,客服问答一般用不到500token。先调通再追求指标,别死磕参数。
4090双卡跑8B LoRA,batch size=4溢出很正常,你这数据量其实不用追求大batch,试试batch size=1加8步累积,等效8的batch对收敛就够了,loss抖大概率是学习率太高,降到2e-4左右看看。LoRA rank 32对客服任务有点浪费,16完全够用,省下来的显存还能塞更长上下文。另外平均500 tokens确实偏长,把超过800 tokens的样本截断或者做下清洗,速度能快不少。
双4090跑8B LoRA,batch size=4爆显存挺正常的,毕竟平均500 tokens确实偏长。我建议你把LoRA rank降到16,然后batch size保持2,gradient accumulation设成8,这样等效batch size是16,比你现在4+4的组合稳很多。另外loss抖不一定全是batch的问题,检查下学习率是不是太高了,LoRA微调一般1e-4到2e-4就够,你试试把lr调到1e-4,warmup steps设个200,收敛会平滑不少。
4090双卡跑8B LoRA,batch size=4溢出很正常,8B模型就算LoRA也得吃不少激活内存,我一般开gradient checkpointing,batch size直接拉满到8,累积步数设2就行。你loss抖大概率不是累积步数的问题,看看学习率是不是太高了,LoRA rank 32其实不算大,但长文本多的话可以试试把max_seq_len砍到384,数据里超过这个长度的直接截断,训练速度能快不少。另外你两张卡有没有开DeepSpeed ZeRO-2?光靠默认的DDP显存利用率上不去。
4090双卡跑8B其实不用追求大batch,单卡batch size=2配合4步累积在数学上等效于8,但loss抖大概率是lr没跟着调,累积步数翻倍的话lr最好也相应调大一点。另外LoRA rank=32对8B来说确实偏高了,尤其你数据量才一万多,降到8或者16试试,效果可能反而更稳。长文本500 tokens其实还好,但你可以看看是不是padding太多浪费了显存,用动态padding能省不少。实在不行就把序列截断到384,客服场景一般用不到那么长。
我之前也遇到过类似情况,4090双卡跑8B用LoRA其实batch size 2+累积4步是够的,但loss抖大概率不是batch的问题,你试试把学习率降到1e-4以下,另外LoRA rank设成16甚至8就够了,32对客服任务有点浪费。长文本平均500 tokens确实吃显存,可以把max length截到384,或者用flash attention省显存。还有个坑是梯度累积步数多了要配合warmup,不然前期loss必抖,你可以把warmup步数调成总步数的10%再试试。