最近在做一个多工具调用的Agent项目,用PyTorch 2.1跑Qwen-2.5-7B。推理时我把历史对话、工具返回结果都拼进上下文,用generate()流式输出。现在发现每跑完一个工具调用,显存就涨几百MB,连续跑十几个工具后直接OOM。我已经试过torch.cuda.empty_cache(),也把梯度关了(with torch.no_grad()),甚至把历史序列截断到4096,但占用还是只增不减。奇怪的是,单次推理后显存会回落一点,但基线就是比刚开始高。有没有大佬遇到过类似情况?是KV cache没释放,还是streamer对象持有张量?或者单纯就是PyTorch的显存碎片化,需要定时重置torch.cuda.memory._record_memory_history()?希望有实战经验的朋友指点一下,谢谢。
PyTorch写Agent循环时显存越跑越高,是代码问题还是框架特性?
全部回复
共 7 条大概率是cache对象没清干净,试试每个循环重建model和tokenizer,或手动释放旧streamer。
建议给generate()传past_key_values=None强制重置,不然基线肯定一路涨。
多半是generate里cache_updated没清干净,试试手动重置past_key_values,顺便检查下streamer是不是存了整段输出。
PyTorch的缓存分配器确实会保留显存块,empty_cache()只是把空闲块还给缓存池,不会真正释放给系统。你试试在循环外先warmup一次,然后监控torch.cuda.memory_reserved()和memory_allocated(),如果reserved持续上涨那就是分配器碎片化问题,可以定期用torch.cuda.memory_allocated对比看看。另外streamer如果持有上个step的logits张量,也会造成引用计数不归零,建议在generate后显式del掉streamer对象再gc.collect()。我之前遇到过类似情况,最后发现是history列表里存了太多张量,改成只存token id就稳定了。
八成是streamer没清干净,试下结束流式后del掉对象再empty_cache,基线应该能降回来。
八成是cache_utils里没清干净,试试generation_config里把use_cache设False,或者手动重置一下streamer的token缓存。
八成是streamer或cache对象没清干净,我之前也这样,手动del加gc.collect能好点。
我之前也踩过类似的坑,最后发现多半不是PyTorch的锅,而是HuggingFace的generate内部在维护past_key_values时没释放干净。你试试把每次流式输出后的streamer对象显式del掉,再配合empty_cache,有时候streamer里会缓存logits或hidden state。另外你截断到4096但KV cache是按层数乘以头数动态分配的,如果模型内部没重新初始化past_key_values,旧的长度可能还占着块,建议每次工具调用后重新build一个新的model instance或者用model.reset()之类的方法。还有个隐蔽点是Qwen的tokenizer在拼接工具返回时可能会生成特殊张量,比如tool_call_ids,这些在batch维度上不会自动清理。我后来直接用paged attention的库比如vLLM跑agent循环,显存控制就稳多了,但要牺牲一点灵活性。你也可以看看是不是数据加载器里保留了上一轮的input_ids,那个引用也会让显存基线抬高。总之先排除代码里的引用泄漏,再怀疑框架特性,毕竟PyTorch本身不会无故吞显存。