最近在调一个BERT分类模型,训练完保存了state_dict,然后单独写了个推理脚本加载。奇怪的是,训练时显存占用大概4G左右,但一跑推理,显存直接飙到接近10G,然后OOM。我明明已经用了torch.no_grad(),也调了model.eval(),还试了torch.cuda.empty_cache(),都没什么效果。输入数据也就是一条文本,batch_size=1。是不是我哪里写得不规范?还是说PyTorch推理本来就这么吃显存?求有经验的朋友指点一下排查思路,或者有没有什么通用优化技巧,谢谢!
PyTorch模型推理时显存暴涨,是哪里写错了?求大佬指点
全部回复
共 98 条试试关掉梯度之后把优化器也清掉,尤其是混合精度下梯度缓存特别占显存。我之前也遇到过,基本就是这原因。
试试关掉梯度计算图里的缓存变量,或者检查下是不是把优化器状态也一起加载了,那玩意吃显存很离谱。
遇到过类似的坑,大概率不是模型本身的问题,而是推理脚本里不小心把梯度图给保留了。比如对输入张量调用了requires_grad_(),或者模型里有dropout以外的层在eval模式下还在构建计算图,试试在no_grad块里把输入也detach一下,再确认下是不是用了transformers的pipeline,那个有时候会额外缓存东西。
另外可以开一下torch.cuda.set_per_process_memory_fraction限制显存,看它是直接OOM还是慢慢涨,如果慢慢涨可能是缓存碎片问题,empty_cache其实治标不治本。我之前用torch.jit.trace把模型转成script模式推理,显存直接降了一半,你可以试试这个方向。
还有个容易忽略的点,就是batch_size=1但序列长度如果没padding到固定长度,动态shape会导致cudnn每次重新申请workspace,显存峰值会很高。建议给输入加个max_length限制,或者用torch.utils.checkpoint把中间激活值换成重计算,虽然慢点但显存能压下来。
这情况我遇到过,大概率不是模型本身的问题,而是推理脚本里把optimizer或者训练状态也load进来了,或者不小心把梯度打开了。你可以查一下是不是有model.zero_grad()或者requires_grad没关干净,再有就是检查下输入有没有被重复包装成多个batch。另外试试把输入和模型都half()一下,显存能省一半,反正BERT推理精度影响不大。
还有个排查思路,你可以在推理时打印一下每个op的显存占用,或者用torch.cuda.set_per_process_memory_fraction限制一下看看是不是真的涨到10G。我之前碰到过类似情况,最后发现是torch.no_grad()没包住所有forward,比如有个地方不小心调了.backward()。如果还不行,就换个思路,把模型转成ONNX或者用torch.jit.trace,不仅快还省显存。
训练时4G推理10G这个现象不太正常,我怀疑你推理脚本里可能不小心把梯度也保留了,比如模型输出后调用了backward或者loss计算,试试在forward前后加个torch.cuda.synchronize()看峰值是不是瞬时冲高的。另外检查下是不是加载模型时把optimizer的state_dict也load进去了,那个会额外占不少显存。我之前遇到过类似问题,最后发现是推理时忘了对输入做no_grad包裹的变量赋值,导致计算图被保留,你可以在forward后加一句del output再手动释放下缓存看看。
你这情况我遇到过,多半不是模型本身的问题,而是推理脚本里把优化器或者梯度相关的状态也加载进来了,或者压根没做梯度清零。试试在eval()之后再加一句optimizer.zero_grad(),另外检查下是不是用了torch.no_grad()但没包住整个forward过程。还有个坑是transformers库的模型有时会缓存中间激活,可以试试把model.config里use_cache设成False,再不行就看看是不是数据加载时把整个数据集都塞进GPU了,先排除这些再考虑优化技巧。
试试关掉梯度保存,把输入也detach一下,还有检查下是不是有个隐藏的优化器没释放。
训练时4G推理10G这个现象不太正常,你检查下是不是把优化器状态或者梯度也一起load进来了,或者模型里还有dropout之外的训练层没关干净。我之前遇到过类似情况,最后发现是推理脚本里不小心把输入也requires_grad了,虽然no_grad包着但显存还是被占着。另外可以试试把batch里的padding都去掉,或者用half精度跑一下,BERT推理一般不会比训练吃更多显存的。
你这情况我遇到过,八成不是推理本身的问题,而是加载权重时把优化器状态或者梯度也带进来了,试试只load_state_dict里model的部分,别整棵模型对象。另外检查下是不是把验证集也喂进模型了,或者DataLoader的num_workers开太多,有时候内存碎片也会导致显存虚高。实在不行就开个profiler看看哪一层炸的,比瞎猜快多了。
我之前也踩过这个坑,10G大概率不是模型本身的问题,而是推理脚本里不小心把梯度图给保住了。试试在加载完权重后加一行model.zero_grad()或者model.requires_grad_(False),还有检查下是不是用了torch.no_grad()但没把输入也设成requires_grad=False。另外,如果用了transformers库,记得把return_dict设成False,有时输出tuple会省不少显存。最后可以看看是不是缓存了所有中间激活值,比如用了torch.jit.trace或者torch.inference_mode()替代no_grad,后者在PyTorch 1.9+里更彻底。
我之前也踩过类似的坑,训练时梯度会释放一部分显存,但推理时如果模型里有dropout或者batchnorm的training状态没关干净,有些层会额外保留激活值。你确认下是不是加载完state_dict之后又调用了model.train(),或者模型定义里某些模块的training属性被意外置True了。另外,BERT这类模型如果用了attention_mask,但输入没做padding到固定长度,动态shape也可能导致缓存分配器碎片化,显存越涨越高。你可以试试把输入序列截断到固定长度,比如512,再比较下显存曲线。还有个隐藏点,就是transformers库的模型有时会默认开启gradient_checkpointing,但推理时它反而会缓存中间变量,你可以在加载后显式设置model.config.gradient_checkpointing=False。至于empty_cache,它只是把空闲缓存还给CUDA,并不解决实际分配问题,可以测一下在推理前后打印torch.cuda.memory_summary(),看看是哪个tensor占了大头。如果实在排查不出来,试下用torch.inference_mode()替代no_grad,它能彻底禁用自动求导和一部分内存跟踪。最后实在不行就换半精度推理,FP16能直接砍一半显存,效果立竿见影。
训练时4G推理却飙到10G,这明显不太正常,我怀疑是推理脚本里把优化器或者梯度相关的东西也加载进来了,或者模型里某个层在eval模式下反而缓存了激活值。你可以试试在推理脚本里只用model.load_state_dict加载权重,然后强制把requires_grad全设为False,再跑一下看显存曲线。另外,torch.no_grad()和empty_cache()确实治标不治本,重点排查下是不是有hidden state被意外保存了,比如某些自定义forward里return了中间变量。实在不行就开一下torch.autograd.detect_anomaly,虽然慢但能定位到具体哪一行爆的显存。
之前也踩过类似的坑,你试试把optimizer和scheduler的state_dict也一起load进来,有时候是它们把显存缓存带起来了。另外检查下推理脚本里是不是不小心把整个验证集或者dataloader的shuffle给打开了,数据加载也会占显存。还有个小技巧,输入前后加torch.cuda.synchronize()看看实际峰值,有时候是显存碎片化导致的假象。
我之前也踩过类似的坑,大概率不是推理本身的问题,而是加载权重时把优化器状态或者梯度也带过去了。你试试在加载state_dict之后,显式清一下optimizer的缓存,或者检查下是不是有hook没移除。
另外,BERT推理时显存暴涨很可能是输入长度没控制,长文本的attention矩阵是按平方增长的,一条文本也可能撑爆。你可以打印下输入的token长度,或者用torch.utils.bottleneck跑一下看是哪一层爆的。
还有个土办法,推理时把torch.cuda.amp.autocast()加上用fp16,显存能直接砍半,我试过效果很明显,不过要小心精度损失。
训练时是不是开了梯度检查点或优化器状态没释放?推理时加载的是完整模型还是带Dropout的?建议先排除模型结构里有没有意外保留训练专用层。
试试把输入和模型都转成半精度,再检查下有没有把整个验证集都塞进显存了。
有没有开gradient checkpointing?有时候推理阶段显存暴涨跟输入长度和attention计算有关,试试缩小max_len。
跑推理前先清一下CUDA缓存,再把模型用half精度加载,能省不少显存。
说到这个我太有共鸣了,之前也踩过一模一样的坑。你试试把optimizer的state_dict也一起load进来,或者检查下是不是在推理脚本里不小心创建了训练时的优化器,有时候这玩意儿会悄悄占掉一大块显存。另外,BERT系列模型跑推理时,如果没关掉gradient checkpointing或者没设max_length,哪怕batch_size=1,显存也可能因为序列长度被padding到512而暴涨,你试着打印下input_ids的shape看看是不是被撑大了。还有一种可能,就是你的模型里带了dropout或者batchnorm,虽然eval()能切换模式,但如果你用了torch.jit.script或者某些自定义层,可能没生效。最实用的排查办法是,在推理循环里用torch.cuda.max_memory_allocated()和torch.cuda.memory_snapshot()看下具体是哪一步分配的,别光看任务管理器。如果实在找不到,可以试试把输入tensor用torch.no_grad()包一层再传入,或者用half精度跑一次,显存能直接砍半。我最后是靠把模型里不必要的缓存变量清掉才解决的,你翻翻模型定义里有没有self.xxx = []这种累积列表的代码。
这个情况我见过不少次,大概率不是PyTorch推理本身的问题,而是推理脚本里某个地方不小心把梯度图给保留了,比如在循环里重复跑模型没清零,或者loss计算链被带进来了。你可以试试在推理循环外面只做一次model.eval(),然后检查一下有没有loss.backward()之类的残留代码。另外,如果用了transformers库,记得把return_dict=False或者检查一下attention_mask有没有正确传,有时候padding长度不一致也会让显存分配变得很夸张。我上次就是漏了个torch.inference_mode(),换成这个比no_grad更彻底,你试试看有没有改善。
看到你这个情况我第一反应是检查一下是不是把模型切到了训练模式但没真正冻结BN和Dropout,不过你说了调了eval应该不是这个。更可能的问题在于你加载state_dict后忘了把模型放到eval模式对应的设备上,或者推理脚本里不小心又跑了一遍forward的梯度计算,比如某些操作比如自定义loss或者梯度累积逻辑被带过来了。另外,一个常见坑是输入数据没有做torch.no_grad包裹下的.detach()或者没有统一放到cuda:0,如果数据还在CPU上,模型参数在GPU上,会触发隐式拷贝导致显存临时翻倍。我建议你直接用torch.profiler或者nvidia-smi监控一下每个op的显存分配,看看是不是某些中间变量比如attention mask没被释放。还有一个很实用的技巧,推理时用torch.inference_mode()替代torch.no_grad(),它能更彻底地跳过自动求导记录,减少控制流开销和显存峰值。如果还是涨,可能是模型本身有缓存机制,比如Transformer里的KV cache被默认保留,你可以试试把config.use_cache = False。最后实在没辙就开torch.cuda.memory_summary()看具体哪个张量占了大头,通常一下就能定位到问题。
这情况有点怪,正常推理不该比训练还吃显存。你检查下是不是加载state_dict时把优化器状态也一起load了,或者模型里还挂着dropout/training相关的缓存?另外试试把输入ids和attention_mask都放到同一个device上,有时候数据在CPU上会触发隐式转换导致额外内存。实在不行就用torch.jit.script或者半精度推理试试,能压不少占用。