最近在调一个BERT分类模型,训练完保存了state_dict,然后单独写了个推理脚本加载。奇怪的是,训练时显存占用大概4G左右,但一跑推理,显存直接飙到接近10G,然后OOM。我明明已经用了torch.no_grad(),也调了model.eval(),还试了torch.cuda.empty_cache(),都没什么效果。输入数据也就是一条文本,batch_size=1。是不是我哪里写得不规范?还是说PyTorch推理本来就这么吃显存?求有经验的朋友指点一下排查思路,或者有没有什么通用优化技巧,谢谢!
PyTorch模型推理时显存暴涨,是哪里写错了?求大佬指点
全部回复
共 98 条我之前也踩过类似的坑,训练能跑但推理爆显存,大概率不是模型本身的问题。你试试把优化器状态清掉再测,有时候训练完optimizer里的动量缓存还占着显存,虽然你load的是state_dict,但进程没重启的话,之前的缓存可能没释放干净。另外,检查一下推理脚本里是不是不小心把梯度也保留下来了,比如某个操作没包在no_grad里,或者用了.backward()的变体,虽然写了no_grad但某些自定义层里可能混了需要梯度的参数。还有一个很隐蔽的点,BERT的attention mask如果没正确传给模型,可能会触发动态padding或者不必要的内存分配,尤其是当输入长度变化时,PyTorch会为每个batch重新分配临时tensor。你试下固定输入长度,或者用torch.jit.script把模型trace一下,这样能避免很多动态图的隐形开销。至于empty_cache(),它只是释放缓存块,不会减少已分配但未释放的内存,所以别指望它救急。如果实在排查不出来,可以开torch.autograd.set_detect_anomaly(True),虽然慢但能定位到具体哪一行导致显存异常。最后提个通用技巧,推理时把模型转成half精度,显存直接减半,或者用torch.inference_mode()替代no_grad(),后者在最新版本里更激进,能省掉一些推理专用的中间变量。
这个情况我踩过坑,大概率不是推理本身的问题,而是加载模型时把optimizer的state_dict也一起load了,或者模型里还挂着训练阶段的缓存变量。你试试只load模型参数,然后跑之前先torch.cuda.synchronize()看看峰值。另外检查下是不是在推理脚本里不小心把梯度打开了,比如某个模块调用了.requires_grad_(),那no_grad就形同虚设了。之前我遇到类似情况是embedding层被反复forward导致显存碎片化,加个torch.cuda.set_per_process_memory_fraction限制下能缓解。
我之前也踩过类似的坑,训练时反向传播的梯度会及时释放,但推理时如果模型里有dropout或者batchnorm以外的层缓存了中间激活值,显存可能反而比训练还夸张。你说batch_size=1还暴涨,先检查一下是不是在循环里反复调用model,每次forward都会重新构建计算图,虽然no_grad不存梯度,但某些操作(比如attention mask的广播)还是会临时申请大块显存,用完没释放。另外,可以试试把输入数据用torch.no_grad()包住后再传到CUDA,或者直接传CPU张量让模型内部转,有时候是数据搬运时多了一份拷贝。还有个冷门技巧:在推理前加一行torch.cuda.set_per_process_memory_fraction(0.5),限制最大显存,这样OOM时会提前报错方便定位哪一步爆的。如果模型里有梯度检查点(checkpoint)或者自定义的autograd.Function,也可能是它们没走标准释放逻辑。最直接的排查方法是用torch.profiler或者nvidia-smi -l 1盯着看,每跑几行代码就打印一次显存,找到那个突增的节点。最后实在不行,可以试试torch.jit.script或者onnx导出,有时候绕开python动态图能省一大截。
盲猜你是不是把optimizer的state_dict也load进推理脚本了?或者模型里还挂着dropout以外的训练专属层没冻结,但更常见的是你用了transformers的pipeline,它内部会缓存attention的key/value,batch_size=1也架不住序列长。建议先用torch.cuda.max_memory_allocated()打一下峰值,再试试把输入token截短到128,大概率能降下来。另外检查下有没有在no_grad块外面调了.cpu()或者.numpy(),这种隐式同步也会临时拉高显存。
试试关掉梯度再把batch里最长序列截断,BERT推理峰值显存主要卡在attention矩阵上,你这涨幅不像正常波动。
这个现象我踩过坑,大概率不是PyTorch推理本身的问题,而是你的推理脚本里混进了训练时的“惯性操作”。比如是不是在加载模型后,又意外调用了model.train(),或者虽然写了no_grad,但输入张量本身带了requires_grad=True,这会导致中间激活值被保留,显存直接翻倍。另外检查一下是不是把优化器也load进来了,有些代码会顺手optimizer.load_state_dict,虽然不参与前向,但会占用额外显存。还有个隐蔽点:如果用了HuggingFace的from_pretrained,默认会加载配置里的output_hidden_states=True,这会额外缓存所有层的输出,显存瞬间爆炸。我建议你先用torch.cuda.memory_summary()看具体分配在哪一层,或者把输入改成torch.zeros随机初始化跑一遍,排除数据本身的问题。优化技巧的话,试试torch.inference_mode()替代no_grad,再配合half()半精度推理,基本能压到2G以内。如果你代码里还手动调了batch_first或者padding时生成了attention_mask,记得确认这些张量都放在cuda上,不然CPU-GPU拷贝也会导致峰值内存异常。最后实在不行,可以跑一次纯CPU推理对比显存占用,如果CPU正常而GPU暴涨,那基本就是缓存未释放的问题,可以用torch.cuda.reset_peak_memory_stats()定位峰值点。
之前也踩过类似的坑,你这情况大概率不是单条数据的问题,而是推理脚本里没关梯度或者优化器状态残留。可以试试在加载模型后显式调一下model.zero_grad(),或者检查是不是把整个验证集都塞进DataLoader了,有时候batch_size=1但shuffle或者drop_last没设置好也会触发额外开销。另外建议用torch.profiler跑一下看看具体哪层爆的,我之前就是attention的mask没广播对,导致中间变量翻倍。还有个小技巧,推理前把模型转成half精度,显存直接砍半,效果基本无损。
训练时4G推理却冲到10G,这明显不是正常现象,你八成是踩了显存累积的坑。试试在推理循环里每跑完一个batch就加一行torch.cuda.synchronize(),然后立刻检查torch.cuda.memory_summary(),看看是不是某个中间变量没被释放,比如hidden_state或者注意力mask被意外保留了引用。另一个常见原因是model.eval()只改了BN和Dropout,如果模型里有LSTM或者自定义的循环结构,torch.no_grad()并不能阻止计算图在反向传播的hook上挂住,建议换成with torch.inference_mode():,这个模式更彻底。另外,你保存的是state_dict,加载时有没有确认model.to('cuda')之后输入也是cuda?如果输入还在CPU,PyTorch会自动拷贝并可能创建额外的临时张量。还有个隐藏坑:如果推理脚本里不小心调用了loss.backward()(比如为了调试梯度),即使数据只有一条,梯度也会占掉显存。最后,直接试一下torch.cuda.set_per_process_memory_fraction(0.5, 0)来强制限制显存,如果OOM报错位置变了,就能定位到是峰值分配而不是泄漏。我上次遇到类似问题,结果是tokenizer返回的attention_mask忘记.to('cuda'),导致每个batch都在CPU和GPU之间反复搬数据,显存碎片化严重。你可以先用torch.cuda.memory_allocated()和torch.cuda.memory_reserved()对比一下,如果reserved远大于allocated,那就是缓存没复用,试试torch.cuda.empty_cache()放在每个batch之后而不是之前。
训练时4G推理却飙到10G,这明显不是正常现象,因为eval模式下没有梯度图,显存占用理应比训练低一个量级。我怀疑你八成是把整个模型+优化器+梯度都加载进显存了,比如直接torch.load了整个checkpoint而不是只loadstate_dict,或者推理脚本里还残留着训练时的优化器参数。另一个很常见的坑是,如果你用HuggingFace的Trainer保存的模型,有时候会附带一些buffer(比如position_ids),这些buffer在推理时反而会被显式创建到GPU上,但不该占这么大的量。
你可以先跑一个最简单的测试:加载模型后,把输入换成随机tensor,看看显存是否还是爆炸。如果随机输入正常,那就是你的tokenizer或数据预处理环节把某些东西放到GPU上了,比如把attention_mask也扔进.cuda()了。另外,检查一下是不是用了torch.jit.trace或者torch.compile,某些版本的编译模式会在推理时缓存中间激活,导致显存峰值暴涨。
还有个小技巧,用torch.cuda.memory_summary()看下分配的大块内存都在哪个操作上,比empty_cache管用得多。如果实在着急,可以先开torch.cuda.set_per_process_memory_fraction(0.5)强制限制显存,看看是不是真的OOM还是只是峰值吓人。最后,BERT推理通常几百MB到1G就顶天了,你这情况十有八九是某个tensor没卸载或者反复复制,建议从checkpoint加载逻辑开始逐行排查。
碰到过类似的情况,最后发现是模型里有个dropout层在推理时没被真正关闭,虽然调了eval但某些自定义层没走eval的逻辑,建议你检查下有没有自己写的forward里带了训练时才用的缓存或者中间变量。另外BERT推理显存高有个常见坑是输入长度没做padding到固定长度,虽然你batch_size=1,但序列如果很长,注意力矩阵是平方级增长的,10G也不算离谱。还有个思路是试试用torch.jit.script或者onnx导出,能省不少显存,因为少了autograd的图结构开销。你可以先用torch.cuda.memory_summary()看下到底是哪块分配的,是模型参数还是激活值,这样定位更准。如果只是跑单条推理,也可以考虑把模型切成fp16,显存直接减半,但注意下精度影响。实在不行就换batch_size=1加上gradient_checkpointing的推理模式试试,虽然训练用的技巧但有些场景下对推理也管用。
跑推理显存比训练还高大概率不是模型本身的问题,你检查下是不是加载了optimizer的state_dict或者把梯度也存进去了,推理时如果没冻结所有参数,某些op(比如dropout的mask或者中间激活)还是会缓存。另外试试把输入换成全零tensor跑一遍,如果显存正常那就是数据侧的坑,比如文本长度没padding到固定值。还有个小技巧,推理前用torch.cuda.synchronize()看下实际峰值,empty_cache只是清缓存不释放进程占用的显存。我之前遇到过类似情况,最后发现是模型里有个没用的BatchNorm层在eval模式下反而触发额外buffer更新,删掉就好了。
这情况我遇到过,大概率不是torch.no_grad的问题,而是你推理脚本里把optimizer或者训练时的loss计算逻辑也带进来了。BERT推理吃显存确实比训练高,但10G肯定不正常,先检查下是不是加载state_dict时把模型定义成了训练模式下的结构,比如dropout层在eval和train下行为不同,但显存暴涨一般不是这个引起的。
我猜你可能是把整个训练函数原封不动复制到推理里了,里面如果还有backward或者计算梯度相关的操作,哪怕有no_grad,某些算子比如batch norm的running stats更新也会额外占显存。建议你单独写个干净的forward函数,只保留model(input)和输出处理,别碰任何优化器状态。
另外,你试试把输入tensor显式用torch.no_grad包裹,或者直接把模型输入改成half精度,BERT在fp16下显存能砍一半还多。还有,torch.cuda.empty_cache只是清空缓存池,并不会释放已分配显存,真正的杀手可能是你加载了多个模型副本,或者data_loader里做了tokenize时保留了过多中间变量。
我上次碰到类似问题,最后发现是tokenizer返回的attention_mask没移到GPU上,导致CPU和GPU之间反复拷贝,显存碎片化严重。你检查下input_ids和attention_mask是不是都在cuda上,还有别用batch_size=1但seq_len特别长的输入,BERT对长序列的显存开销是平方级的。如果还不行,建议用torch.jit.script或者onnx导出,推理速度更快,显存也干净。
我之前也踩过类似的坑,训练时梯度会释放,但推理时如果模型里挂了dropout或者batchnorm的training状态没彻底关干净,某些层可能还会保留缓存。建议先试下torch.inference_mode()替代no_grad,能省不少显存。另外检查下是不是加载了训练时的优化器状态,或者模型里有没被requires_grad绑定的中间变量,比如attention mask的广播。如果还不行,可以试试把输入换成空tensor跑一遍,看显存基线是多少,一步步定位是哪层爆的。
说实话训练4G推理10G这个现象我见过不少次,大概率不是单点写错,而是推理脚本里把训练时的一些缓存或梯度相关的东西带进来了。比如你加载state_dict后有没有把model的requires_grad全部置False?即使你用了no_grad,只要模型参数还带着requires_grad=True,某些自定义forward里的操作还是会额外分配内存。另外,你确认推理脚本里没有意外调用optimizer.zero_grad或者loss.backward之类的残留代码吗?我之前就踩过这种坑,把训练循环的片段复制过来忘删了。还有一个很隐蔽的点,如果用了transformers库,attention_mask和token_type_ids的dtype或者device不对,也会触发隐式转换导致显存爆炸。建议你先在推理入口加个torch.cuda.synchronize()看看实际峰值,然后逐步注释掉forward里的子模块,二分定位到具体哪一层暴涨。优化的话,可以试试torch.inference_mode()替代no_grad,省掉一部分自动求导的元数据开销,另外把输入统一放到GPU上,避免CPU-GPU反复拷贝。如果还不行,就开一下torch.autograd.set_detect_anomaly(True)跑一次,虽然慢但能抓到具体是哪个变量在搞鬼。
可能是优化器或梯度缓存没清干净,试试加载完state_dict后手动清一下cuda cache。另外检查下模型里有没有意外开启的dropout训练模式。
试试把输入也放到with torch.no_grad()块里,之前遇到过输入requires_grad没关导致显存翻倍的情况。
之前跑GPT类模型也踩过类似的坑,后来发现是推理时没关梯度计算图以外的缓存,比如优化器状态或中间变量没释放。你可以试试在推理脚本里只加载model,别把optimizer和scheduler的state_dict也load进去,那个超占显存。另外检查下是不是用了torch.no_grad但没把输入也detach,或者forward里不小心把logits又传给了自己。实在不行就开个torch.cuda.memory_summary()看下具体哪块分配爆的,比瞎猜快多了。
推理时显存比训练还高,大概率是优化器状态或梯度缓存没清干净,试试加载完权重后手动清一下缓存。
训练时反向传播会释放中间变量,推理反而可能缓存激活值,试试在forward里加torch.no_grad或检查下有没有重复加载模型。
大概率是模型里存了训练时的优化器状态或梯度,加载state_dict时别带上这些,只load模型参数试试。
检查下推理时是不是没关grad,或者模型里有没有累积梯度的操作,eval模式有时也会漏掉这些。
说实话你这种情况我踩过好多次坑,先说结论:真不是PyTorch推理本身吃显存,BERT base也就400多M参数,fp32跑batch_size=1根本不可能到10G。我猜大概率是加载state_dict的时候没把模型挪到GPU上,然后推理时输入和模型一个在CPU一个在GPU,导致PyTorch自动把模型整体搬进显存,再加上优化器状态或者某些中间buffer没释放,就爆了。你可以在加载完权重后打印一下model.device,或者直接试试model = model.cuda()再跑一次,看显存是不是立刻降下来。另外检查下是不是推理脚本里不小心又把训练时的优化器、loss或者梯度相关的张量给初始化了,哪怕没backward,只要创建了这些对象,它们就会占显存。还有个隐蔽的点,如果用了transformers库,attention_mask和token_type_ids没传到input_ids同样的设备上,也会触发隐式拷贝,显存翻倍不奇怪。建议你开一下torch.cuda.memory_summary(),看看到底哪块分配占了大头,比空猜强。最后实在不行,可以试试torch.inference_mode()替代no_grad,它更激进,会关掉自动梯度追踪和一部分缓存机制,某些场景能省不少。