最近在做一个简单的ReAct Agent,基于PyTorch和HuggingFace的LLM,每次调用模型生成回复后,我把历史对话拼到prompt里继续下一轮。但跑了十几轮之后,显存占用直线上升,从4G涨到12G,最后直接OOM。我试过清空缓存(torch.cuda.empty_cache())和调低max_new_tokens,都没什么效果。是不是每次拼接prompt时,历史输出被重复计算了梯度?还是说推理时也需要detach?我看一些Agent框架(比如LangChain)好像没提这个问题,是我写法有问题吗?希望大佬指点一下,最好能贴一段简单的处理逻辑,谢谢!
用PyTorch写Agent时,多轮对话的显存越堆越高怎么处理?
全部回复
共 174 条断句保存历史时不带梯度就行,用no_grad包一下再存,显存就不会跟滚雪球似的涨了。
大概率是历史token全在计算图里没释放,推理时记得包torch.no_grad(),再配合past_key_values缓存就行。
多半是历史token全进图算梯度了,推理记得包torch.no_grad(),再不行就手动截断前几轮对话。
大概率不是梯度的问题,推理模式下按理说不会存梯度,但你如果没写torch.no_grad(),模型内部某些缓存可能一直在累积。我遇到过类似情况,把prompt拼接改成只保留最近几轮,或者用KV cache的剪枝,显存会稳很多。另外检查下是不是history列表本身没清理,每次把完整token序列都传进去了,这个最容易被忽略。可以试试把每轮的输出只存文本,下次重新编码,别直接复用旧tensor,显存能降一半。
这问题我碰到过,大概率不是梯度的问题,是你把历史token全塞进输入,KV cache越攒越大,显存自然就爆了。试试每轮对话后把prompt里的旧token截断,只保留最近几轮,或者用滑动窗口。另外推理时记得用torch.no_grad()包起来,虽然不影响KV cache,但能省点中间变量。LangChain没提是因为它默认帮你做了history压缩,你可以看看它的ConversationBufferWindowMemory是怎么实现的。
这问题我之前也踩过,坑不在显存缓存,而在你每次把历史输出拼进prompt时,那些token的gradient其实还挂在计算图里。虽然你是在推理,但只要没显式包torch.no_grad(),PyTorch就会默认保留中间变量用于反向传播,所以显存才会只涨不降。你试试把生成部分整体包进with torch.no_grad():里,然后每轮对话结束后手动把当前轮的input_ids和attention_mask从计算图上detach出来,再拼到历史里。另外,LangChain没提是因为它底层默认用了pipeline或者直接调model.generate(),那个本身就不跟踪梯度,但你自己写循环就容易忽略这个。还有个小技巧,如果对话轮次特别多,建议做个滑动窗口,只保留最近几轮,别无限拼,既省显存又能控制prompt长度。你那个empty_cache其实没啥用,它只释放未使用的缓存块,计算图还占着就没辙。最后检查一下是不是把tokenizer的返回结果直接存了,那个里面也有grad_fn,记得对tensor单独取出来做深拷贝。
这问题我上周刚踩过坑,跟你描述的一模一样,从6G一路飙到爆。核心原因就是你把整段历史对话都塞进prompt,PyTorch默认会从输入到输出构建整张计算图,虽然inference模式下不反传,但中间激活值全被保留了,尤其是attention的KV cache,每轮都在翻倍累积。你只清空显存缓存没用,因为那些张量还在计算图里被引用着。正确做法是在每次生成前用torch.no_grad()包住,并且把prompt里的历史输出统一用detach()切掉,更关键的是要手动管理KV cache,HuggingFace的generate接口里有个use_cache参数,你把它设为False试试,虽然慢点但显存能稳住。另外如果你用的是GPTQ或AWQ量化模型,还得注意past_key_values会以tuple形式存在,得在每轮对话后显式del掉再gc.collect()。LangChain不提这问题是因为它内部默认做了缓存清理,或者他们用的框架自动帮你处理了。我自己最后是写了个循环,每轮结束把past_key_values置空,再调用empty_cache,现在跑五十轮也就稳定在6G左右,你可以试试。
试试把历史输出token拉满缓存,推理时用torch.no_grad()包一下,prompt拼完就释放旧变量。
这问题我遇到过,核心确实不是缓存的事,而是你每次把历史输出拼进prompt时,那些token在底层会带着之前的计算图一起走,显存自然只增不减。你试试在每次生成完把logits或者hidden_state detach掉,或者干脆用with torch.no_grad()包住整轮推理,尤其是只做生成不反向传播的时候。另外我习惯每轮对话后手动把input_ids拷贝到cpu再转回gpu,相当于切断旧的graph,实测显存能稳定住,你可以参考下这个思路。
你八成是推理时没开no_grad,导致每轮都在累积计算图,跟max_new_tokens关系不大。我之前也踩过这坑,后来学乖了,生成完直接把past_key_values置空,或者用model.generate的时候传入use_cache=False(虽然慢点但省显存)。更省事的做法是每轮只保留最近几轮对话,别全量拼接,反正模型也记不住太远的内容。
显存涨这么快大概率是历史tensor没释放,empty_cache只是清碎片,治标不治本。你可以试试把每轮的输出转成纯文本存下来,下一轮再重新tokenize,不要直接复用上一轮的tensor。还有个小技巧,写个循环里定期del掉中间变量再手动gc.collect(),我之前这么搞完基本能控制在8G以内。Lang
这问题我踩过一模一样的坑,你大概率不是prompt拼接的问题,是HuggingFace的model.generate()在内部把整个输入序列都包进了计算图,哪怕你设了torch.no_grad(),只要没在生成前手动把梯度关干净,历史token的中间激活值就会一直挂在显存里。我之前自己写Agent的时候,每轮对话完都得把当前轮的输出和prompt整体重新编码一遍,但真正吃显存的是之前所有轮次的KV cache,PyTorch的缓存机制不会自动帮你释放这部分。你可以试试在每轮生成前显式调用model.eval(),然后包一层with torch.no_grad(),但更关键的是要定期清掉past_key_values——如果你用的是自回归接口,建议手动维护一个固定长度的滑动窗口,超了就直接截断最老的turns,别让历史无限增长。另外torch.cuda.empty_cache()只是把未使用的缓存块还给驱动,不是真的释放已分配的张量,所以你感觉没用是正常的。LangChain没提是因为他们默认用text-generation-inference或者vLLM这类服务端推理,底层有paged attention和KV cache管理,你本地用transformers就得自己处理。我最后是直接改成每轮都重新用tokenizer把完整对话编码成新tensor,然后只保留最后两轮的KV cache,显存立刻稳住了。
试试在每轮拼接前把历史tensor拷贝出来再detach一下,或者直接用list存文本别堆张量,应该能解决。
这问题我踩过一模一样的坑,核心真不在empty_cache,而是你每次把整段历史拼进prompt时,之前生成的那些token的hidden state其实还挂在计算图里。虽然你用的是推理模式,但HuggingFace的generate默认会保留中间变量用于beam search之类的操作,多轮叠加下来就是显存黑洞。我试过最有效的办法是每轮对话结束后,把当前prompt和生成的回复整体做一个detach的副本存下来,下一轮直接用这个纯tensor去拼新的输入,别让旧的历史参与任何梯度计算。或者更粗暴点,用torch.no_grad()包住整个生成过程,再把生成的ids转成cpu上的list存历史,每轮重新tokenize,虽然慢点但显存稳如老狗。另外你提到的LangChain其实内部做了类似处理,他们不会把整个历史都塞进同一个连续tensor,而是用不同的message对象分开存,这样每次只对新增部分做forward。你可以试试把历史对话拆成独立的分块,每轮只对最后一段做encode,前面的部分直接用缓存下来的key-value,这样能省一大截。如果不想改架构,至少把max_new_tokens调小,同时设置use_cache=True,并且每轮结束后手动把optimizer.zero_grad()和gradient checkpointing打开,虽然推理时不用optimizer但这么干能强制释放一些中间buffer。还有个小细节,检查一下是不是你把token_type_ids或者attention_mask也拼进历史了,那个东西累积起来也占内存。最后实在不行就定期截断历史,比如只保留最近三轮,反正ReAct的推理一般也用不到太久远的上下文。
老哥试试每轮只保留最近几轮对话,或者用KV cache来存历史状态,不然prompt越来越长显存肯定爆。
跑完一轮记得把输入输出都detach一下,虽然推理不存梯度,但历史tensor还在图上堆着,手动释放掉能省不少。
大概率不是梯度的问题,你推理时本来就没开grad,多半是kv cache在作祟。每次拼接历史后重新过一遍全量prompt,之前生成的key/value都会留在显存里,不释放自然越堆越高。
你可以试试把历史对话的token截断,或者用HuggingFace的past_key_values手动管理,每轮只传新增部分。另外LangChain其实也遇到过这问题,只是它默认帮你做了缓存清理,你直接裸调模型才暴露出来。
简单处理的话,每轮生成完把outputs的past_key_values丢掉,然后强制设model.eval(),再配合torch.cuda.empty_cache(),应该能缓解不少。如果还不行,就考虑换用vLLM或TGI这类推理优化框架,它们对多轮显存管理做得更省心。
试试把历史对话截断或者做embedding缓存,gradient别留,inference时用torch.no_grad()包一下应该能压住。
这问题大概率不是梯度的事,推理模式下你只要确保model.eval()并且没开grad,prompt拼接本身不会累积梯度。真正吃显存的是你每轮都把完整历史token重新过一遍模型,KV cache又没释放,建议自己维护一个定长窗口,只保留最近几轮对话,或者用HuggingFace的cache工具手动清理一下past_key_values。我之前也遇到过,后来改成只存最近5轮加个摘要,显存就稳住了。
多半是梯度图没断,推理时记得用torch.no_grad()包一下,再清下计算图缓存就行了。
这问题我也踩过坑,多半不是梯度的问题,而是你把整段历史都塞进prompt,token长度线性涨,KV cache也跟着涨,显存自然就爆了。试试每轮只保留最近几轮对话,或者用HuggingFace的pipeline时把past_key_values传下去,这样能复用之前的计算。我写Agent时是手动维护一个定长的对话队列,超了就把最老的pop掉,显存基本稳定。另外推理模式记得加torch.no_grad(),虽然不解决根本问题,但能省点内存。
试试把历史对话截断到最近几轮,再给输入加个no_grad包一下,我这么改完显存稳多了。
这问题我当初搞Agent的时候也撞过,后来查了半天才发现跟梯度没半毛钱关系,纯是缓存没释放的问题。你每次把历史输出拼进prompt,模型在生成新token时会对整段输入做attention计算,那些中间激活值全留在显存里了,empty_cache只清碎片不清这些。关键是你推理时得用torch.no_grad()包起来,然后每轮结束后把logits和past_key_values手动置None,再调一下cache_allocation_config之类的东西。我之前试过最有效的办法是每轮调用完直接把model的输入tensor和输出tensor都del掉,再empty_cache,能压回3G左右。另外LangChain不提是因为它默认用HF的pipeline,内部已经帮你管理了这些,你直接调model.generate()反而容易漏。你要是想省事,可以直接把历史对话截断到最近几轮,别无限拼,显存增长就变成线性的了。