最近在部署一个BERT分类模型,单个样本推理时显存占用大概2G,但连续跑几百个样本后,显存直接飙到10G+,最后直接OOM了。我试了torch.no_grad(),也调用了del和torch.cuda.empty_cache(),但好像效果不大。代码里主要用DataLoader分批加载,每批32条,推理完把结果append到列表里。想问下这种情况一般是哪里没清理干净?是不是需要把每个batch的输入和输出都手动清空?还是模型本身有动态图缓存?另外,有没有什么工具能实时监控每个张量的引用计数,方便定位问题?求有经验的大佬指点一下。
PyTorch模型在推理时显存一直涨,是哪里没释放?
全部回复
共 167 条这种情况我遇到过,大概率不是pytorch本身没释放,而是DataLoader的num_workers在作祟,多进程加载数据时每个worker都会拷贝一份模型副本,显存自然就炸了。你可以试试把num_workers设为0或者1,看看是不是瞬间稳定。另外append列表如果存的是tensor或者detach后的结果,也会累积计算图,最好转成cpu的numpy再存。检测工具的话,pytorch自带的torch.cuda.memory_summary()能看个大概,但更细的引用计数可以用pytorch_memlab或者gc模块配合debug。
试试用torch.cuda.synchronize()强制同步一下,或者看下DataLoader的num_workers是不是设多了。
你这情况我遇过类似的,大概率不是没del变量,而是DataLoader的worker进程和PyTorch的缓存分配器在搞鬼。试试把DataLoader的num_workers设为0或者pin_memory=False,看显存还涨不涨。另外可以装个pytorch_memlab或者nvidia-ml-py3,在推理循环里打log看每一步的显存变化,比空猜靠谱。还有个小技巧,每次迭代完显式调用torch.cuda.synchronize()再清缓存,有时能解决异步操作残留的问题。
这种情况我遇到过,多半不是单个张量没释放,而是DataLoader的num_workers在推理时也会累积缓存,尤其是每次迭代创建新变量没及时回收。你可以试试把batch里的inputs和outputs显式赋None,或者在每个batch后调用torch.cuda.synchronize()看看有没有改善。监控引用计数的话,用gc.get_objects()加上torch.is_tensor()筛一下,能快速定位到哪些多余张量还活着。另外如果用了transformers库,记得把model.eval()和no_grad()放在循环外面,不然每次调用也会重新构建计算图。
这种问题我也踩过坑,主要是DataLoader在遍历时,每个batch的输入张量如果没有显式释放,计算图虽然被no_grad截断了,但张量本身还占着显存。建议你在循环里每个batch结束后手动把inputs、labels设成None,再调一下torch.cuda.empty_cache(),实测能压住内存增长。监控工具的话可以用nvidia-smi配合pytorch的memory_summary(),或者试下pytorch_memlab这个库,能直接看每行代码的显存分配。另外检查下是不是把结果列表存在了GPU上,append会导致列表里的张量一直持有引用。
这种情况我遇到过,多半不是单纯没释放,而是DataLoader的pin_memory或者梯度缓存的问题。你可以在每个batch推理后加个torch.cuda.synchronize()试试,有时候显存统计会滞后。另外关注下模型内部的dropout或者BatchNorm层,推理模式切换后某些缓存没清干净也会累积。监控的话我习惯用pytorch的memory_summary(),能看到每个张量的分配情况,比看引用计数直观。
说实话这个问题挺典型的,我之前用BERT做序列标注也被坑过。你用了torch.no_grad()和手动清缓存但效果不大,那大概率不是计算图缓存的问题,而是DataLoader和列表append这块在悄悄积累东西。每批32条推理后,你把结果append到一个列表里,如果列表本身是Python对象,但里面存的是GPU上的tensor,那这些tensor的引用计数一直没归零,显存就不会释放。我建议你推理完立刻把结果从GPU搬到CPU上,比如用.cpu().numpy()或者直接转成Python列表再存,这样GPU上的tensor就能自动回收了。
另外,你说的“动态图缓存”其实在torch.no_grad()下基本不会累积,但可以检查一下模型里有没有用了requires_grad=True的缓存变量,比如一些中间层的embedding。还有一个容易忽略的地方:DataLoader的num_workers如果大于0,子进程可能会持有一些显存不释放,建议先设成0测试一下。至于监控张量引用计数的工具,可以用torch.cuda.memory_summary()看当前显存分配情况,或者试试pytorch的memory profiler,能显示每个张量的生命周期。我之前用pytorch_memlab这个库,可以实时打印每个位置的变量引用数,找泄露挺方便的。总之核心思路就是确保每个batch的输入输出都不在GPU上停留,并且不要有全局变量一直引用着它们。
这种情况我遇到过类似问题,大概率不是显存没释放,而是计算图没清干净,尤其是把中间变量append到列表里时,列表会一直持有张量的引用,导致整个计算图无法释放。建议你在每个batch推理完后,手动把输入和输出都detach()一下再存,或者用torch.cuda.synchronize()强制同步后再清缓存。另外可以试试pytorch的torch.cuda.memory_summary()看下哪部分占用高,比单纯看显存占用更直观。
试试在每轮batch后把loss和变量都设成None,光靠del有时确实清不干净。
这种情况大概率是DataLoader的pin_memory或者num_workers没处理好,导致内存堆积在CPU侧没有及时释放,建议先把这两个参数关掉试试。另外你每次append列表时,结果如果是tensor且带了梯度,即使no_grad也会保留计算图,推理完记得用.detach().cpu()转成numpy再存。追踪引用计数的话可以用pytorch的torch.cuda.memory_summary(),能看到每个CUDA分配点的详细情况,比手动排查快很多。
试试把梯度全关了,用torch.inference_mode()替代no_grad,另外检查下DataLoader的num_workers是不是设太高了。
试试在append结果时用detach(),列表里存张量会累积计算图,显存当然炸了。
试试在推理循环里加个torch.cuda.synchronize(),可能梯度缓存没清干净。
试试把梯度清零和中间变量显式删除,可能DataLoader的pin_memory也有影响。
试试把每批结果从GPU移到CPU再append,可能是列表里累积了太多GPU上的梯度计算图。
这种情况我也遇到过,问题大概率出在DataLoader的num_workers或者梯度缓存上。虽然推理时用了no_grad,但模型如果开启了训练模式,或者某些层(比如BatchNorm)在推理时仍会维护状态,显存就一点一点堆上去了。可以试试把model.eval()加上,然后检查下DataLoader的pin_memory是不是开了,关掉有时能缓解。至于监控工具,可以用torch.cuda.memory_summary()看分配细节,或者用nvidia-smi配合pytorch的memory_snapshot来抓快照。另外结果append到列表这个操作本身不占显存,但别忘了在循环里显式把每个batch的输入输出变量设成None。
试试把梯度显存优化关掉,或者推理时加个with torch.inference_mode(),比no_grad更彻底。
试试把每个batch的loss或输出detach()一下,再手动清空计算图,可能是梯度积累导致的显存泄漏。
这种情况我之前也踩过坑,大概率是DataLoader加载时没有关闭pin_memory,或者模型里某些层(比如LayerNorm)的缓存没清理。你可以试试在推理循环里显式调用torch.cuda.synchronize(),然后把每个batch的input/output用del删掉再清缓存,但别太依赖empty_cache,它有时候是假释放。监控引用计数的话可以用objgraph库,但我觉得更直接的是跑一个batch就打印一下torch.cuda.memory_summary()看看哪部分在涨。
这种显存持续上涨的问题我也踩过类似的坑,你试的那些方法确实常见但往往不彻底。我个人感觉核心可能出在DataLoader的num_workers上,如果开了多进程,每个worker都会缓存一些张量副本,即使你在主进程里del了,worker的显存依然没释放,可以试试把num_workers设为0或者1看看有没有改善。另外你提到结果append到列表,那个列表本身如果拼接成一个大tensor或者一直保留在内存里,也会阻止显存回收,因为列表里的张量可能还持有计算图的钩子,可以考虑把每个结果立刻转成numpy或者直接写入磁盘。关于工具,pytorch自带的torch.cuda.memory_summary()能打印当前显存分配明细,还能看到哪些张量没有被释放,比单纯看占用更精确。还有个办法是给每个batch的输入输出加个作用域,比如用with torch.no_grad()包裹整个推理循环,同时每次迭代后手动把batch数据设为None,配合gc.collect()强制回收,我试过这个组合对减少碎片化挺有效的。模型动态图缓存确实也可能有影响,特别是用了transformers库的模型,有时会缓存中间激活值,可以检查一下有没有开启model.eval()或者设置return_dict=False。你那个OOM大概在多少样本之后出现?如果前半段涨得慢后半段突然飙升,很可能跟DataLoader的预加载机制有关。