最近在试着用LoRA微调Qwen2.5-7B,想让它懂一点我们公司内部的技术文档(大概5000条QA对)。机器是4090 24G,一开始用的QLoRA + 4bit,batch_size=1,梯度累积设了8,结果跑了不到500步显存就满了,直接OOM。我看网上教程说4bit + LoRA应该很省显存啊,是我哪里设置不对吗?另外,我用的transformers + peft,是不是gradient_checkpointing没开导致的?还是说7B模型本身做指令微调就需要更高的显存底线?有没有大佬分享下自己微调7B的显存配置,或者推荐更小一点的模型(比如3B/4B)来试水?感谢!
用LoRA微调Qwen2.5-7B做领域问答,显存只够跑4bit但还是爆了,正常吗?
全部回复
共 8 条gradient_checkpointing大概率就是罪魁祸首,我跑13B的时候没开它也是几百步就爆,开了之后显存直接砍半。不过4bit+LoRA在24G上跑7B指令微调确实有点极限,尤其你的序列长度如果超过1024的话。建议先开gradient_checkpointing再加个8x的梯度累积试试,还不行就换3B吧,效果其实差不了太多,至少能让你把batch提上去。
4090 24G跑7B QLoRA按理说不会500步就OOM,你这个情况八成是gradient_checkpointing没开,再加上序列长度可能设太长,context一旦超过2k,激活值直接起飞。我之前用同样配置跑Llama-3-8B,batch_size=1,max_seq_len=2048,开gradient_checkpointing,显存峰值大概在18G左右,能稳跑,但你要是把seq_len拉到4096,那24G肯定扛不住。另外注意下你5000条QA对的数据长度分布,如果很多长文档截断后还是超长,建议先做下长度分析,把超过2048的过滤掉或者拆成多轮。还有个小坑,peft的target_modules别全选,只挑q_proj和v_proj能省不少显存,虽然效果可能略降,但前期调试够用。如果你实在想省心,换Qwen2.5-3B或者4B先跑通流程,7B的显存优化后面再慢慢调,毕竟数据量和训练目标才是关键,模型大小只是手段。
gradient_checkpointing必须开,24G跑7B QLoRA稳稳的,你试试把序列长度也砍到512。
gradient_checkpointing确实得开,这玩意儿不开的话24G跑7B的4bit LoRA基本就是极限操作,我上次开满序列长度2048也差点爆。另外你把梯度累积砍到4或者2试试,显存压力会小很多,5000条数据其实不用那么大的累积步数。实在不行就换Qwen2.5-3B吧,效果差距没想象中大,至少能稳定跑完。
gradient_checkpointing没开的话,24G跑7B的4bit LoRA确实容易爆,这玩意儿基本是必选项,开了能省一半左右。另外建议你把序列长度砍到512或768,5000条QA对其实不需要太长上下文,还有attention的显存占用是平方增长的。我同样配置跑过类似的活儿,batch_size=1+梯度累积8,峰值能压在18G以内,你再检查下是不是加载了完整的tokenizer或额外缓存。要是还不行,换Qwen2.5-3B试水绝对够用,效果差距没你想的那么大,毕竟领域数据才是关键。
4090 24G跑7B QLoRA按理说是够的,但你这个OOM大概率就是gradient checkpointing没开,这玩意儿在peft里默认是关的,开了之后激活值内存能砍掉一大截,尤其你序列长度如果超过1024,差距会非常明显。另外你说5000条QA对,步数才500就爆,那可能不是峰值显存问题,而是累积了太多梯度状态,试试把gradient_accumulation_steps降到4,然后配合gradient_checkpointing一起用,应该能稳住。还有个坑是transformers加载4bit时如果没设low_cpu_mem_usage=True,有时候会额外吃显存,虽然不太常见但值得排查一下。至于模型选择,如果你只是做领域问答,其实Qwen2.5-3B配合LoRA效果也不会差太多,特别是你的数据量才5000条,7B反而容易过拟合,3B或者4B训练起来更从容,推理也快。我之前用7B跑过类似任务,batch_size=2,开gradient_checkpointing,峰值大概16G左右,你参考下这个余量。最后建议你先把torch.cuda.max_memory_allocated()打出来看看峰值到底出现在哪个环节,是前向还是反向,别瞎调。
4090 24G跑7B QLoRA按理说挺稳的,你这个问题大概率出在没开gradient_checkpointing上,这玩意儿能省一大半激活显存,不开的话batch_size=1也可能爆。另外你检查下sequence length是不是设太长了,5000条QA里要是有些长文档,padding到2048甚至更长,显存直接吃满很正常。我自己的经验是7B 4bit + LoRA,开gradient checkpointing,seq_len控制在1024,24G能跑batch_size=4,梯度累积反而可以降下来,省时间。还有个小坑,peft的target_modules记得只选q_proj和v_proj,全选所有线性层会让可训练参数变多,显存压力也大。如果实在调不动,可以先拿Qwen2.5-3B跑通流程,效果差不了太多,毕竟你们内部文档领域性强,数据量5000条对7B来说也不算大,3B微调后可能更不容易过拟合。最后建议你盯着nvidia-smi看下是不是被别的进程占了显存,我遇到过好几次这种灵异事件。
4090 24G跑7B的QLoRA按理说应该够,但你500步就OOM大概率是gradient_checkpointing没开,这玩意儿不开的话激活值能吃掉好几个G,加上你梯度累积8其实本质是把batch撑大,显存峰值反而更高。我自己的经验是7B 4bit微调,开gradient_checkpointing + 8bit优化器状态,batch=1,序列长度控制在1024以内,大概峰值在15-18G,你可以先试试这个组合。另外你5000条QA对其实数据量不算小,7B全量微调确实需要更多余量,如果你把seq len压到512,或者用deepspeed zero2把优化器状态offload到CPU,24G应该能稳。实在不行就换Qwen2.5-3B,效果对于内部文档问答差距没那么大,而且你能把batch提到4-8,收敛速度反而可能更快。你检查下是不是序列长度太长,或者attention实现没走flash attention?把这两项搞定再试一次。