最近想用LoRA微调一个LLaMA-3-8B做代码补全,数据集大概几万条。我用的是QLoRA,量化到4bit,但batch size设为4就OOM了(单卡A100 80G)。我看有些教程说可以跑16甚至32的batch,是我哪里姿势不对吗?还是说必须要用gradient checkpointing?我用的transformers+peft,序列长度设了2048。另外想问一下,微调代码模型和微调对话模型在超参上有什么大的区别吗?新手入坑,求大佬指点。
用LoRA微调LLaMA-3,显存总爆,是我的batch size设太大了吗?
全部回复
共 177 条80G还爆基本是没开gradient checkpointing,开了之后8的batch随便跑,32得看显存余量。代码模型学习率调低点,1e-4到5e-5之间更稳。
80G跑4的batch就爆,大概率不是batch size的锅,你序列长度2048加上4bit量化后,激活值才是吃显存的大头,开gradient checkpointing能省一半以上。代码补全任务跟对话微调差别挺大的,代码数据通常更吃长序列和更低的学习率,建议试试1e-5起步,另外可以看看flash attention有没有开,这个对显存优化也很关键。你用的是peft的默认配置吗?LoRA的rank和target modules可能也得调一下,有时候把attention层的q、k、v都加上反而比只加q、v更省显存。
开gradient checkpointing,显存能省一半,batch 4配A100肯定够,代码模型学习率调低点试试。
80G还爆基本不是batch size的问题,你序列长度2048加上4bit量化后激活值才是大头,开gradient checkpointing能省一半以上。代码补全任务建议把学习率调低到1e-4到2e-4之间,warmup比例可以比对话模型高一些,因为代码分布更陡。另外你试试packing或者把序列截到1024,很多代码token其实没那么依赖长上下文。
A100 80G跑8B模型4的batch按理说不该爆,除非你开了eval模式或者dataloader里有额外显存占用。你检查下是不是transformers版本和peft的兼容问题,有时候旧版本会把LoRA的weight缓存到显存里。代码微调我习惯用cosine schedule加5%的warmup,但千万别用对话模型那套0.1的dropout,代码任务很容易欠拟合。
开梯度检查点吧,batch4爆显存不正常,A100跑8b qlora上16没问题,序列长度2048也还好。代码模型lr可以调低点试试。
80G跑4batch还爆,肯定不是batch的锅,你试试开gradient checkpointing,序列长度砍到1024。代码模型学习率一般比对话模型低,5e-5到2e-4区间多试几次。
开gradient checkpointing吧,显存能省一半,batch 8应该稳了。代码模型学习率可以比对话模型低点,1e-4试试。
80G还爆基本就是序列长度和attention的锅,开gradient checkpointing至少能省一半。代码补全建议lr调低点,数据少的话epoch别贪多。
你这配置单卡A100 80G跑4bit QLoRA,batch size=4就OOM确实不太正常,大概率是没开gradient checkpointing,加上序列长度2048对激活内存压力挺大的。可以先试着开一下gradient checkpointing,然后把batch size降到2,用梯度累积到16,显存占用能降一大截。另外代码补全任务跟对话微调差别挺大,学习率建议调低一点(比如1e-4到2e-4),max length其实可以砍到1024,代码不像聊天需要那么长的上下文,能省不少显存。你用的是transformers的trainer还是自己写的训练循环?有时候PEFT的默认配置会在某些层上额外开dropout,也会吃显存。
说实话batch size 4在80G上OOM不太正常,除非你的序列长度实际远超2048,或者数据集里有些样本被pad得很长。QLoRA本身已经省了很多显存,但如果你没开gradient checkpointing,那激活值照样能吃满整个显存,这个开关在peft里默认是关的,强烈建议打开,基本能把激活显存砍掉一半以上。另外你试试把4bit的double quantization打开,还有attention实现换成flash attention,这几个组合下来batch size 8应该没问题。至于教程里说16甚至32,那多半是用了更短的序列长度或者小模型,LLaMA-3-8B在2048长度下真没那么轻松。代码模型和对话模型超参上,代码补全通常学习率要低一点,比如1e-4到2e-4,warmup比例也小,因为代码数据分布更稳定,不像对话那么多样。你还可以考虑用packing把不同样本拼到接近2048,减少padding浪费,但注意别让模型跨样本学出奇怪关联。最后确认下你的数据加载有没有做streaming,几万条其实不大,但预处理不当也可能导致内存碎片。
80G还OOM大概率不是batch size的锅,你序列长度2048加上4bit量化后激活值才是大头,开gradient checkpointing能省一半以上显存,batch 4完全够用。代码补全和对话微调差别挺大,代码任务学习率可以稍微调高一点(比如2e-4),但warmup要拉长,因为代码分布和自然语言差太多,收敛更慢。另外你试试把attention实现换成flash attention,显存占用能再降一截,还有你数据集的padding策略是不是没设对,动态padding能省不少浪费。
80G单卡跑4bit的8B模型,batch4就爆肯定不正常,问题大概率出在序列长度和attention计算上,2048的上下文对代码补全来说确实吃显存。gradient checkpointing建议直接开,能省一半以上显存,但代价是训练速度慢30%左右,你如果觉得能接受就开。另外代码模型和对话模型的超参差别挺大的,代码任务通常需要更小的学习率(1e-4到2e-4),而且warmup步数要拉长,因为代码分布更陡,步子太猛容易崩。你试试把batch降到2,同时开gradient checkpointing,再把序列长度砍到1024,应该能稳跑,后面再逐步调上去。
A100 80G跑4bit还爆显存,八成是序列长度和attention缓存吃满了,开gradient checkpointing能省一大截,batch先降到2试试。
我也是从这坑里爬出来的,8B模型用QLoRA跑4bit,A100 80G按理说batch size 4不该爆,你查下是不是梯度检查点没开,这玩意儿在长序列下省显存效果极其明显,开了之后我2048长度能跑到8甚至12。另外注意一下是不是把输入标签也一起算进显存了,代码补全任务如果序列长度固定2048,实际有效token可能远小于这个,但显存还是按最大长度分配的。gradient checkpointing不是可选项,是必须项,尤其你数据量几万条,不开的话肯定难搞。模型方面,代码微调跟对话模型区别挺大,代码任务学习率可以稍微调低点,比如1e-4到2e-4,warmup steps也要长一些,因为代码分布更尖锐,容易过拟合。还有个容易忽略的点,你peft里target_modules要选对,LLaMA-3的attention层和mlp层都加上,只调默认的q_proj和v_proj会省显存但效果差很多。最后建议你试试把序列长度砍到1024跑一轮看显存峰值,如果还是爆,那就是别的问题了,比如transformers版本和bitsandbytes的兼容性,有时候是库的bug。
80G都能爆,大概率不是batch size的锅,你的序列长度2048加上4bit量化后中间激活值才是大头。gradient checkpointing必须开,这能省掉一大半显存,batch size开到8甚至16都不成问题。代码模型和对话模型超参差异挺大的,代码补全一般学习率可以调高一点(2e-4左右),但warmup steps要加长,因为代码分布更陡峭。你试试packing(把短样本拼到2048)没?几万条数据如果不packing,很多序列实际很短,浪费严重。
顺带问下,你用的 transformers 版本是新的吗?老版本对QLoRA的显存优化有bug,建议升到4.40以上。我上次用同样配置跑7B模型,开了gradient checkpointing后,batch 16都稳得很。代码补全还有个坑,就是loss要只算在completion部分,不然会把prompt的交叉熵也算进去,模型容易学歪。
A100 80G跑4bit的8B模型,batch size=4按理说不该炸,大概率是序列2048太长,加上没开gradient checkpointing,激活值占了大头。先把checkpointing打开,batch size调到8试试,显存占用能掉一大截。代码补全和对话微调确实不太一样,代码任务学习率可以稍微调低点(比如1e-4到2e-4),warmup步数也短一些,因为代码分布更结构化,收敛快。另外你用的transformers版本和peft版本如果比较新,可能默认把attention实现改成flash attention了,没装对应kernel也会显存翻倍,这个坑我也踩过。
80G单卡跑4bit的8B模型,batch4就爆肯定不正常,问题大概率出在没开gradient checkpointing,这玩意儿对显存占用影响比量化还大,开一下能省一半以上。另外你sequence length 2048对代码补全来说有点长,如果数据集里大部分样本没那么长,可以试试动态padding或者截到1024,能省不少。代码模型微调跟对话模型确实不太一样,代码任务通常学习率要低一点(1e-4左右),而且warmup步数可以短一些,因为代码分布更集中。你参考的那些教程可能是用了flash attention或者已经开了checkpointing,别只盯着batch size看。
说实话80G跑4bit的8B模型还爆,八成不是batch size的锅,你查下是不是序列长度2048配合打包逻辑把显存撑爆了。我猜你直接用了padding到2048,几万条代码里很多样本根本没那么长,这样浪费特别狠,试试看能不能动态padding或者按长度分组。gradient checkpointing肯定要开,LoRA虽然省了主干优化器状态,但激活值还是全量存的,不开的话batch size 4在80G上确实悬,开了之后16应该没问题。另外你提到代码补全,这个任务和对话微调差别其实挺大的,代码数据里长依赖特别多,学习率建议比对话任务低一点,比如1e-4到2e-4,warmup比例也可以稍微拉长。还有个坑,代码模型微调时经常会把特殊token比如缩进换行处理得不好,你检查下tokenizer是不是把空格折叠了。最后问一句,你数据集里有没有做过去重和过滤,代码数据质量比数量重要多了,不然微调完反而可能变笨。
试试gradient checkpointing,再配paged optimzer,8batch稳得很。序列长度2048其实挺吃显存的。
8B上QLoRA开2048序列长度,batch 4爆显存挺正常的,很多教程都是拿短序列或者小模型吹的。gradient checkpointing基本是必须开的,开了之后省一半还多,另外可以试试把batch降到2然后梯度累积拉高,效果差不多。代码补全和对话微调超参差异挺大的,代码任务学习率可以稍微调低点,warmup比例也建议调小,因为代码分布和自然语言不太一样,我试过直接用对话那套参数效果会飘。你数据集几万条其实不算大,可以先用几千条跑一版看loss曲线,确认没问题再全量上。