最近想用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设太大了吗?
全部回复
共 179 条80G跑4的batch还爆,大概率不是batch的问题,你检查下是不是序列长度2048加梯度检查点没开,这俩叠加显存消耗很夸张。我跑7B模型时开gradient checkpointing能省一半以上显存,然后batch直接翻倍。代码模型和对话模型超参其实差别挺大,代码补全任务建议学习率调低点(1e-4左右),warmup步数拉长,另外把max_grad_norm设到0.3试试,收敛会稳很多。另外你数据几万条的话,不用一上来就全量微调,先拿几千条跑个实验,确认不爆显存再扩,省得反复调参浪费时间。
8B上LoRA还爆显存,大概率不是batch size的锅,你试试开gradient checkpointing,这玩意儿能省一半多显存,4 batch直接变8没问题。序列长度2048对代码补全来说确实偏长,可以砍到1024试试,代码的局部依赖没那么强。代码模型微调学习率一般比对话模型低一点,1e-4到2e-4之间比较稳,对话模型可以稍微激进些。另外你确认下是不是把全量参数都冻住了,有时候peft默认配置会漏掉某些层。
80G跑4的batch都爆,大概率不是batch size的锅,你查下是不是序列长度2048加上4bit量化后中间激活值没释放干净。gradient checkpointing基本是必须开的,不然长序列下激活值能把显存吃穿,开了之后batch翻倍没问题。代码模型和对话模型超参差异还挺大的,代码补全一般学习率可以调高一点,warmup步数也短些,但主要还是看你的数据分布。你试试把gradient checkpointing开了再把batch调到8,应该能稳。
80G跑4的batch就爆,大概率不是batch size的锅,你查下是不是序列长度2048配合4bit量化后,激活值吃满了显存。gradient checkpointing必须开,这玩意儿能省一半以上显存,代价就是慢点。代码补全和对话微调差别挺大的,代码任务学习率可以调低点,比如2e-4到5e-4,warmup比例也建议拉长,不然loss容易震荡。你可以试试把batch降到2,开checkpointing,然后梯度累积到8,效果应该比硬怼batch size稳定。
80G的卡跑4bit的8B模型,batch=4就爆确实不太正常,我自己用同样配置跑7B,batch=8加gradient checkpointing都没问题。你试试把gradient checkpointing打开,这个基本是标配了,能省不少显存,另外注意一下是不是序列长度2048里有些样本特别长导致padding过多,可以试试动态padding。代码补全和对话模型超参确实有区别,代码任务通常学习率可以稍微高一点,warmup步数也短一些,但主要还是看loss曲线调,别太迷信教程里的数值。
这情况太典型了,A100 80G跑4bit还爆显存,八成不是batch size的锅,是序列长度2048配合gradient checkpointing没开导致的。你试试开一下gradient checkpointing,batch size可以不动,显存能省一半以上。代码补全和对话模型超参确实有区别,代码任务学习率可以稍微调高一点(比如2e-4),warmup steps也要跟上,因为代码数据分布更陡峭。另外你几万条数据的话,其实可以试试paged optimizer,能压一点显存峰值。
80G跑4的batch还爆,大概率不是batch size的锅,你序列长度2048加上4bit量化后激活值还是很吃显存,gradient checkpointing必须开,开了之后8甚至16都能跑。代码模型和对话模型超参确实有区别,代码补全任务学习率可以稍微调高一点,1e-4到2e-4之间试试,warmup steps也适当加长,另外你数据集几万条的话,epoch数别贪多,1-2个就够,多了容易过拟合。对了,你检查过transformers版本没,老版本对LLaMA-3支持有bug,也可能导致显存异常。
A100 80G跑4bit的QLoRA还爆显存,大概率不是batch size的锅,sequence length 2048加上attention的显存占用其实很夸张。gradient checkpointing基本是必须开的,能省下将近一半的显存,然后把batch size调到8试试。代码补全和对话微调最大的区别在learning rate和warmup,代码任务一般需要更小的lr(1e-4左右)和更长的warmup,因为代码分布更陡峭。还有个坑是数据格式,代码补全最好把文件路径和上下文拼成一段,别用对话模板那套。
说实话80G跑4bit的8B模型batch size4就爆,大概率不是batch的锅,而是序列长度和attention的计算量在作祟。2048的序列长度对代码补全来说其实挺长的,尤其如果数据里有很多长函数或者大文件,KV cache会吃掉大量显存,你可以先试试把序列砍到1024或者512,很多代码补全任务根本不需要那么长的上下文。gradient checkpointing肯定要开,这玩意儿能省一半以上的激活显存,开了之后batch size翻倍基本没问题,代价就是训练慢个20%左右,但总比OOM强。另外你那个教程说能跑16甚至32,人家大概率是用了8bit或者做了序列打包,或者数据本身平均长度很短,别太迷信教程里的数字。至于代码模型和对话模型的超参区别,代码任务一般学习率可以稍微调高一点,比如2e-4到5e-4,同时warmup steps要短,因为代码数据分布相对稳定,不像对话那样多样;但最重要的还是要把loss关注在completion部分而不是prompt上,不然模型会学着去预测你自己的指令模板。还有个小坑,代码数据里经常有大量重复的token比如缩进和括号,建议用mask掉padding和prompt的loss,不然模型容易只顾着学格式。你先试试序列缩短加gradient checkpointing,八成能跑起来,再不行就把batch size降到2,但别急着改学习率,稳定训练比追求大batch重要得多。
说实话你这个配置OOM挺奇怪的,A100 80G跑4bit的8B模型理论上batch size 8甚至12都该没问题。我怀疑问题不在batch size本身,而是你sequence length拉到2048之后,attention的显存占用是平方增长的,加上QLoRA虽然量化了权重,但激活值还是全精度,这块开销很容易被忽略。gradient checkpointing确实建议开,它能用一点计算换大量显存,开了之后batch size 8基本稳的,16的话看情况。另外你用的transformers+peft,记得把gradient_accumulation_steps配合起来用,别硬顶batch size。至于代码模型和对话模型,我觉得最大的区别是学习率和warmup——代码补全任务通常更吃低学习率(1e-4到2e-4左右),而且序列长度本身就长,反而要注意截断策略,别让pad token占太多计算。我自己试过在CodeLlama上做类似任务,感觉数据质量比超参重要多了,几万条如果重复度高,反而容易过拟合到某些模式上。你那个数据集的代码风格统一吗?如果不统一,可能先做一下去重和格式化,比调batch size收益更大。
80G跑4的batch还爆,大概率不是显存容量问题,是峰值显存被激活值吃满了。gradient checkpointing建议直接开,序列2048的话这步基本是必须的,显存能省一半以上。另外你可以试试把优化器状态offload到CPU,或者用paged_adamw_8bit,这几个组合下来16batch应该没问题。
代码模型和对话模型超参差异主要在训练目标和数据分布上,代码补全更吃长序列和更低的learning rate,我一般用1e-4到2e-4,对话模型反而可以稍微激进点。你几万条数据的话,跑2-3个epoch就够了,别死磕大batch,效果真不一定差。
说实话你这个配置OOM挺正常的,A100 80G看着大,但QLoRA跑8B模型序列长度2048时,激活值才是真正的显存杀手,batch size只是压垮骆驼的最后一根稻草。gradient checkpointing必须开,这个能省掉一大半激活显存,代价就是训练慢个20%-30%左右,但总比OOM强。另外你还可以试试把梯度累积步数加上去,batch size降到1或者2,照样能模拟出大batch的效果,别迷信教程里那种16、32的数字,人家可能用的是DeepSpeed ZeRO-3或者把序列长度砍到1024了。
至于代码模型和对话模型的超参区别,我感觉主要在学习率和warmup步数上。代码补全任务对局部上下文更敏感,学习率通常要比对话模型低一点,比如2e-4到4e-4之间,warmup可以稍微拉长,让模型先稳住底层表示。另外你数据集几万条不算多,LoRA的rank可以设小一点,8到16就够,不然容易过拟合,尤其是代码这种结构化文本,正则化稍微强一点没坏处。
我还有个疑问,你用的是transformers自带的trainer还是自己写的训练循环?如果是后者,注意一下是不是把labels也放到显存里了,有时候这个细节会莫名其妙多占几个G。你试试把序列长度临时改成512跑一下,如果显存占用断崖式下降,那基本就是激活值的问题,不是batch size本身。
80G跑4的batch还爆?先把gradient checkpointing开开,这玩意儿能省一半多显存。
8B模型4bit量化后本身权重就占5-6G,但LoRA的梯度、优化器状态和激活值才是大头,你序列拉到2048,单条样本的激活内存是随长度线性涨的,batch4爆掉不奇怪。gradient checkpointing又不是什么玄学,它就是把前向激活扔了反向再算一遍,显存能省一半以上,代价是慢30%左右,A100跑这个规模完全能接受,开了之后batch8甚至12应该没问题。至于教程里说16、32的,人家多半是开了梯度累积加checkpointing,或者序列长度只有512,你拿2048去比肯定不对。代码模型和对话模型超参差挺多的,代码补全更吃长序列和低学习率,因为要捕捉跨行的结构依赖,LoRA rank可以试16到32,对话模型反而常用8,另外代码任务里warmup步数要拉长,loss经常会有突变,建议观察几个step的loss曲线再停。你试试把序列裁到1024,开checkpointing,batch4加梯度累积到等效16,如果还爆那就是数据预处理时padding没做对,看看有没有把eos token硬塞进标签导致计算异常。
A100 80G跑bs=4还OOM确实有点怪,我4090 24G用QLoRA跑7B、seq 1024都能上bs=8。先确认下是不是gradient checkpointing没开,这玩意对显存影响巨大,开了基本能省一半。另外你optimizer用的啥?如果没换成paged_adamw_8bit,优化器状态也会吃掉不少显存。代码补全和对话微调区别挺大的,代码任务一般lr要低一点,序列里有效token占比高,梯度累积步数也得跟着调,不然容易过拟合。
序列长度2048挺吃显存的,试试gradient checkpointing加梯度累积,batch小点也能跑。
A100 80G跑4bit QLoRA还OOM,batch size 4确实有点不对劲。你序列长度2048、数据几万条,激活值占用不小,gradient checkpointing基本是必备的,不开的话显存翻倍很正常。另外确认下有没有把optimizer设成paged_adamw_8bit,普通adam状态也很吃显存。代码补全和对话微调差别挺大,前者通常序列更长、更依赖完整上下文,学习率可以低一点,epoch也别太多,不然容易过拟合。
2048序列太长了,显存大头在这,试试降到512再开gradient checkpointing。代码补全和对话微调差别不大,主要看数据格式。
batch size=4在2048序列长度下爆80G显存有点不太正常,建议先查一下是不是把optimizer的state也算进去了,或者试试梯度累积来绕过显存瓶颈。gradient checkpointing基本是QLoRA标配了,开了能省不少显存,代价就是慢一点,但A100上完全扛得住。代码补全和对话模型超参差别其实不大,主要看数据格式和loss mask怎么设,代码任务一般序列更长但有效token更密集。另外检查下peft版本,老版本确实有显存泄漏的坑。