最近想用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还爆基本不是batch size的锅,你seq len 2048加4bit量化后,activation才是大头。gradient checkpointing必须开,能省一半多,然后把gradient accumulation用起来,batch size设1都行,效果等效。代码模型和对话模型超参差挺多的,代码补全建议lr调低点(1e-4到2e-4),warmup steps拉长,因为代码分布更陡,收敛节奏不一样。你试试开checkpointing后batch 2-4,应该稳了。
A100 80G跑4bit的8B模型,batch size 4就爆基本可以确定是没开gradient checkpointing,这个必须开,能省一半还多的显存,开了之后batch size 8甚至16应该没问题。另外你序列长度2048对代码补全来说偏长了,很多代码tokenizer压缩率不高,实际信息密度低,可以试试1024,显存压力会小很多。代码模型微调的话,学习率一般要比对话模型低一点,我习惯用1e-4到2e-4之间,warmup比例也可以调小些,因为代码数据分布更稳定,不需要那么多预热。还有个坑是max_length别设成2048就真的每条都填满,最好看看你数据集里真实长度分布,很多教程里的batch size都是基于短序列算的。
80G跑4的batch还OOM肯定不对劲,我怀疑你序列长度2048加上4bit量化后,KV cache和中间激活值才是真正吃显存的大头。你可以试试开gradient checkpointing,能省一半以上,然后batch size调到8或16没问题。至于代码模型和对话模型的超参,代码任务lr可以稍微调低点,比如1e-4到2e-4,warmup步数拉长,因为代码分布更陡峭,训太猛容易loss震荡。另外你确认下是不是paged_optimizer没开,QLoRA那个参数很关键,不开的话内存碎片化也会导致莫名爆显存。
你这个问题我上周刚踩过坑,batch size不是主因,主要是attention的中间张量在2048长度下太大了。建议你先把gradient checkpointing打开,再把batch size降到2,用gradient accumulation凑到等效batch size,显存占用能直接砍掉一半多。代码微调的话,我一般lr设5e-5到1e-4,对话模型可以稍微激进点,但代码补全任务对过拟合更敏感,warmup ratio建议调到0.1以上。顺便问下,你用的peft是不是最新版?老版本对LLaMA-3的LoRA target modules支持有bug,也会导致显存异常。
我倒是觉得你sequence length
说实话80G跑4bit的8B模型还爆显存,大概率不是batch size的锅,而是你的序列长度和attention计算占了太多激活内存。2048的序列长度对8B模型来说算挺长的了,加上QLoRA虽然量化了权重,但前向传播时的中间激活值还是全精度,这部分才是真正的显存大头。
gradient checkpointing确实得开,不开的话激活值会一直存到反向传播,开了之后能省掉大概60%到70%的显存占用,代价就是训练速度慢个20%左右,但总比OOM强。另外你可以试试把序列长度先砍到1024跑通流程,再慢慢往上加,毕竟代码补全很多场景下上下文也不需要那么长。
至于batch size,QLoRA本身对batch size不敏感,4和8在效果上差别不会很大,关键是梯度累积步数要跟上,比如batch size=4加上8步累积,等效batch size就是32,很多教程说的16或32应该是指这个等效值,不是单卡物理batch。
微调代码模型和对话模型最大的区别在于学习率和warmup策略。代码任务通常需要更低的峰值学习率,比如1e-4到2e-4之间,而且warmup步数可以稍微长一点,因为代码数据的分布更尖锐,学太快容易震荡。另外代码模型对loss的收敛判断跟对话不一样,建议多看看验证集上的pass@k指标,别光盯训练loss。
还有个小坑,transformers+peft在加载QLoRA模型时,如果忘了设置trust_remote_code或者tokenizer的padding side,可能会导致隐式增加序列长度,这个也会爆显存,你可以检查下数据collator是不是把padding加到了最长样本。先试试把gradient checkpointing开起来,然后把max_length降到1024,batch size保持4,应该就能跑起来了。
A100 80G跑4bit的8B模型,batch size 4就爆肯定是没开gradient checkpointing,这玩意儿对LoRA来说基本是必开的,能省一半多显存。另外你序列长度2048对代码补全来说可能偏长了,如果数据里大部分样本没那么长,可以试试动态padding或者截断到1024。代码模型和对话模型超参确实有区别,代码任务学习率一般可以稍微调大点,warmup步子也短一些,因为代码分布更结构化,收敛快。你试试开checkpointing然后把batch提到8,应该就稳了。
开梯度检查点吧,显存能省一半多,batch4够用了,别人16那是短序列。
代码模型学习率得调低点,1e-4试试,对话模型那套参数直接搬过来容易炸。
80G的卡跑4bit的8B模型,batch size 4就爆基本可以确定不是batch的问题,因为光模型权重加激活值远没到极限。你sequence length拉到2048,这个对显存的消耗比batch size更致命,尤其代码数据里长序列占比高的话,attention的显存是平方级增长的。gradient checkpointing肯定要开,它能在几乎不影响效果的情况下把激活显存砍掉大半,代价只是慢一点。另外你用的transformers+peft,记得把unsloth或者flash-attention装上,这俩对长序列的显存优化非常明显,有时候比调batch size管用得多。至于代码模型和对话模型的微调区别,我感觉代码模型对学习率更敏感,稍微大一点就容易把已有能力冲掉,建议从1e-4往下调,而且代码数据里不同语言和框架的分布差异大,最好按文件粒度采样而不是按行数,不然容易让模型偏好某些常见模式。我自己的经验是,代码补全任务用LoRA的话,rank设16到32就够,再高收益很小但显存和过拟合风险都上去了。你试试开gradient checkpointing加上flash-attention,batch size提到8到12应该没问题,如果还爆就查一下是不是dataloader里把整个数据集都load进显存了,有些人会无意中干这事。
80G还爆的话肯定不是batch size的锅,你序列长度2048加上4bit量化后激活值才是大头,开gradient checkpointing能省一半以上显存。代码模型和对话模型超参差别挺大的,代码补全建议学习率调低点(1e-4左右),warmup步数拉长,另外别用对话模板直接喂原始代码就行。你试试把batch降到2加梯度累积到8,效果应该差不多。
说实话能跑4已经很不错了,我同样的配置batch size 2都经常爆,gradient checkpointing肯定得开,能省不少显存,另外可以试试把序列长度砍到1024,代码补全不一定非要2048。至于batch size 16或32那些教程,多半是用了deepspeed或者多卡,单卡别太当真。代码模型微调的话,学习率可以稍微调低一点,1e-4到2e-4之间比较稳,对话模型倒是经常用5e-5,感觉代码任务对参数扰动更敏感。
你用的是transformers的trainer还是自己写循环?有时候DataLoader那边num_workers设太高也会吃显存,特别是几万条数据的情况下。
显存爆八成是seq_len太长,2048吃显存很凶,开gradient checkpointing能省一半,batch先降到2试试。
80G都爆?开gradient checkpointing吧,能省一半多,batch能翻倍。
代码补全跟对话微调差别挺大,学习率调低点,序列长点。
80G跑4的batch就OOM大概率不是batch size的锅,你序列长度2048加上4bit量化后激活值才是大头,gradient checkpointing必须开,开了之后batch 8甚至16应该没压力。代码补全和对话微调差别挺大的,代码任务学习率可以稍微调高一点(比如2e-4),而且warmup步数不用太长,因为代码分布更稳定。你试试packing把短序列拼起来,能大幅提升吞吐,不然几万条数据很多是padding浪费算力。
试试开gradient checkpointing,然后batch降到2,梯度累积补回来,A100跑8B绰绰有余。
代码模型学习率一般比对话模型低,我习惯用2e-4,你参考下。
开梯度检查点吧,显存直接掉一半,另外8B模型2048长度batch4已经很极限了,别信那些截图。
说实话80G跑4bit的8B模型batch size=4还OOM,我觉得大概率不是显存容量的问题,而是transformers默认把激活值全存下来了。gradient checkpointing几乎是必须开的,开了之后显存占用能掉一半以上,代价就是训练慢个20%左右,但换来的batch size提升绝对值很划算。你看到的那些跑16甚至32的教程,多半是开了checkpointing外加用了paged optimizer,QLoRA里那个4bit的double quantization本身也吃显存,别小看。另外序列长度2048对代码补全来说有点奢侈,如果数据集里大部分样本没那么长,可以试试把max length砍到1024甚至512,显存压力会小很多,而且收敛速度反而可能更快。至于代码模型和对话模型的超参区别,代码任务通常学习率要稍微低一点,比如2e-4到5e-4之间,warmup步数也不用太长,因为代码分布更结构化,loss下降比较快。我自己的经验是,先拿500条数据小跑一版,盯着显存曲线调,比直接上全量省心多了。你用的peft版本是0.11以上吗?老版本对4bit的适配有时候会有bug导致显存泄漏,升级一下可能就解决了。
80G跑4的batch还爆,确实不对劲,但我赌五毛你没开gradient checkpointing。这玩意儿对LoRA来说基本是必选项,开了之后显存能掉一半以上,4batch稳得一批。至于教程里吹的16、32,多半是序列长度短或者用了DeepSpeed ZeRO,别太当真。代码模型和对话模型超参差距挺大的,代码补全学习率可以稍微调高一点(比如2e-4),但warmup steps要拉长,因为代码分布比对话更陡。另外你几万条数据其实不算多,建议先拿1万条跑个epoch看下loss曲线,别一上来就全量。
8B上4bit还爆显存?先开gradient checkpointing,再调小batch,这两个不冲突。
A100 80G跑4bit的8B模型,batch size 4就爆大概率不是显存容量问题,而是序列长度2048带来的激活值暴涨。gradient checkpointing确实得开,能省一大半激活显存,另外你试试把gradient accumulation加上,batch size降到1或2,效果一样但显存压力小很多。代码模型和对话模型超参差异挺大的,代码补全任务学习率可以稍微调高一点,warmup steps也要适当增加,但主要还是得看你的数据分布和loss曲线,建议先跑个小样本实验看看loss有没有明显下降再说。
80G跑4的batch还爆,八成是序列长度和梯度检查点的事,开个gradient checkpointing能省一大截显存。
代码模型lr可以调低点试试,对话模型那套超参直接搬过来容易飘。
开梯度检查点吧,8batch没问题,代码模型lr调低点试试。