最近在调一个BERT分类模型,训练完保存了state_dict,然后单独写了个推理脚本加载。奇怪的是,训练时显存占用大概4G左右,但一跑推理,显存直接飙到接近10G,然后OOM。我明明已经用了torch.no_grad(),也调了model.eval(),还试了torch.cuda.empty_cache(),都没什么效果。输入数据也就是一条文本,batch_size=1。是不是我哪里写得不规范?还是说PyTorch推理本来就这么吃显存?求有经验的朋友指点一下排查思路,或者有没有什么通用优化技巧,谢谢!
PyTorch模型推理时显存暴涨,是哪里写错了?求大佬指点
全部回复
共 98 条八成是推理时忘了关梯度或者把优化器状态也load进来了,试试只load模型参数再清一下cuda缓存。
检查下是不是开了gradient checkpointing或者模型里有dropout没设成eval,有时候是中间变量没释放。
推理时显存比训练还高,大概率不是模型本身的问题,而是推理脚本里某个环节把显存“粘”住了。比如你加载state_dict后有没有把模型也放到cuda上?或者输入数据是不是忘了加.cuda()导致CPU和GPU之间反复拷贝?另外,BERT这种模型即使batch=1,如果序列长度没限制,attention矩阵也会随长度平方增长,建议检查一下tokenizer有没有截断。还有个冷门坑:如果你用了torch.no_grad()但模型内部有dropout或batch_norm的training=True残留,也会额外缓存激活值,试试加载后显式model.train(False)再跑一次。最后实在不行,可以开个profiler看看具体是哪一层爆的,比瞎猜快。
试试关掉梯度保存,检查下是不是加载了优化器状态或者dropout没关干净。
训练时是不是开了gradient checkpointing?推理时反而把激活都缓存了,查下模型里有没有eval模式不生效的缓存逻辑。
检查下是不是把梯度也存了,推理时关掉requires_grad或者用torch.inference_mode试试。
跑推理前先清一下缓存,另外看看是不是模型里有dropout或batch norm在eval模式下没生效。
之前跑BERT推理也遇到过类似的坑,最后发现是优化器状态没清干净——如果你加载state_dict的时候顺手把optimizer的checkpoint也load进去了,即便不调用optimizer.step(),显存里也会常驻一整套momentum和variance的buffer,那玩意儿比模型本身还占地方。建议你直接打印一下model.parameters()和optimizer的显存占用对比看看,或者干脆只加载model的权重,别碰optimizer。另外torch.no_grad()只管梯度计算,但如果你在推理脚本里不小心把输入也requires_grad=True了,或者模型里某层有显式创建临时张量的操作(比如attention mask的扩展),一样会吃显存。还有个容易忽略的点是transformers库的tokenizer会返回attention_mask和token_type_ids,如果这些没转移到GPU上,某些版本会自动隐式转换,反而多占一份显存。至于empty_cache,它只是释放缓存块,并不会压掉已分配但还没释放的张量,所以基本没用。最直接的排查办法是加个显存快照工具,比如torch.cuda.memory_summary(),看看到底是哪个操作把显存撑爆的,通常能看到具体分配堆栈。另外一个小技巧是推理时把batch_size=1的输入也pack成连续内存,避免碎片化,有时候显存暴涨纯粹是碎片太多导致的。最后如果还不行,试试half精度推理,显存直接砍半,但注意某些层要手动保持float32防止数值溢出。
我猜可能是推理时把输入和模型参数都放到了同一个device上但没注意梯度图残留?虽然no_grad了,但某些操作比如自定义forward里的中间变量或attention mask没显式指定device,可能隐式创建了临时张量。建议用torch.cuda.max_memory_allocated()打一下峰值,看看是不是真的涨在显存还是显存碎片化。另外试试把batch_size降到1的同时,输入token长度截断到128,看能不能压下来,我之前碰到过类似问题,最后发现是位置编码没设buffer导致每次前向都重新计算。
或者你检查下是不是加载state_dict时用了model.load_state_dict(torch.load(...))但没map_location='cuda',导致CPU和GPU各存了一份?还有个偏方,推理前先跑一个dummy输入做warmup,有时候能触发一些lazy初始化,但你这个显存暴涨幅度有点大,更像是有个中间变量被重复计算了,比如在forward里循环了序列长度。总之先看看峰值分布,再考虑用torch.jit.script或者onnx导出跑推理,能省不少显存。
- 我之前也踩过这个坑,大概率不是代码问题,而是推理时把整个计算图都保留了,试试在输入上加
torch.no_grad()包裹一下,或者检查下是不是用了torch.tensor而不是torch.from_numpy。 - 另外,显存暴涨也可能是推理时模型内部缓存了中间激活值,比如BERT的attention mask没处理好,建议用
torch.jit.script或者torch.inference_mode()代替no_grad,能省不少内存。 - 实在不行就开
torch.cuda.set_per_process_memory_fraction限制一下,看看是不是有别的显存碎片问题,不过你这个10G确实有点夸张,正常bert-base单条推理不该超过2G。
我之前也踩过类似的坑,光靠no_grad和eval其实拦不住显存暴涨,关键得看推理脚本里有没有不小心把梯度图给保留了。比如你加载state_dict之后,有没有对模型参数做任何操作?像是对某个tensor调用了requires_grad_,或者把输入也设成了可训练参数,这种都会导致显存翻倍。还有个小细节,BERT的position_ids和attention_mask如果是从输入动态生成的,有时候会隐式地创建计算图,建议把所有输入都包在with torch.no_grad()外面,而不是只包forward那一步。另外你试过用torch.inference_mode()吗?这个比no_grad更彻底,能禁掉所有自动求导机制,我换了这个之后显存直接降了30%。还有个排查技巧,在推理循环里每跑一次就打印一下torch.cuda.max_memory_allocated(),看看是不是第一次forward就爆了,还是累积到后面才爆,这能区分是模型结构问题还是数据流问题。如果确认不是代码问题,可能是你加载的checkpoint里带了优化器状态或者某些buffer,用load_state_dict(strict=False)看看有没有多余的键,顺便清理一下。最后实在不行,可以试试把模型转成半精度fp16推理,显存能砍一半。
我猜你是不是在推理时把输入也送进了GPU,但忘了对输入做input_ids = input_ids.cuda()之外的with torch.no_grad()包裹整个forward过程?有时候model.eval()只管BN和Dropout,但no_grad如果只包了部分计算图,中间变量还是会留在显存里。另外可以试试把torch.cuda.empty_cache()放在每个batch之后,或者检查一下是不是推理时无意中把梯度也计算了,比如调用了loss.backward()。我之前遇到过类似问题,最后发现是attention_mask的维度不对,导致中间张量爆炸,你用torch.no_grad()包住完整推理流程再试试,不行就打印一下每层的显存占用。
查一下是不是把优化器或者梯度也load进来了,光load state_dict不该吃这么多显存。
试试把输入和模型都half()转精度,能省一半显存。
试试关掉梯度再把输入也detach一下,另外看下是不是模型里缓存了中间变量没清。
训练时梯度累积的缓存释放不掉,推理时模型里某些层可能保留了activations,查下有没有用torch.jit或者混合精度。
我之前也遇到过一模一样的情况,训练4G推理反而爆显存,最后发现是优化器状态和梯度释放的问题。你训练时显存包含了优化器参数,但推理时按理说应该远小于训练才对,除非模型里还保留了dropout或batchnorm的缓冲区,或者你加载state_dict时不小心把训练时的额外缓存也带上了。建议先排查一下是不是在forward里用了类似torch.save(tensor)或列表累积中间变量,BERT的中间层输出如果不及时释放,batch_size=1也能堆出几个G。另一个很隐蔽的点是,如果你在推理脚本里忘了model.eval(),但代码里又用了torch.no_grad(),这俩不冲突,可如果模型里有自定义的buffer(比如位置编码的cache),它会在每个step不断增长,我踩过这个坑。你试试在推理循环里隔几步打印torch.cuda.memory_allocated(),看是线性增长还是跳变,线性增长基本就是缓存泄漏。还有torch.cuda.empty_cache()只是把空闲缓存还回去,并不会阻止新分配,真正要控制的是别在循环里创建新tensor。如果急着跑,可以先转成半精度model.half(),显存直接砍半,但注意输入也要转。最后实在不行,试试torch.inference_mode()替代no_grad,它连autograd的跟踪都彻底关了,能省一点。
说实话看到你这个情况我第一反应是,训练能跑通但推理OOM,大概率不是模型本身的问题,而是你推理脚本里某些代码隐式地构建了计算图。虽然你加了no_grad,但如果你在模型外面套了比如梯度累积、或者用了某些自定义loss函数里的中间变量,这些都可能让显存爆炸。我建议你先用torch.cuda.memory_summary()看看具体是哪一层分配的显存,这样定位最快。
另外有个很常见的坑,就是你把整个验证集或者测试集一次性load进显存做推理,即便batch_size=1,如果数据加载器里没设置pin_memory=False或者num_workers太大,也会导致临时缓冲区占用飙升。你可以试试把推理循环里的inputs和labels都显式地放到cpu上,再调用empty_cache,看看峰值有没有降下来。
还有啊,BERT模型本身如果用了分词器的padding策略,一条文本也会被pad到最长序列长度,你检查一下max_length是不是设的特别大。我之前遇到过类似问题,最后发现是tokenizer默认把序列padding到512,但训练时是动态padding,推理时忘了改,显存直接翻倍。实在不行你试一下torch.utils.checkpoint,把前向传播切成几段,能省不少显存,就是会慢一点。
这问题我踩过一模一样的坑,而且当时比你更蒙,因为我连训练都没跑,直接加载别人给的权重做推理都能爆显存。你试的那三板斧其实都是常规操作,但大概率问题出在模型本身或者输入数据上,我猜你八成是没关掉梯度缓存之外的什么东西。比如,如果模型的forward里用了像dropout或者batch norm之外的某些层,在eval模式下依然会保留训练时的激活值缓存,尤其BERT里那些attention的中间变量,batch_size=1也可能因为序列长度太长导致显存爆炸。另一个常见坑是,你是不是在推理脚本里也把优化器或者loss相关的模块加载进来了?哪怕你没调用,只要这些对象还持有计算图引用,显存就释放不掉。我的建议是,先别管empty_cache,用torch.cuda.max_memory_allocated()打印一下峰值分配,看看是不是在第一次前向就爆了,如果是,就检查输入ids的attention_mask是不是不小心传了全1,或者有没有把整个tokenizer的output都塞进模型。还有就是试试把模型包装成torch.compile或者用half精度跑,能显著降显存。最后,如果实在排查不出来,直接上torch.profiler看每一层的显存分配,那个最直观,我当时就是靠它发现是某个自定义的layer norm实现写了缓存变量没清理。
这情况我也踩过坑,重点查下推理脚本里是不是把优化器状态也load进来了,或者模型里还挂着dropout和BN的training模式。另外试试把输入tensor用torch.no_grad()包起来,再手动跑一次torch.cuda.synchronize()看看到底哪一步显存峰值最高。还有个隐蔽点,如果用了transformers库,记得关掉output_attentions和output_hidden_states,这俩默认可能开着,一条文本也能炸出几个G的中间变量。
这情况我碰到过,多半不是模型本身的问题,而是推理脚本里把梯度也带进去了。你试试在加载完state_dict之后,给所有参数加个requires_grad_(False),或者干脆用torch.inference_mode()代替no_grad(),它能彻底关掉自动求导的跟踪机制,显存能降一大截。
另外检查一下是不是把整个tokenizer和模型都塞进GPU了,有时候输入ids的维度没对齐,或者attention mask忘了传,也会导致隐式创建超大中间张量。还有个小技巧,推理时把batch_size设成1但别开grad,同时用torch.cuda.max_memory_allocated打个峰值,看看是不是有碎片化累积的问题。
如果还是高,直接用torch.jit.script或者onnx导出试试,BERT这种结构优化空间挺大的。反正4G到10G这个跳变不正常,肯定不是正常推理的消耗,八成是哪里不小心把训练模式的部分逻辑带进来了。
这种情况我去年也踩过,最后发现是HuggingFace的BertModel在forward里默认返回了tuple,你如果直接拿output[0]去做分类,其实整个hidden_state都被保留了。可以先试试把模型的return_dict设成False,或者只用最后一层的pooler_output,能省不少显存。另外你说训练才4G,推理反而10G,很可能是推理脚本里把optimizer和梯度相关的缓存也加载进来了,毕竟state_dict只存参数,但如果你脚本里不小心创建了optimizer实例,它会额外申请显存。还有个坑是torch.no_grad()只管autograd,但如果你用了model.eval()之后又调用了model.train(),某些层比如Dropout和LayerNorm的缓存行为会变。建议你用torch.profiler跑一下,看具体是哪一层分配的显存最多,我之前就是这么定位到是attention的score矩阵爆了。实在不行就开amp混合精度推理,float16能直接减半,但注意要转换输入数据。最后检查下你的输入是不是被自动拼接了多个batch,有时候DataLoader的collate_fn会在batch_size=1时还做padding到最大长度,也会导致显存虚高。
我也遇到过类似的情况,最后排查下来其实不是模型本身的问题,是推理脚本里忘了把optimizer和梯度相关的缓存清掉。你训练完保存state_dict之后,如果是在同一个进程里直接切到推理模式,optimizer里的动量、梯度缓存还占着显存,再加上CUDA的缓存碎片,显存翻倍很正常。建议你试试把推理放到一个全新的脚本里,或者至少用torch.cuda.reset_peak_memory_stats()看下峰值到底分配在哪。另外,BERT这种模型如果开了gradient_checkpointing,推理时反而可能因为要重建中间激活而更吃显存,你检查下模型配置里有没有这个开关。还有个容易忽略的点是model.eval()只影响dropout和BN,但如果你用了torch.no_grad()却还是把输入传进了model,其实中间变量还是会被保留在计算图里,除非你显式用with torch.no_grad():包住整个forward。如果这些都排除了,那大概率是CUDA缓存没释放,试试torch.cuda.set_per_process_memory_fraction(0.8)限制一下,或者用torch.cuda.memory_summary()看看具体哪一层分配最多。我最后是用torch.jit.trace把模型脚本化之后推理,显存直接降了30%,你可以考虑下。
碰到过类似的坑,大概率不是模型本身的问题,而是推理脚本里不小心把梯度图给保留了。可以检查下输入有没有设requires_grad=False,或者模型里有没有dropout之类在eval模式下还正常工作的层。
另一个常见原因是加载state_dict时把整个优化器状态也带上了,或者用了model.train()后又忘记切回eval,这些都会让显存翻倍。建议用torch.inference_mode()替代no_grad,它能彻底禁用梯度跟踪,有时候empty_cache没效果是因为显存碎片化,可以先试试在推理前强制释放一次。
如果还不行,试着把输入移到CPU上跑一次对比,或者用torch.profiler看下具体哪一层爆的,通常定位到问题就快了。