最近在调一个BERT分类模型,训练完保存了state_dict,然后单独写了个推理脚本加载。奇怪的是,训练时显存占用大概4G左右,但一跑推理,显存直接飙到接近10G,然后OOM。我明明已经用了torch.no_grad(),也调了model.eval(),还试了torch.cuda.empty_cache(),都没什么效果。输入数据也就是一条文本,batch_size=1。是不是我哪里写得不规范?还是说PyTorch推理本来就这么吃显存?求有经验的朋友指点一下排查思路,或者有没有什么通用优化技巧,谢谢!
PyTorch模型推理时显存暴涨,是哪里写错了?求大佬指点
全部回复
共 108 条试试关掉梯度之后把优化器也删了,推理脚本里别加载optimizer状态,多半是这玩意儿占的。
试试把优化器状态也存下来,有时是优化器占着显存没释放。还有检查下有没有不小心把梯度传进模型了。
讲真,你这个情况我遇到过好几次,训练4G推理10G基本可以排除是模型本身算力需求,问题大概率出在推理脚本的写法上。最可疑的是你加载state_dict之后有没有把模型也挪到gpu上,但更常见的是你没关掉optimizer或者没把一些训练时才需要的buffer和梯度缓存清干净。我猜你是不是用了类似model.zero_grad()或者保留了某些hook?另外可以查一下是不是在forward里用了类似torch.no_grad()包着但内部有创建计算图的op,比如某些自定义的attention mask或者position encoding用了原地修改导致autograd记录。一个很实用的排查办法是,在推理循环里每跑一次就打印一下torch.cuda.memory_summary(),能看到是哪个tensor占了大头,我之前发现是tokenizer的padding到max_len的tensor忘了detach,结果整个中间变量被保留。还有个小坑,如果你调了model.eval()但没把dropout的training flag传下去,某些层可能还在走训练路径。实在不行就试一下torch.jit.script或者把输入先移回cpu再forward,看显存是否还涨,这样能定位是不是gpu端缓存问题。优化方面,你可以试试开启torch.backends.cudnn.benchmark=True,然后把输入序列长度固定,避免动态shape导致缓存碎片化。
我之前也踩过类似的坑,训练4G推理10G这个量级差得太离谱了,肯定不是正常的。你试试把推理脚本里的输入也放到GPU上,然后确认一下是不是在no_grad块里做的forward,有时候加载完模型忘了model.cuda(),输入在CPU上反而会触发一些隐式的数据传输和缓存。另外,BERT的attention机制在推理时如果开启了grad,哪怕只是默认的torch.is_grad_enabled()为True,它也会为中间激活值分配内存,光model.eval()并不会关掉autograd,必须确保torch.no_grad()包住了整个推理循环,而不是只在某个函数里写了一下。还有个容易忽略的点,你是不是用了model(x)之后还取了loss或者调用了backward?如果只是预测,直接拿logits就行。empty_cache只是把缓存还给PyTorch分配器,并不是释放给系统,所以看着没用很正常。我建议你加两行打印看看:torch.cuda.memory_summary()和torch.cuda.max_memory_allocated(),能直接定位是哪一层爆的。如果确认代码没问题,那可能是你用half精度或者past_key_values没清干净,BERT的cache tensors会随着序列长度线性增长。最后实在不行,可以试试torch.inference_mode()替代no_grad,这个更激进,能省掉不少推理时的额外开销。
推理比训练还吃显存,八成是保存和加载的方式有问题。你试试直接torch.save(model, path)整个模型存下来,加载的时候torch.load(path, map_location='cuda'),别用state_dict重新构建,有时候中间变量没释放干净就会这样。另外检查一下是不是忘了model.half()或者输入张量还在requires_grad状态。我之前也遇到过类似情况,最后发现是dataloader里pin_memory和num_workers设太大,推理根本不需要这些。
这个现象挺典型的,我之前也踩过类似的坑。你训练时4G是因为有梯度、优化器状态这些额外开销,但推理时反而涨到10G,大概率不是no_grad的问题,而是加载模型的方式有鬼。很多人保存state_dict后推理时直接重新实例化一个完整模型,如果config里某些参数跟训练时不一致,比如hidden size或者num labels对不上,PyTorch有时会静默重建导致显存异常。还有一种可能是你加载权重时没加map_location,权重先落到CPU再复制到GPU,中间临时占用翻倍。另外check一下是不是在__init__里就把模型推到了cuda,然后又load_state_dict,这样会多一份副本。建议打印一下torch.cuda.memory_allocated和memory_reserved的差值,看看是不是碎片或者缓存没释放。实在不行就用torch.cuda.memory_summary()看看到底哪一层吃掉了显存,比瞎猜快多了。
推理显存暴涨一般不是no_grad没写对,更可能是你把模型放到GPU之后又保留了计算图或者缓存没释放。先确认下加载state_dict时有没有用map_location,以及有没有不小心把优化器状态也load进来了。还有个常见坑是输入张量带着梯度历史,检查下dataloader或预处理里有没有detach。建议用torch.cuda.memory_summary看看具体是哪块占的,比瞎猜快多了。
可能是梯度没关干净,建议检查一下推理脚本里有没有漏掉的requires_grad=True,或者某些层没被no_grad包住。也有可能是加载模型时没指定map_location,权重被复制到了多张卡或者CPU和GPU之间来回倒腾。另外BERT这种模型可以试试用torchscript或者onnx导出再推理,显存能降不少。我之前遇到过类似情况,最后发现是dataloader里pin_memory没关,虽然batch=1但预取逻辑还在跑。