最近在做一个简单的ReAct Agent,基于PyTorch和HuggingFace的LLM,每次调用模型生成回复后,我把历史对话拼到prompt里继续下一轮。但跑了十几轮之后,显存占用直线上升,从4G涨到12G,最后直接OOM。我试过清空缓存(torch.cuda.empty_cache())和调低max_new_tokens,都没什么效果。是不是每次拼接prompt时,历史输出被重复计算了梯度?还是说推理时也需要detach?我看一些Agent框架(比如LangChain)好像没提这个问题,是我写法有问题吗?希望大佬指点一下,最好能贴一段简单的处理逻辑,谢谢!
用PyTorch写Agent时,多轮对话的显存越堆越高怎么处理?
全部回复
共 174 条你这问题我也遇到过,核心确实是每次拼接prompt时历史输出没做detach,导致计算图越堆越长,显存就炸了。推理时记得在生成后对tensor调用.detach()或者用torch.no_grad()包一下,另外也可以每轮只保留最近几轮对话,别把全部历史都塞进去,这样能省不少显存。
这个问题我踩过一模一样的坑,核心原因就是历史tensor没有被释放。PyTorch推理时默认会保留计算图用于可能的反向传播,哪怕你只是做生成,每次拼接prompt后,之前的对话tensor和attention的key/value缓存都还在显存里累积着。你试的empty_cache只清空闲散内存,但没法释放还被变量引用的显存块。
最直接的解法是在每轮推理前用torch.no_grad()包裹,并且在模型生成后手动把输入tensor和past_key_values置为None,然后调用del显式删除变量。像这样:with torch.no_grad(): outputs = model.generate(...) 之后把input_ids、attention_mask这些全删掉。另外如果你用的是HuggingFace的pipeline,可以设置clean_up_tokenization_spaces=True并且每次重新初始化pipeline对象。
还有一个容易被忽略的点:如果prompt拼接时用了列表或元组来存储历史记录,那这些容器本身也会持有tensor引用。建议每轮只保留纯文本字符串,等到下一轮再重新tokenize,这样旧tensor就能被GC回收掉。LangChain其实内部做了类似的事情,它默认不会在内存里保留所有历史tensor,而是每次都重新编码文本。
我自己写了个简单的循环,每次调用前先清一下上下文管理器:model.eval() + torch.inference_mode() + 手动释放,基本能把显存控制在单轮推理的1.5倍以内。你可以试试把历史对话存成字符串列表,下一轮开始前用tokenizer重新编码整段文本,别用累积的tensor拼接。
我也遇到过这个问题,其实是每次拼接prompt时,没把历史输出的梯度上下文断开。用torch.no_grad()包裹推理部分,或者对生成结果调用.detach()就能解决。另外建议别把整段历史都拼进去,可以只保留最近几轮对话,显存涨太多的时候直接裁剪一下历史列表。
大概率是历史输出没detach导致计算图累积,推理时用torch.no_grad()包一下就行。
这个问题我上周刚踩过坑,核心原因就是历史拼接时没有detach。PyTorch默认会构建计算图,每次把之前输出的token拼进新prompt,梯度就会一直积累,显存自然越堆越高。你可以在每次生成后手动把tensor detach,或者直接用with torch.no_grad()包裹推理段,再配合torch.cuda.empty_cache()清一下碎片,基本就能稳住。LangChain其实内部做了类似处理,只是没在文档里强调。
这个问题我踩过一样的坑,核心就是每次拼接prompt时,之前生成的token的key/value cache没有释放,而且梯度默认还在计算图里。你推理时加个torch.no_grad(),然后每次生成完手动把past_key_values置为None,或者直接用model.generate()并设置use_cache=True,这样显存就不会持续堆积了。我之前写了个简单循环,每轮用tokenizer把新对话拼进去后,记得把input_ids整体传给模型,别用变量累加的方式。
这问题我碰到过一模一样的情况,核心就是历史对话的token embeddings会一直留在计算图里。你推理时把整段历史prompt重新过一遍模型,梯度虽然不更新但中间变量没释放,显存自然越堆越高。解决思路是在每轮生成完手动把历史输出的requires_grad设为False,或者直接对历史输出调用detach()断开连接。我自己的做法是每次拼接prompt前,把之前生成的token ids存到list里,重新构造输入时只保留文本内容,不保留任何张量引用。另外建议把torch.no_grad()包在整个推理循环外面,能省不少显存。
试试把history里的输出用torch.no_grad()包一下,推理时不需要跟踪梯度,能省不少显存。
这问题我踩过一样的坑,核心确实是推理时没detach导致计算图不断累积。你每次拼接历史输出时,tokenizer返回的input_ids会带着之前轮次的梯度信息,即便只是推理,PyTorch默认也会保留。解决方案很简单:在每轮生成前对输入做with torch.no_grad():,或者手动把input_ids.detach()一下。另外建议用KV cache管理历史,比如每次只保留最近几轮的关键信息,别把全部历史拼进去,既能省显存又能防止长上下文性能下降。
你这问题我太有同感了,之前也踩过同样的坑。关键其实不在缓存,而是你的历史对话拼接后,整个prompt的token数在不停增长,而PyTorch默认会为所有中间变量保留计算图用于反向传播,即便你只是在做推理。一个比较直接的解法是在模型调用前加with torch.no_grad():,这样就不会追踪梯度,显存占用能明显下降。另外,建议每次生成完回复后,把历史对话中不需要梯度的部分显式detach一下,比如把input_ids和attention_mask从计算图中分离出来再拼接到下一轮。还有个容易被忽略的点是,HuggingFace的model.generate()内部默认会缓存past_key_values,多轮对话下这些缓存也会越积越多,像GPT-2、Llama这类模型你可以手动调用model.clean_cache()或者设置use_cache=False来清掉历史key-value缓存。如果不想每次都丢缓存,也可以按固定窗口截断历史,只保留最近几轮对话,这样显存就稳住了。LangChain其实内部有自动管理历史token的机制,只是没明说,你可以参考下它的ConversationBufferWindowMemory实现。
试试把历史对话的梯度关掉,推理时用torch.no_grad()包一下,显存瞬间就降下来了。
你遇到的这个问题很典型,核心原因确实是推理时没有对历史输出做detach。PyTorch默认会构建计算图,即使你只是生成文本,历史token的梯度信息也会保留,导致显存越堆越多。一个简单的做法是在每次生成后,把输出张量detach掉,或者直接用torch.no_grad()包装推理过程。我自己写Agent时习惯在循环开始前定义一个空列表存prompt,每次只把新生成的文本转成字符串拼进去,避免保留整个张量。另外检查一下你是否在eval模式下调用的模型,虽然不影响显存累积,但能省点开销。
你这个问题我遇到过,大概率是历史对话里的token embeddings一直在计算图中挂着没释放。推理时用torch.no_grad()包裹一下,或者每轮对话结束后显式把prompt的grad_fn链断开,比如prompt_tensor = prompt_tensor.detach()。另外可以试试每轮只保留最近几轮对话,别让prompt无限膨胀,不然缓存清空也救不了。
你遇到的这个情况我跑agent时也踩过坑,关键确实在于推理时没有detach。每次拼接prompt时,历史token的gradient会一直累积在计算图里,即便只是推理,PyTorch默认也会保留中间变量用于反向传播。可以试试在model.generate()外面包一个torch.no_grad(),或者显式调用new_tokens.detach(),另外建议每轮对话后手动清一下计算图,比如把past_key_values设为None。我后来参考了LangChain的源码,发现它对每一步的输入都做了clone().detach()处理,这样显存就不会一直涨了。
这个问题我之前也踩过坑,核心原因确实是你猜的那样——每轮拼接历史对话时,之前的输出默认是带着计算图的,哪怕你只是用model.generate推理,它也会把整段prompt的梯度信息保留下来,积少成多显存就炸了。一个简单的做法是在每轮生成前对输入做with torch.no_grad()包裹,或者在构建输入张量时手动detach一下,比如input_ids = input_ids.detach()。另外我习惯每轮结束后主动把中间变量删掉,再调empty_cache,虽然不能根治但能缓解。LangChain其实底层也处理了这些,只是它封装好了你没看到,你可以参考它的对话缓冲区实现,本质上就是只保留tokenized的列表而不保存梯度。还有一个容易被忽略的点:如果你用past_key_values做缓存,那随着对话变长,key-value缓存本身也会占大量显存,可以考虑限制历史轮数或者用滑动窗口截断。我自己的做法是在每轮生成前,把history里超过指定长度的部分截掉,再重新tokenize,这样显存增长就稳定多了。
你这问题大概率是历史token没做截断,试试每次只保留最近几轮对话,再手动清一下计算图。
这问题我踩过一模一样的坑,核心不在detach,而是你每次把完整历史拼进prompt时,所有历史token的KV cache都被重新计算了,而且PyTorch默认会保留计算图里的中间变量,即便你是inference模式,只要没包在torch.no_grad()里,那些中间激活值就会一直占着显存。empty_cache只是释放未使用的缓存块,并不会主动清掉还在引用里的张量。你试试在生成那一步明确加上with torch.no_grad(),然后每次只保留最新的KV cache,用past_key_values传进去,别每次都从头算。还有,如果用的是HuggingFace的generate,记得设use_cache=True,并且把历史输出的embedding直接存下来,别整个token序列都塞回model。LangChain没提是因为它默认帮你做了这些,或者它用的框架自动清理了。最简单的验证方法:每轮结束后打印一下所有tensor的refcount,你会发现历史输出被多个变量引用着,手动del掉旧prompt和output,再调gc.collect(),效果立竿见影。
这问题我踩过一模一样的坑,大概率不是梯度的问题,而是你每次把完整历史拼进prompt时,之前生成的token的KV cache没有释放,又在新一轮里重复计算。试试在每轮推理后手动把past_key_values置空,或者用with torch.no_grad()包住model.generate,这样能切断历史输出的计算图。另外,如果用的是HuggingFace pipeline,可以检查一下是否默认保存了所有中间状态,我之前就是改成每轮只保留当前轮的输出,显存就稳定了。你用的模型是不是没设use_cache=True?这个参数对多轮很关键。
这问题我之前也踩过,核心就是推理模式下你虽然没开grad,但历史token的KV cache会一直累积在显存里,跟梯度没啥关系。empty_cache只释放未使用的缓存块,救不了这个。建议你每轮只保留最近N轮对话截断一下,或者用HuggingFace的past_key_values手动管理,每轮把旧的KV传给下一轮,别每次都重新算全量prompt。我之前写了个简单方案就是固定最大历史长度,超出就把最早的消息丢出去,显存立刻稳住了。
推理模式下记得用torch.no_grad()包住生成,不然历史拼接的图会一直累积梯度,显存当然只涨不降。