最近在试着微调一个7B的chat模型,用的LoRA,单卡A100 40G。训练的时候loss降得挺正常,但一跑推理就报CUDA OOM,而且加载模型的时候就要吃掉快25G显存。我试了FP16和int8量化,感觉也没好多少。看别人的部署教程好像很轻松,是不是我的batch size或者max length设置有问题?还是说7B模型本来就这德行,得用vLLM或者什么其他框架才行?求有经验的大佬指点一下,给个大致的方向就行,谢谢了!
部署7B大模型微调后显存爆了,求大佬看看是不是我配置有问题?
全部回复
共 49 条说实话你这个情况我太熟了,7B模型加载权重本身就差不多要14G(FP16),加上KV cache和中间激活,25G起步真不算离谱。LoRA微调时显存占用高是因为反向传播要存梯度,但你推理时还这么吃显存,大概率是max length或者batch size设太大了,你可以试试把max length砍到512,batch size设成1,应该能压到20G以内。另外你说的int8量化没效果,我猜你是用transformers的load_in_8bit,那个其实只在加载时省了权重内存,但推理时的KV cache还是按FP16算的,所以显存降得有限。vLLM确实是个方向,它主要优化了KV cache的显存管理,7B模型在40G上跑并发应该挺轻松的,不过我建议你先用transformers的torch.compile配合静态缓存试试,有时候不用换框架就能解决。还有个小坑,你检查下是不是把微调时的padding策略带到推理了,如果左padding或者固定长度padding,显存会白白浪费一大截。最后想确认下,你加载模型时是不是还顺便加载了adapter权重?如果LoRA的target modules设得太多,合并后的模型会比原版大不少,这可能也是显存异常的一个原因。
说实话25G的加载占用不太正常,7B的FP16权重也就14G左右,你是不是max length设太长了导致KV cache爆炸?LoRA微调后合并权重再推理会好很多,另外试试把tokenizer的padding和truncation统一一下,有时候隐性padding会白白吃掉几个G。vLLM是能省不少显存,但你这情况更像是推理配置的问题,先把max length砍到2048看看,别急着上框架。
说实话25G加载7B模型确实有点高了,我怀疑你transformers加载时把梯度 checkpoint 或者缓存没清干净,试试model.eval()加torch.no_grad(),然后确认下是不是把训练时的lora权重也一起load进去了。另外A100 40G跑7B推理理论用不到一半,vLLM能省不少但本质还是得看你的max length和batch,8k以上长度就别指望裸跑舒服了。你可以先跑个最简单的单条短文本看看显存占用,如果还高就查查代码是不是有什么buffer没释放,我上次就是忘了删optimizer导致推理时占用翻倍。
推理OOM基本是KV cache的锅,跟微调配置关系不大,换vLLM能省不少显存。
加载就吃25G挺正常的,FP16的7B光权重就14G了,加上KV cache和中间激活,推理时OOM不奇怪。你LoRA微调完有没有把adapter merge回去?如果没merge直接挂adapter跑,显存占用反而更高。建议试试vLLM或者llama.cpp这类推理框架,PagedAttention对KV cache管理好很多,batch size和max length也能动态调。另外int8量化要生效得用bitsandbytes的load_in_8bit,光转权重格式没用。
推理OOM跟训练配置关系不大,先试试vLLM部署,加载25G正常,但推理爆显存大概率是KV cache没管好。
推理OOM跟训练本身关系不大,主要是加载和KV cache吃显存。你试试vLLM或者TGI这类框架,它们对显存管理比transformers默认的generate好太多,PagedAttention能省不少。另外max length别设太大,推理时KV cache是随batch和长度线性涨的。int8量化如果只量化权重、激活还是fp16,省得有限,可以看看GPTQ或AWQ。
推理OOM跟训练两码事,试试vLLM或TGI部署,加载25G挺正常的,7B fp16权重就14G了。
推理OOM跟训练配置没关系,是加载方式的问题,换vLLM或者用device_map分片加载会好很多。