最近在试着用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 条4090双卡跑8B LoRA,batch size=2加梯度累积8步其实更稳,你设4步累积等效batch才8,对一万条数据来说偏小了。rank32不算高,但长文本多的话,可以把max_seq_len砍到384,显存立刻松快不少。loss抖大概率是学习率太高,试试降到1e-4配warmup,另外确认下是不是序列padding没开attention mask,这个也容易让loss乱跳。实在不行就换QLoRA,4bit下batch size能翻倍。
同款配置,我之前微调Qwen也遇到过这问题。你先别急着调累积步数,4090跑8B模型batch size=4爆显存挺正常的,LoRA虽然省显存但激活值照样吃内存。我后来是把sequence length从512砍到384,batch size直接上8,一个epoch从14小时降到6小时,loss曲线反而稳了。你那个rank=32对8B模型确实偏高,8到16基本就够用,而且你平均500 tokens的话,长文本里pad太多会拉低有效batch,试试开gradient checkpointing加8bit优化器,显存能省将近一半。至于累积步数,我建议先固定batch size=1,累积步数设8,等效batch=8跑通流程,然后再慢慢往上加,别一上来就追求大batch。另外你loss抖可能不是batch的问题,是学习率太高,LoRA微调一般3e-4起步,你这个数据量1万条其实不大,rank=16加4步累积应该能压住。最后问一句,你数据里是不是有很多重复的客服模板?如果类别分布不均,batch大小影响会特别明显,试试把数据shuffle打散再分段采样。
显存爆了先别调rank,试试gradient accumulation配大batch,loss抖大概率是lr没跟着调低。
双4090跑8B其实挺尴尬的,显存卡在中间。你试试batch size=2再加8步累积,等效batch 16,loss抖多半是lr没跟着调,把学习率降到1e-4左右看看。LoRA rank 32对8B确实偏高了,降到16能省不少显存,效果基本不掉。另外你长文本多的话,可以把max_length截到384,客服场景没必要吃满500。我之前跑类似数据,1万条一个epoch大概4小时,你参考下。
试试把rank降到16,累积步数改2,4090上batch大小看显存余量调到6左右,loss抖多半是lr太高。
显存不够就把seq_len砍到400,rank降到16,累积步数4没问题,loss抖多半是lr太高了。
4090双卡跑1万条数据十几个epoch正常,先砍一半数据调通流程再说。
4090双卡跑8B用LoRA,batch size=4溢出挺正常的,毕竟8B的激活值在长文本下很吃显存。我建议batch size=2,梯度累积设8步,等效batch size=16,loss抖动大概率是学习率没跟着调,试试从2e-4降到5e-5。rank=32对1万条客服数据其实偏高,降到16甚至8效果可能更稳,而且训练速度能快不少。长文本这块你可以先按tokens做动态padding或截断到512,省下的显存能换更大batch。另外你用的是PEFT还是transformers自带脚本?有时候显存溢出是优化器状态没走LoRA,检查下target_modules有没有设对。
4090两张跑8B还爆显存,大概率不是batch size的锅,你查下是不是把序列长度设成512以上了,LoRA虽然省了梯度但激活值照样吃显存,平均500 tokens的话建议把max length砍到384,数据里长的截断短的补齐,显存能省将近一半。至于gradient accumulation导致loss抖动,我怀疑你learning rate没跟着调,等效batch翻四倍的话LR也得相应往上拉,不然梯度估计噪声太大,试试从2e-4起步往5e-4调,同时warmup步数加到200。rank=32对8B模型确实偏高了,你这数据量rank=8就够用,降下来显存和速度都会改善,而且泛化反而更好,可以对比下验证集loss。另外你一个epoch十几个小时太离谱了,检查下是不是没开flash attention和bf16混合精度,两张卡数据并行的话batch size=2每卡其实等效4,配合gradient checkpointing基本不会炸。长文本多的话建议先按长度排序做动态batching,同batch内序列长度接近能大幅减少padding浪费。最后别死磕单epoch,先跑个2000步看loss趋势,调参效率比跑完一轮高多了。
两个4090跑8B其实没必要硬上大batch,我试过类似配置,bs=2加8步累积效果就挺稳的,你那个4步累积loss抖大概率是lr没跟着调,累积步数翻倍学习率也得相应提一点。LoRA rank 32对一万条数据确实偏高,降到16能省不少显存,而且长文本场景下把max_seq_len裁到384或512效果反而更好。另外你确认下是不是用了梯度检查点,这个开关能省将近一半显存,bs=4应该就能塞下了。
4090双卡跑8B LoRA,batch size=4溢出有点反常,你检查下是不是seq_len没截断,500 tokens确实偏长,建议统一padding到256或512,显存能省一大截。梯度累积设成8步试试,等效batch=16,loss抖大概率是lr太高了,降到1e-5左右配合warmup会稳很多。LoRA rank=32对8B来说不算高,但你这数据量其实16就够用了,效果不会差太多,省下来的显存还能加大batch。另外你试试DeepSpeed stage 2,双卡能直接塞下batch=8,比累积靠谱。
双4090跑8B LoRA,batch size=4爆显存挺正常的,我一般单卡开8就极限了,你试试batch size=2加gradient accumulation=8,等效16的batch,loss抖大概率是lr太高,降到1e-4左右再看看。rank 32确实偏大了,客服任务8到16足够,长文本多的话可以把max length截到384,处理速度能快不少。另外你数据一万条其实不算多,十几个小时一个epoch确实离谱,检查下是不是dataloader的num_workers没调好,或者attention实现没开flash attention。
4090双卡跑8B长文本本来就吃紧,你试试把max_length砍到384或者256,八成能解决一大半。LoRA rank 32对于一万条数据确实偏高了,降到16甚至8,loss会更稳,泛化也未必差。至于gradient accumulation,步数设2就够,4步的话学习率得相应调低,不然优化器步长和真实batch对不上,loss当然抖。我上次微调类似体量数据,batch size=2加2步累积,lr用2e-4,效果挺稳的,你可以参考下。
4090双卡跑8B LoRA,batch size=4爆显存其实挺正常的,毕竟8B模型光权重就占16G,加上激活值和梯度,两张卡总共48G显存真不算宽裕。你试试单卡batch size=1,然后gradient accumulation设成8,这样等效batch size=8,但显存压力小很多,速度反而可能比你现在batch size=2还快。loss抖动不一定是累积步数的问题,你先确认一下learning rate,LoRA微调一般建议1e-4到3e-4,太大就会导致loss震荡。另外rank=32对8B模型来说确实偏高了,尤其你这个数据量只有一万多条,rank=8到16完全够用,太高反而容易过拟合还费显存。长文本平均500 tokens也是个关键点,你可以试试把输入截断到256或者384,或者用Flash Attention减少显存占用,能省不少。我之前微调类似规模的模型,都是单卡4090,batch size=1,累积8步,跑1.5万条数据一个epoch大概三四个小时,你参考下这个配置,应该能快很多。还有个小技巧,把LoRA只挂在attention层别挂MLP层,能再省点显存,收敛也会更稳定。
4090双卡跑8B,batch size=4溢出很正常,你降到2然后累积4步等效batch=8其实方向是对的,但loss抖大概率是lr没跟着调,累积步数翻倍的话lr也得相应提一点。LoRA rank=32对8B来说确实偏高,试降到16,省显存还能加快速度,效果一般不会差太多。另外你平均500 tokens确实偏长,可以试试把超过512的截断,或者用打包方式把短样本拼一起,能有效提吞吐。先别急着崩,这配置跑通完全没问题,就是得磨参数。
gradient accumulation不是让你直接改累积步数就完事的,关键得配合学习率调整,你从bs=4降到等效bs=8,lr不跟着降loss肯定会抖。另外LoRA rank=32对8B模型确实偏高了,试试16或者8,显存压力小很多,收敛也稳。长文本这块建议把超过512 token的样本截断或者做滑窗,不然attention计算量太夸张。我自己的经验是双卡4090开gradient checkpointing + LoRA rank=16 + bs=2累积8步,一个epoch大概4小时,loss曲线很平滑。你那个batch size=4溢出大概率是没开gradient checkpointing,开了之后显存占用能降一半还多。
看到你这个配置和报错,我第一反应是大概率问题不在batch size本身,而在你那个rank=32上。8B模型用LoRA,rank16一般就够用了,32的参数量直接翻倍,再加上4090的24G显存其实有点尴尬,单卡跑4的batch确实容易爆。我建议你先把rank降到8试试,loss抖动可能跟这个关系更大,毕竟你数据集才一万条,rank太高反而容易过拟合。
gradient accumulation确实能等效大batch,但4步累积配合batch=2,等效batch=8,这个数值对8B模型来说其实不小了,loss抖可能是学习率没跟着调,你试试把学习率降到1e-5以下,或者用cosine调度带warmup。长文本这块,平均500 tokens确实偏长,但更关键的是你padding有没有做对,如果每个batch都按最长序列填充,显存肯定浪费很多,建议开dynamic padding。
另外你两个卡是用的单卡数据并行还是张量并行?如果只是数据并行,那得确认一下是不是显存没均摊,有时候代码里没设device_map,两张卡等于白挂。我自己的经验是,一万条数据、8B模型,LoRA训练一个epoch控制在4-6小时比较合理,你跑十几个小时肯定哪里不对劲,建议先看下显存利用率和GPU占用是不是真的打满了。
还有个思路,你试试把序列长度截到384或512,太多长样本对客服任务未必有用,反而拖慢速度。调参这事别急,我上次也卡了一周,最后发现是tokenizer的pad_token没设,导致每个batch都在重新分配内存。你先从rank和lr这两点改起,大概率能解决一大半问题。
4090双卡跑8B还不敢开梯度累积4,试试bf16+8bit量化,rank砍到16,loss抖大概率是lr太高了。
说实话你这配置跑8B模型batch size=4爆显存太正常了,4090单卡24G,LoRA虽然省了全量微调的内存,但激活值和attention中间变量照样吃显存,我一般单卡就设batch size=1或2,靠gradient accumulation凑到8或16。你设rank=32对8B模型来说确实偏大了,尤其数据量才一万多条,rank=8到16就够用,不然LoRA引入的额外参数反而让优化更难,loss抖可能跟这个也有关系。另外你说长文本平均500 tokens,那等于每个样本的显存占用比普通对话高出一大截,建议先把输入截断到384或256试试,或者用flash attention和梯度检查点,能省不少内存。gradient accumulation设4步,等效batch size=8其实没问题,但loss抖得厉害,我怀疑你学习率没跟着调,等效batch变大后学习率得稍微降一点,比如从2e-4降到1e-4。还有你一个epoch十几个小时,是不是没开bf16混合精度?4090对bf16支持很好,开了之后速度能翻倍。最后建议你用deepspeed stage 2或者zero offload,两张卡可以张量并行,batch size=2每卡,累积4步,效果应该比你现在稳得多。
看到你4090两张还爆显存,八成是序列长度和注意力计算吃爆了,平均500 tokens其实不小。我建议你先试试batch size=2加上gradient accumulation=8,等效batch 16,但学习率得相应调低一点,比如从2e-4降到1e-4,loss抖多半是学习率没跟着调。LoRA rank 32对8B模型确实偏高,降到16试试,很多场景下8-16效果反而更稳。另外可以开gradient checkpointing,能省不少显存,虽然慢一点但至少不崩。
4090双卡跑8B其实挺尴尬的,单卡16G显存LoRA也得精打细算。你试试batch size=1,gradient accumulation设8,等效batch=8,loss抖动可能是lr没跟着调,accumulation步数变大后学习率得适当降一点。另外rank=32对8B确实偏高,砍到16试试,省下的显存足够你开长文本截断到512,你这平均500tokens确实容易爆。我之前调类似任务,lr=2e-4配warmup ratio=0.1,loss曲线就稳很多,你可以参考下。