最近想用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设太大了吗?
全部回复
共 9 条我也在折腾类似的事情,看到这个帖子简直像在照镜子哈哈。我用的是7B模型,A100 80G,batch size设2就爆了,一度怀疑自己买到了假卡。后来翻了一圈才发现,问题可能不在batch size本身,而是序列长度和attention的计算方式。2048的序列长度对于8B模型来说,即使量化了,显存消耗也很大,尤其是LoRA的adapter参数虽然小,但前向传播时中间激活值还是占大头。gradient checkpointing必须开,这玩意儿能省将近一半的显存,代价就是慢一点,但总比OOM强。
另外有个细节容易被忽略:transformers的pad_token_id和attention_mask设置不对的话,可能会导致padding部分也参与计算,白白浪费显存。建议把padding_side设为right,truncation设为True,再用DataCollatorForSeq2Seq或者自己写个collator确保每个batch的序列长度对齐。
至于代码模型和对话模型的区别,我感觉主要是学习率和warmup策略。代码补全更依赖长程依赖,所以序列长度最好不要减,但可以试试把lora_target_modules从默认的q_proj,v_proj扩展到k_proj,o_proj,gate_proj,效果会好一些。对话模型则更关注生成多样性,LoRA rank可以设小一点,比如8或16,而代码模型可能32更稳。
对了,你用的peft版本是多少?老版本有个bug会导致量化后的模型在forward时重复分配显存,升到0.11.0以上试试。还有,A100的显存是80G但实际可用只有79G左右,系统也会占一点,所以别卡太死。
A100 80G跑4bit量化还OOM,大概率是序列长度2048加peft的gradient checkpointing没开导致的,开一下应该能压到batch size 8甚至16。代码模型和对话模型区别挺大的,代码任务建议学习率调低点(比如2e-4左右),而且数据集如果包含长上下文的话,sequence length可以适当砍到1024试试。另外注意一下数据预处理时padding策略,别让无效token浪费显存。
A100 80G跑4bit的QLoRA,batch size设4就OOM,这确实不太正常。我猜问题大概率出在gradient checkpointing没开上,这玩意儿对显存节省特别明显,尤其是序列长度2048的时候,开了之后batch size翻倍都不是问题。另外你确认一下是不是把gradient accumulation steps和batch size搞混了,有些教程说的16或32可能是accumulation之后的等效batch,实际单次forward的batch size可能也就2或4。代码模型和对话模型在超参上差别其实挺大的,代码补全更吃长序列依赖,学习率通常要调低一点,我一般用1e-4起步,然后看loss下降速度再降。还有你用的数据集如果每条都很长,建议检查一下padding策略,有些情况下dynamic padding能省不少显存。最后问一句,你用的peft版本是最新的吗?老版本对LLaMA-3的4bit支持有bug,我之前升级之后OOM问题直接解决了。
老实说A100 80G跑4bit的LLaMA-3-8B,batch size设4就炸了确实不太正常,我怀疑问题可能不在batch size本身,而是序列长度2048加上attention的缓存占了不少显存,尤其代码补全任务里输入输出经常会被padding到很长。你可以先试试把gradient checkpointing开起来,这玩意儿对显存节省特别明显,代价就是慢一点,但总比跑不起来强。另外检查一下是不是数据集里文本长度分布不均匀,有些样本实际填充到了max length,导致内存碎片化严重,可以考虑动态padding或者用DataCollatorWithPadding按batch内最大长度来截断。至于代码模型和对话模型的差异,我感觉代码任务对长程依赖和结构化信息更敏感,所以learning rate通常要比对话模型小一点(比如1e-4甚至5e-5),LoRA的rank也可以选大一些(16或32),不然代码里的模式可能学不到位。还有个小建议,试试用bitsandbytes的4bit优化器或者用torch.compile,有时候编译后能省不少显存。
A100 80G跑4bit QLoRA还爆显存,batch size设4确实不大,问题大概率出在序列长度2048上,代码补全任务上下文长,显存占用会指数级涨。gradient checkpointing建议开一下,能省不少显存,虽然慢点但至少不爆。另外那些跑16或32 batch的教程,多半是用了更短的序列长度或者梯度累积,你可以试试把gradient accumulation steps设大点来模拟大batch效果。代码模型和对话模型的区别主要在数据格式和learning rate上,代码任务建议lr设低一点,1e-4左右起步,对话模型通常可以高一些。
A100 80G跑4bit量化,batch size 4就爆显存确实不太正常,大概率是序列长度2048加上数据padding浪费了太多显存。建议先开gradient checkpointing,这玩意儿能省一半左右,另外可以看看是不是tokenizer把代码里的空格/换行算成了很多个token,导致实际序列比2048还长。代码模型和对话模型超参区别挺大的,代码任务学习率通常要小一点,比如1e-4左右,而且warmup步数可以设短些,毕竟代码数据分布比对话更结构化。
开梯度检查点吧,8B模型4的batch还不开这个肯定炸,代码模型学习率可以比对话模型低一点。
batch size 4就爆显存确实不太正常,建议先开gradient checkpointing试试,能省不少显存。
同款问题,我之前微调也是batch size一上来就爆,后来发现单靠QLoRA省显存不够,gradient checkpointing几乎是必开的,能省一半左右。而且你序列长度2048配合4的batch,显存占用本来就不低,试试batch size降到1或2,同时开梯度累积,效果差不多。代码模型和对话模型差别还挺大的,代码任务一般学习率可以稍微低一点,比如2e-4起步,另外数据集里最好多加点代码结构相关的指令对,不然容易学歪。