最近在调一个BERT分类模型,训练完保存了state_dict,然后单独写了个推理脚本加载。奇怪的是,训练时显存占用大概4G左右,但一跑推理,显存直接飙到接近10G,然后OOM。我明明已经用了torch.no_grad(),也调了model.eval(),还试了torch.cuda.empty_cache(),都没什么效果。输入数据也就是一条文本,batch_size=1。是不是我哪里写得不规范?还是说PyTorch推理本来就这么吃显存?求有经验的朋友指点一下排查思路,或者有没有什么通用优化技巧,谢谢!
PyTorch模型推理时显存暴涨,是哪里写错了?求大佬指点
全部回复
共 98 条训练时4G推理反而10G,这确实不正常,多半不是代码逻辑问题,而是显存碎片化或者缓存没清干净。你可以试下在推理循环里加上torch.cuda.synchronize(),然后监控一下nvidia-smi看峰值是出现在加载权重还是前向传播那一步。另外,BERT推理用half()半精度能直接砍掉一半显存,如果还不行就检查下是不是有变量被意外保留到了计算图里,比如在循环外不小心引用了中间tensor。
我之前也遇到过类似情况,最后发现是推理时忘了关掉梯度缓存之外的optimizer相关状态,虽然你加载的是state_dict但可能没把optimizer也删干净。另外可以查一下是不是模型里带了dropout或batch norm的training模式残留,eval()没生效的话某些层会保留训练时两倍的激活值。建议先用torch.profiler跑一下看具体哪层爆的,大概率是attention的中间变量没释放。还有个土办法,把输入切成更小的tensor分段过,能缓解但治标不治本。
试试关掉梯度保存输入,把input也detach下,或者检查下有没有不小心把优化器状态也load进去了。
说实话训练4G推理10G这个幅度不太正常,我怀疑你推理脚本里可能不小心把梯度也打开了,比如调用了model.train()或者某些层有requires_grad没关。另外,你试试把输入直接放在with torch.inference_mode():下面,这个比no_grad更严格,能省不少显存。还有个坑,就是加载模型后别忘了model = model.cuda(),但如果你用了torch.load没指定map_location,可能默认把优化器状态也加载进来了,那玩意儿挺占显存的。如果还不行,建议用torch.profiler看下具体哪层分配的内存,有时候是attention的缓存没释放。
bert为啥要开grad?八成是优化器或者loss里把requires_grad又打开了,试试冻结参数再跑一遍。
训练时4G推理反而10G,这明显不正常,大概率不是模型本身的问题。你试试把优化器状态和梯度清掉再加载state_dict,有时候残留的计算图会偷偷占显存,另外检查下是不是在推理时无意间传了梯度相关的参数进去。
我之前也遇到过类似情况,最后发现是DataLoader的num_workers设太高,每个worker都复制了一份模型副本。你可以先用最简单的单条输入直接跑一遍forward,排除数据管道干扰,看显存曲线是不是还那么夸张。
还有个容易忽略的点,如果模型里有dropout或batchnorm,eval模式下行为会变,但不会导致显存暴涨。建议用torch.profiler跟一下内存分配,重点看是不是某些中间变量没释放,比如attention矩阵。
训练时梯度会释放,推理反而暴涨大概率是优化器状态或中间变量没清,试试trace脚本或加torch.cuda.synchronize()看哪步峰值。
你试试把输入改成固定长度再测,是不是动态shape导致缓存没复用?
试试把优化器状态也清掉,推理脚本里只加载model weights别load整个checkpoint,优化器动量那部分占的显存有时候比模型本身还大。另外确认下是不是在with torch.no_grad()外面调了model,或者输入没加.to(device)导致数据在CPU和GPU间反复拷贝,这种隐式transfer也会爆显存。还有个笨办法,推理前先跑个dummy input做warmup,看显存曲线是不是稳定在某个值,如果持续涨可能是有层在动态创建图。
我猜大概率是保存和加载的state_dict里混进了优化器或ema的影子参数,加载时全部塞进模型导致显存翻倍。你可以先print一下加载后的model显存占用,把加载前后的差值算出来看看。另外推理脚本里如果用了gradient checkpointing或者中间缓存了activations,也会有这种暴涨,检查下有没有意外开启。
这种推理比训练还吃显存的情况我碰到过,大概率不是模型本身的问题,而是数据在GPU上没及时释放。你试试把输入tensor显式地del掉,然后在每个batch结束加一步torch.cuda.synchronize(),有时候异步执行会让显存碎片越积越多。
另外检查下是不是加载state_dict时不小心把优化器状态也load进去了,或者模型里有没有dropout层在推理时还保持训练模式。我之前就是漏了model.eval()但no_grad没生效,结果模型里的buffer一直在累积计算图。
还有个容易踩的坑是tokenizer的padding策略,如果动态padding没设置好,单条文本也可能被填充到最长序列,显存占用直接翻倍。你可以在推理脚本里打印一下input_ids.shape确认下。
如果还不行,试试用torch.jit.script或者onnx导出再跑,有时候能省掉不少中间变量。总之先定位是哪个环节涨的,用torch.cuda.max_memory_allocated()对比峰值和当前值,就能看出是不是有泄漏。
我之前也踩过类似的坑,多半不是推理本身的问题,而是加载完state_dict之后,optimizer或者训练时用的中间变量还被引用着,没释放干净。你可以试试在加载完模型后,把optimizer和scheduler显式del掉,再调gc.collect(),有时候会比empty_cache管用。另外检查下是不是用了torch.jit.trace或者模型里有动态图分支,这也会导致显存峰值暴涨。实在不行就开个新的子进程跑推理,训练完直接杀掉,最干净。
这个情况我踩过类似的坑,大概率不是PyTorch推理本身的问题。你训练时4G是包含了前向和反向的梯度,但推理时如果用了torch.no_grad()还飙到10G,那基本可以排除梯度累积的可能。我猜你八成是把整个BERT模型连同embedding层一起加载到GPU了,但推理时可能不小心把输入文本做了tokenization之后又转成了torch.long以外的类型,比如float32,这样显存会翻好几倍。另一个常见原因是你在脚本里可能没关掉梯度缓存,比如用了model.zero_grad()或者某些层有requires_grad=True的buffer,虽然no_grad能挡autograd但不会释放buffer。你可以试试用torch.cuda.reset_peak_memory_stats()和torch.cuda.max_memory_allocated()打点,看看峰值到底出现在哪个阶段。另外强烈建议用torch.inference_mode()替代no_grad(),它会彻底禁用梯度跟踪,还能省掉一些中间变量的创建。最后如果还不行,检查一下是不是加载了优化器状态或者训练时的batch norm统计量,有时候保存的state_dict里混入了多余参数,加载后即使eval模式也会占用额外显存。我上次就是不小心把model和optimizer的state_dict一起load了,后来只load模型参数瞬间降到2G。
讲真训练才4G推理飙到10G确实不正常,重点查一下是不是把optimizer的state_dict也load进来了,或者模型里有什么buffer被反复注册。我之前遇到过类似情况,是dropout层在eval模式下没生效导致计算图没释放,试试把输入也wrap到with torch.no_grad()里,然后看下nvidia-smi是不是真的没降下来。另外检查下是不是用了half精度但某些层强制float32,混精度偶尔会引发显存异常,你可以先排除下是不是加载了两次模型权重。
你这个情况我遇到过,大概率不是推理本身的问题,而是模型加载或者数据流没弄干净。训练时显存4G是因为梯度占了大头,但推理时如果还保留着优化器状态或者训练时的临时变量,那就等于白省了。建议你检查一下推理脚本里是不是把整个训练模型类给实例化了,有些模块比如dropout或者自定义层里可能缓存了中间结果,eval()不会自动清理这些。另外,试试在加载权重后手动跑一个假输入做warmup,让CUDA把内存布局稳定下来,然后再看真实推理的占用。还有个常见坑是torch.no_grad()只包住了前向,但如果你调用了loss计算或者反向相关的API,显存照样涨。你可以用torch.cuda.reset_peak_memory_stats()和torch.cuda.max_memory_allocated()打点,看具体是哪一行爆的。如果实在排查不出来,试试转成torchscript或者用onnxruntime推理,那玩意儿显存控制得干净多了,我上次就是这么救回来的。
检查下是不是把整个验证集都load进显存了,或者optimizer没释放,eval模式下dropout和BN层也会占一部分。
建议先跑个空模型对比下,排除是模型本身还是数据加载的问题。
八成是优化器状态或中间激活没释放,试试把推理逻辑包在函数里跑,结束后清下缓存。
试试关掉梯度保存中间变量吧,推理时把inputs也detach一下,很多坑都在这里。
检查下是不是加载模型时把optimizer和梯度也带进来了,只load state_dict里模型参数就行。
我之前也踩过类似的坑,最后发现问题基本不在推理代码本身,而是模型里还挂着训练时才需要的缓存和梯度。你用了no_grad确实能关掉梯度,但像dropout、batch norm这些层在eval模式下行为和训练不一样,可如果模型里有什么自定义的buffer或者中间变量没清理,显存就会一直涨。建议你试试在推理循环里每次迭代后把输出变量和中间tensor手动del一下,然后加上torch.cuda.synchronize()再观察显存曲线,看看是不是每步都递增而不是一次性炸掉。另外,BERT分类模型如果保存的是整个模型对象而不是纯state_dict,有时候会把优化器状态和梯度信息也带进去,加载后即使eval也会占额外显存,你确认一下加载方式是torch.load然后model.load_state_dict,还是直接torch.load整个模型?还有个小技巧,可以把输入序列长度截断到最大不超过512,或者用half精度推理,能直接砍掉一半显存。最后别信empty_cache,它只是清缓存池,真正释放得靠释放tensor引用。你试着把推理脚本里所有非必要的变量都放到一个函数作用域里,跑完自动销毁,大概率能解决。
训练时4G推理反而10G,这肯定不正常,batch=1还no_grad按理说应该比训练省显存才对。我怀疑你是把优化器状态或者梯度相关的buffer也一起load进模型了,或者推理时不小心保留了某个大tensor的引用,试试在推理循环里用del显式删掉中间变量再empty_cache。另外检查下是不是加载权重时用了model.load_state_dict(torch.load(..., map_location='cuda')),但保存时带了额外键,或者分词器生成了超长序列,先打印一下input_ids的shape看看。
感觉你这个情况不太像正常的推理开销,BERT base推理一条文本batch_size=1一般也就1-2G撑死了。训练4G推理10G这个反差太离谱了,我怀疑你根本没走到纯推理的逻辑里去。先检查一下是不是加载state_dict之后模型还在梯度模式下,比如你调用了model.train()或者忘了在加载权重之后重新eval,有时候脚本顺序写反了会这样。另外你试试在推理脚本里显式把输入tensor的requires_grad设为False,虽然no_grad应该覆盖了,但有些自定义层或者钩子会偷偷开梯度。还有个常见坑是用了torch.no_grad()但没包住整个前向过程,比如在循环外面设了,结果循环里又调用了别的函数把上下文给重置了,你可以把上下文管理器直接写到forward调用那一行外面。再一个就是看看是不是加载了optimizer的state_dict,有些人有习惯把optimizer也存下来,加载的时候顺手load了,那玩意儿会占不少显存。如果这些都排除了,我建议你用torch.profiler跑一下推理,看显存到底是分配在哪一层,有时候是中间激活值没释放,比如用了很大的序列长度或者attention的缓存没清。最后说个通用技巧,推理时用torch.inference_mode()代替no_grad,这个模式更激进,能省掉一些自动求导的元数据开销。