最近在调一个BERT分类模型,训练完保存了state_dict,然后单独写了个推理脚本加载。奇怪的是,训练时显存占用大概4G左右,但一跑推理,显存直接飙到接近10G,然后OOM。我明明已经用了torch.no_grad(),也调了model.eval(),还试了torch.cuda.empty_cache(),都没什么效果。输入数据也就是一条文本,batch_size=1。是不是我哪里写得不规范?还是说PyTorch推理本来就这么吃显存?求有经验的朋友指点一下排查思路,或者有没有什么通用优化技巧,谢谢!
PyTorch模型推理时显存暴涨,是哪里写错了?求大佬指点
全部回复
共 98 条讲真,你这种情况我遇到过好几次,最典型的坑其实是模型里残差连接或者中间变量没释放干净,尤其是BERT这种深层transformer,反向传播的缓存图和推理时的临时张量完全不是一回事。你虽然加了no_grad,但如果你在forward里手动写了某些操作,比如把多个tensor拼接后存到list里,或者用了python原生循环去迭代层输出,这些中间结果照样会驻留在显存里,跟gradient无关。我建议你先用torch.profiler或者nvidia-smi盯着看,把推理分成几段跑,看看到底是哪一行代码之后显存跳上去的,大概率是某个算子生成了超大临时张量,比如attention的score矩阵,你试试把序列长度砍半对比一下就知道是不是这个原因了。另外,empty_cache只是清空缓存池,不是释放显存,它只影响后续分配的复用率,对峰值占用没卵用。还有个冷门点,你保存state_dict时如果是在torch.no_grad之外用model.state_dict(),那它会把当前所有参数拷贝到CPU,但如果推理脚本加载时用了map_location不一致,或者模型定义里有个大buffer比如位置编码没注册成persistent=False,也会莫名其妙多占一块。最后提醒一句,训练4G推理10G完全不合理,正常同一模型推理应该比训练小一半以上,所以肯定不是正常现象,你查一下是不是推理时误开了梯度,比如把某个模块的requires_grad置True了,或者模型里有BN层在eval下还在更新running stats,那个会额外保存每层统计量。
你这个情况我太熟了,之前调GPT2做生成的时候也撞见过一模一样的,训练4G推理10G,差点以为见鬼了。先说结论,PyTorch推理正常情况下显存是训练的三分之一都不到,你这不是常态,肯定有地方漏了。最大的嫌疑是加载state_dict的时候,模型有没有被完整地放到GPU上,有时候忘了调.to(device)或者输入没同步,会导致某些层在CPU上算完再拷回来,显存直接翻倍。另一个很隐蔽的点是,如果你在模型里写了dropout或别的训练专用层,eval模式是关掉了,但如果你用了torch.no_grad()却仍然在with块里创建了新的张量,比如把中间结果append到list里,那这些张量不会自动释放,会一直累积到脚本结束。我建议你先用torch.cuda.memory_summary()看看到底是哪一步分配的显存,八成会看到缓存分配器那块有个巨大的保留块。还有个土办法,把推理脚本改成每跑一条就打印一次显存,定位到具体是forward还是后处理爆的。如果实在找不到,试试加载模型后先跑一次空输入做warmup,再跑真实数据,有时候是cuDNN第一次启动时建算法的开销。另外,检查一下是不是用了类似transformers库的模型,有些版本在eval模式下还是会缓存hidden_state,需要显式传return_dict=False或output_hidden_states=False。别急着调batch_size,你这条文本单条4G都不到,问题大概率在代码结构上。
试试把优化器状态和梯度清掉再加载,推理脚本里如果还保留了训练时的优化器或者loss计算逻辑,显存会翻倍。另外检查下是不是模型里不小心带了dropout或者BatchNorm的training模式,eval()没生效的话某些层会保留额外缓存。我之前遇到过类似情况,最后发现是没把输入也放到同一个device上,导致CPU和GPU之间反复拷贝,显存碎片化飙升。
排查思路其实很简单,先看下推理时是不是不小心把梯度也带进去了,比如模型输入有没有设置requires_grad=False,或者有没有在no_grad作用域外调用模型。另外,BERT的KV cache如果没手动管理,长文本下显存会翻倍,你试试把max_length限制一下或者用half精度推理,一般能降一半多。我之前也遇到过类似情况,最后发现是误把整个验证集循环里的中间变量都保留下来了,建议检查一下有没有把logits或hidden_state存到list里。
之前也踩过类似的坑,大概率不是推理脚本本身的问题,而是加载state_dict的时候把优化器或者梯度相关的缓存也带进来了。你可以试试在加载完权重后手动调一下optimizer.zero_grad(),或者干脆在推理前把model.zero_grad(set_to_none=True)加上。另外,检查一下是不是数据加载时把整个验证集一次性塞进GPU了,有时候DataLoader的pin_memory或者num_workers设置不对也会导致显存异常。还有个冷门但实用的点,如果用了transformers库,试试把model.config里的output_hidden_states和output_attentions都设成False,默认可能是开的。实在不行就开个torch.cuda.memory_stats()看下是哪块分配暴涨,比瞎猜快。
我之前也踩过这个坑,大概率不是推理代码的问题,而是模型里有个别层(比如attention里的dropout或者LayerNorm)在推理时仍然保留了训练模式的某些计算图,试试看把torch.set_grad_enabled(False)放到加载模型之前,或者干脆用torch.inference_mode(),这个比no_grad更彻底。另外检查下是不是把optimizer的state_dict也一起load了,那个东西有时候会悄悄占显存。如果还不行,可以用torch.profiler看一下具体是哪一层分配的内存,别盲目清缓存,那个基本没用。
你这种情况我怀疑是模型里带了batch_first或者位置编码的缓存没清掉,尤其是BERT系列,有些层在forward时会动态生成mask矩阵,batch_size=1也会占用和序列长度平方相关的内存。建议把输入做一下torch.tensor的连续内存拷贝,或者试下model = torch.compile(model, mode="reduce-overhead"),我这边有时候能省一半。另外确认下是不是用了apex的混合精度,那个在推理时偶尔会重复分配显存。
说个反直觉的,你训练时4G可能是因为反向传播把中间变量及时释放了,但推理时如果开了torch.no_grad(),某些算子反而会保留中间结果用于可能的梯度计算,虽然理论不该这样。
检查下优化器或loss是否还被引用着没释放,推理时把optimizer置空试试。
八成是优化器状态或者中间激活没释放,试试推理前把optimizer置空再清下缓存。
看到你这个情况我第一反应是推理脚本里可能没关梯度,但你说已经用了no_grad,那就要看是不是模型里还有没冻结的buffer在反向传播,比如BN层的running_mean这些,在eval模式下按理说不会更新,但如果你显存是逐层涨上去的,可以试试把模型每个层的输出显存打一下,定位到具体哪一层爆的。另外还有个常见坑,就是加载state_dict时如果用了strict=False,可能有些层没加载上,导致某些参数还是requires_grad=True,推理时虽然不反向,但PyTorch为了保险还是会为这些参数保留梯度空间。还有一个我踩过的雷,就是推理时如果用了torch.utils.checkpoint或者gradient_checkpointing,它默认会缓存中间激活值,即使no_grad也不会自动释放,得手动把模型里的checkpoint禁用掉。你可以先排查一下是不是数据加载时把整个数据集都放GPU上了,有时候DataLoader的pin_memory或者num_workers设得不对也会造成显存异常。最后实在不行,试试torch.inference_mode()替代no_grad,这个更彻底,连autograd的图都不构建,能省不少开销。
你这情况我遇到过,大概率不是PyTorch推理本身的问题,而是模型里某个算子(比如attention mask或者长序列的中间变量)在推理时没走对分支,导致显存峰值比训练还高。可以先试试在推理脚本里把torch.cuda.empty_cache()放到加载模型之后、跑数据之前,再给输入加上torch.no_grad()包裹整个前向,排除是缓存碎片导致的。另外检查下有没有不小心把requires_grad留在某个参数上,或者用了model.train()里的dropout还在生效。如果还不行,就开个torch.profiler看下具体是哪一层分配了显存,一般能直接定位到问题。
这个情况我也踩过坑,多半不是推理本身的问题,而是你加载完state_dict之后,优化器或者训练时的缓存变量还留在显存里没释放。试试在加载模型后把optimizer显式删掉,再调一下torch.cuda.synchronize()看看。另外检查下是不是模型里有用到dropout之外但训练时才会创建的张量,比如一些buffer或临时变量,推理时没清干净。我之前是发现推理脚本里不小心把整个训练好的模型又包了一层DataParallel,结果显存直接翻倍。
遇到这种训练正常推理爆显存的情况,我第一反应是查一下是不是把整个模型复制到了多个GPU上,或者有显存碎片化的问题。你试过torch.cuda.empty_cache()没用很正常,它只是释放缓存块,不代表进程占用会还给系统,真正要看的其实是峰值占用。我建议你直接在推理脚本里加个torch.cuda.max_memory_allocated(),分步打印出哪一行开始显存暴涨,这样能定位到是模型前向、还是数据加载、或者是某个操作导致的。另外,BERT推理时如果没关掉梯度,哪怕有no_grad,某些自定义层里的buffer或者中间变量也可能被保留,你可以检查一下模型里有没有用requires_grad=True的buffer。还有个容易忽略的点,就是推理时如果用了model(x)而不是model.forward(x),并且模型内部有dropout之外的动态图操作,可能触发autograd记录,试试把输入也detach()一下。如果实在排查不出来,可以试试用torch.jit.trace把模型脚本化,有时候能规避掉一些奇怪的显存问题。最后,你训练用的是混合精度吗?如果训练时用了AMP但推理时没设置,某些层会默认用FP32,显存翻倍也不是没可能。
这情况我也踩过坑,大概率不是模型本身的问题,而是推理脚本里没关梯度或者缓存了中间变量。你可以查查是不是用了类似output = model(input)之后还取了output.requires_grad或者把loss的梯度图给带进来了,试试在no_grad块里把输入也.detach()一下。另外如果用了transformers库,检查下有没有把return_dict设成False,有时候它会额外保留一些激活值。实在不行可以开个profiler看看是哪一层爆的,比瞎猜快多了。
有个点你大概率忽略了:训练时显存是动态分配的,但推理时如果模型里带了dropout或者某些层在eval模式下没完全关闭,可能还是有额外开销。另外检查下是不是加载完state_dict后忘了调用model.cuda(),或者输入数据没放到GPU上,导致CPU和GPU之间反复拷贝,这也会让显存异常膨胀。还有一个常见坑是模型里如果有循环或者重复调用了某些子模块,推理时梯度虽然关了,但中间变量没释放,试试用torch.jit.script或者trace一下模型,看能不能压下去。实在不行就开着nvidia-smi盯着,逐行打印每层输出尺寸,定位到具体哪一层爆的。
看你这描述,我盲猜八成是推理脚本里把优化器或者梯度相关的状态也load进来了,或者模型里还挂着训练时才需要的buffer。你只保存了state_dict的话,按理说不该这样,但可以检查下是不是在eval模式下还跑了一些会构建动态图的op,比如某些自定义的forward里用了python控制流或者对tensor做了in-place修改。另外,torch.no_grad()只关梯度,不关显存缓存,如果模型里有大矩阵乘法或者attention的中间结果没释放,batch_size=1也可能爆——BERT自己也就几百MB,你试下加载后先跑一次空的forward,看显存基线是多少。还有个小技巧,把torch.cuda.set_per_process_memory_fraction设小一点,强制它别把显存吃满,这样至少能看报错。我之前遇到过一次,是torch.utils.checkpoint的细节没处理好,推理时反而把中间激活全存了,你可以查查代码里有没有类似的东西。实在不行就用torch.profiler跑一遍,看哪一行分配显存最多,比瞎猜效率高。
训练时4G推理却飙到10G,大概率不是模型本身的问题,而是推理脚本里把梯度也带进去了。检查下输入有没有requires_grad=True,或者模型参数有没有被意外设置成可训练,有时候model.eval()并不会自动关掉所有层的梯度计算。另外可以试试用torch.inference_mode()替代no_grad(),这个能彻底禁用autograd,省掉不少缓存。再不行就在推理前手动清一下CUDA缓存,但别指望它解决根本问题,关键还是看有没有多余的中间变量被保留下来。
训练时4G推理却飙到10G,这肯定不正常,batch=1还OOM大概率不是显存不够而是内存碎片化或者缓存未释放。你可以试试在推理循环里加上torch.cuda.synchronize()看是不是异步执行导致的峰值,另外检查一下是不是模型里带了dropout或BN层在eval模式下没关干净。还有个容易踩的坑是,如果你用transformers库,记得把return_dict=False或者显式清理一下past_key_values,有时候这些中间变量会一直驻留显存。之前我遇到过类似问题,最后发现是dataloader的num_workers开太多,每个worker都复制了一份模型参数,推理时全挤到显存里了。
我之前也踩过类似的坑,训练正常但推理爆显存,最后发现是模型里有个dropout或者batchnorm的buffer没冻结,虽然你调了eval,但某些自定义层可能没生效,建议检查一下模型结构里所有forward里创建的临时变量是不是都赋给了self。另外,你试试用torch.jit.script或者直接把模型转成onnx跑推理,有时候PyTorch的autograd图虽然被no_grad包住了,但某些算子还是会缓存中间结果,尤其是BERT里的attention mask和position id,如果每次动态生成也可能累积显存碎片。还有个很隐蔽的点,就是推理时如果用了torch.cuda.synchronize或者显式调用了.cuda(),可能会触发额外的内存分配,你可以用torch.cuda.memory_summary()打印一下看具体是哪一层在涨。我上次是发现有个Tensor被重复创建了,比如在循环里把输入又expand了一下,导致临时变量没被释放,你可以试着把输入先固定成CPU上的tensor,然后一步步打印每层的显存变化来定位。另外,如果实在找不到原因,直接开AMP混合精度推理,显存能降一半还多,不过要小心某些层不支持fp16。总之你这个情况肯定不是正常现象,BERT-base推理batch_size=1撑死2G,你查一下是不是加载了优化器状态或者梯度图被意外保留了。