最近在试着用LoRA微调一个7B的模型做代码生成,显卡是4090 24G。我参考了一些开源项目的写法,设置lora_r=8,lora_alpha=16,batch_size=1,gradient_accumulation_steps=8,max_seq_len=2048。结果跑起来显存直接飙到23G+,偶尔还会OOM。我看很多教程说LoRA应该很省显存,但我这个好像没比全参数微调省多少。已经试过把seq_len降到1024,也只降了2G左右。是不是我的target_modules选太多了?还是说7B模型本身就这样,24G就是勉强?有没有实际跑过类似配置的朋友能分享下你们的显存占用和参数设置?
LoRA微调7B模型显存一直爆,是batch size问题还是我代码写错了?
全部回复
共 111 条说实话你这个问题我太有同感了,之前用3090跑7B的时候也踩过一模一样的坑。LoRA省显存是相对全参数微调而言的,但7B模型光加载权重就要14G左右,加上梯度和优化器状态,24G其实挺吃紧的,你这个占用数字我看着就觉得正常。target_modules确实会影响显存,但你这配置lora_r才8,影响有限,更关键的是你的max_seq_len=2048,注意力计算和KV cache才是大头,降到1024只省2G也印证了这点。我建议你开一下gradient_checkpointing,能省不少,代价就是慢个20%左右,另外检查下是不是把模型的全部参数都设成requires_grad了,LoRA应该只训练低秩矩阵。我之前跑7B用batch_size=1,seq_len=1024,gradient_checkpointing开着,显存能压在16-18G,你可以试试这个组合。还有个小技巧,把optimizer换成AdamW的8-bit版本,能再抠出1-2G。最后如果你只是想调通代码,可以先拿2B模型跑通流程,再换回7B,省得一直被OOM打断调试思路。
24G跑7B LoRA确实紧,但你这占用不太正常。我之前用4090跑7B,seq_len 2048,batch 1,lora_r=16,显存大概16G左右,你这23G明显偏高。建议查下是不是把embedding和lm_head也加进target_modules了,这俩参数量巨大,LoRA微调一般不加。另外看看有没有开gradient_checkpointing,这个能省不少。
这配置看着挺正常的,问题大概率出在max_seq_len和梯度检查点上。7B模型即使LoRA,激活值也占大头,2048长度对24G卡确实吃紧,我跑13B时开gradient_checkpointing后显存能降30%左右,你可以试试。另外target_modules别全选,只挑q_proj和v_proj,能再省一点。还有,看看是不是用了flash_attention,没开的话显存差距很大。
这配置正常,7B+2048长度24G就是极限,val数据记得关gradient checkpointing试试。
24G跑7B+LoRA,seq_len 2048确实会到临界点,你这占用其实算正常,不算代码问题。我试过lora_r=16、seq_len 1024,batch=1,峰值大概18G左右,但一旦开gradient checkpointing能再省3-4G,你试试这个开关。另外target_modules别贪多,我通常只挑q_proj和v_proj,全选的话激活内存会明显涨。7B这体量,想舒服点还是得上48G的卡,24G属于能用但得抠配置。
24G跑7B LoRA这个占用其实挺正常的,我拿3090试过类似配置,seq_len 2048时基本也是贴着上限走,你降到1024能省2G已经不错了。OOM大概率不是target_modules的问题,而是attention的KV cache在长序列下太吃显存,你可以试试开gradient checkpointing,能省不少。另外检查下是不是把adapter也加载进显存计算了,有些库默认会保留全量权重做forward。我自己的经验是7B想舒服点微调,24G还是得上8bit量化,不然只能接受小batch加短序列。
我之前跑7B也遇到过类似情况,24G看着大但LoRA的显存大头其实在激活值和优化器状态上,seq_len砍半只降2G挺正常的。你试试gradient_checkpointing有没有开,这个能省不少,再配合8bit优化器,基本能把峰值压到15G左右。另外target_modules如果像q,k,v,o,gate,up,down全选了,计算图也会变大,可以先只选q和v试试。
2048的seq_len在7B上确实挺吃显存的,LoRA省的是优化器状态和梯度那部分,但激活值该占还是占。你试试开gradient_checkpointing,再把flash attention打开,这两个能省不少。另外target_modules如果qkv全加了,参数量上去激活也会涨,可以先只挂q_proj和v_proj看看。我4090上跑7B LoRA,seq_len 1024加checkpointing大概16-18G浮动,2048的话23G确实悬。
2048的seq_len在7B模型上确实挺吃显存的,尤其是你开了gradient checkpointing没有?如果没开的话,activation那部分能吃掉不少。我之前用4090跑Qwen-7B的LoRA,r=16,seq_len=1024,batch=1,grad_accum=4,显存大概在18-20G左右浮动,开checkpointing能省个2-3G。你23G+的话,先确认下是不是把embedding和lm_head也放进target_modules了,有些人图省事直接all-linear,那显存直接起飞。另外优化器如果用adamw,7B的LoRA参数虽然不多,但optimizer state还是会占一点,换成adafactor或者8bit adam能再挤出来一些。还有个容易忽略的点是dataloader的num_workers和pin_memory,这些也会吃显存,虽然不多但积少成多。建议你先用nvidia-smi或者torch.cuda.memory_summary看下到底是哪块占大头,别光靠猜。
2048的seq太吃显存了,试试flash attention加梯度检查点,24G跑7B LoRA本来就紧。
你这显存占用不太正常,2048序列长度加batch 1不该这么夸张。我跑7B LoRA的时候4090上大概16-18G,你重点看下是不是把embedding和lm_head也加进target_modules了,或者用了fp32加载模型。另外检查下gradient_checkpointing有没有开,这个对显存影响很大,但会稍微慢一点。