推理时显存不够,核心思路不是死磕显存容量,而是把“推理痕迹重计算”这个隐藏的显存杀手单独拎出来优化当你把KV Cache和中间激活从“保存”变成“按需重算”,同样的显卡能多塞几倍的模型上下文,这才是大模型推理显存节省的关键路径。
推理痕迹重计算到底是什么在吃掉你的显存
很多朋友在部署大模型时,盯着参数量买显卡,结果一跑起来就OOM,这时候你打开nvidia-smi,会发现显存占用远高于模型权重文件的大小,多出来的部分,就是推理痕迹(intermediate states)和KV Cache。推理痕迹重计算这个动作,本质上是用时间换空间,但具体怎么换、换哪部分,很多人没搞明白。
从一次完整的推理过程看显存流向
假设你让模型生成“今天天气如何”这句话,模型需要逐token预测,在生成第1个token时,Transformer每一层都要计算Q、K、V矩阵,经过多头注意力、前馈网络、LayerNorm,每一层都会产生临时的中间激活值(activation),这些激活值如果全存下来,显存占用会随着batch size和序列长度暴涨,业内专家指出,在长序列推理场景下,中间激活甚至能占到总显存消耗的40%以上。
这里就出现一个分岔路口:
- 保存所有痕迹:下次计算梯度或做反向传播时直接用,速度快但显存巨大。
- 丢弃痕迹,需要时重算:只保留最原始的输入和权重,每次需要某层的梯度或结果时,从原始输入重新前向计算一次。
推理阶段没有反向传播,所以很多人以为不需要保存激活值,但实际推理时,KV Cache本身就是一种“痕迹”它保存了所有历史token的K和V,而推理痕迹重计算,指的是在生成新token时,不再把整条序列的KV全部缓存,而是只缓存一部分,缺失的部分根据之前的原始输入重新计算,这种思路在长上下文场景下特别有效。
推理痕迹重计算的两种典型实现路径
- 按层重算:每生成一个新token,只保留最后一层的KV Cache,前面所有层的KV在需要时从初始token重新前向计算,这适合显存极小但能容忍延迟的场景。
- 按块重算(chunked recompute):把序列切成块,已生成的块保留KV,新块产生时重算旧块中部分数值,这是目前主流框架如vLLM、SGLang里常见的分页KV策略的底层逻辑之一。
为什么重计算能节省显存,代价又是什么
核心逻辑很简单:显存显示的是峰值占用,而重计算把峰值压平了,不重算时,一次性存下所有痕迹,峰值极高;重算时,每次都只保留一小块临时缓冲,用完就释放,峰值大幅下降。

显存节省的量化感知
我们不用编造具体百分比,但可以给出一个直观的场景对比,假设你有一张24GB的显卡,跑一个70B量级的模型(需要量化),如果全量保存KV Cache和所有中间激活,可能上下文长度只能开到2048,开启推理痕迹重计算后,中间激活几乎零存储,KV Cache只保留最近几层,上下文长度能轻松开到8192甚至更多。这是数量级的差距,不是省个10%的小打小闹。
延迟和吞吐的权衡
代价是计算量上去了,重算意味着同一段前向传播要跑两次甚至更多次,对于实时交互场景,比如聊天机器人的首token延迟,重计算会增加几十毫秒到几百毫秒的等待,但对于离线批量推理、文档分析这类对延迟不敏感的任务,多花点算力换显存容量,成本上非常划算。
行业共识认为,推理痕迹重计算在以下场景中收益最大:
- 长上下文阅读(如合同、论文、代码库)
- 单卡部署超大模型(如8K以上上下文)
- 需要高并发但显存有限的GPU集群
实操:在你的推理框架中激活重计算
大部分主流推理框架已经把重计算做成了开关选项,你不需要自己改CUDA代码,只需要找到正确的配置项。
vLLM中的开启方法
vLLM从0.4版本开始支持--recompute-kv参数(具体名称可能随版本变化,请查阅你所用版本的文档),启动命令类似:
python -m vllm.entrypoints.openai.api_server --model /path/to/model --max-model-len 32768 --recompute-kv
开启后,vLLM会丢弃不再频繁使用的历史KV块,在需要时根据原始token重新计算,如果你在用llm = LLM(model="...", recompute_kv=True)的编程接口,效果类似。
llama.cpp中的做法
llama.cpp的CPU推理时,通过--no-mmap和--mlock可以控制内存映射,而针对KV Cache的重算,可以通过--keep参数控制保留的token数量,比如只保留最近256个token的KV,更早的上下文在生成下一个token时重新计算:
llama-cli -m model.gguf -c 8192 --keep 256
这个命令会让模型每生成一个新token,都重新处理前面被丢弃的token,显存占用接近恒定。
自己实现时最关键的三个步骤
如果你用的是自定义推理脚本,手动实现重计算并不复杂,关键路径如下:
- 定位痕迹存储点:找到代码里保存每一层KV Cache和中间激活的列表,通常在
forward函数的循环里。 - 设置阈值:定义什么条件下触发重算,比如当序列长度超过某个值,或者缓存区的显存占用超过预设上限。
- 重算入口:从最初的
input_ids和位置编码重新跑一遍前向,直到填充缺失的断层,注意要复用已经计算出的最近层KV,避免全量重算。

实际操作中,推荐你写一个缓存对象,内部维护“已保存的token索引区间”和“最近一次计算的layer depth”,生成时如果需要的token不在区间内,就触发一次局部重算。
结合量化与重计算:单卡跑长上下文的最优解
推理痕迹重计算不是孤立的技巧,它和量化、张量并行是黄金搭档,很多人在百度搜索“大模型推理显存不够怎么办”时,得到的第一答案是量化,但纯粹量化只压缩权重,对KV Cache和中间激活的压缩有限,如果你想在单卡4090部署Qwen2.5-72B并支持长上下文,方案组合如下:
- 权重使用4-bit或8-bit量化,降低基底占用
- 开启KV Cache量化到8-bit或4-bit
- 同时开启推理痕迹重计算,把不常用层的前向结果丢弃
三者叠加后,24GB显存可以部署72B模型并支持8K上下文,这在量化+全缓存的时代几乎不可能,你可以在HuggingFace的Transformers库中通过use_cache=False禁用默认缓存,再手动实现部分重算,效果会更可控。
一个可验证的实验路径
想验证重计算到底省了多少显存,你可以做这样一组对比实验:
- 用相同的输入序列长度(比如4096),关闭重计算,记录峰值显存
nvidia-smi读数。 - 开启重计算,保持同样的上下文长度,记录峰值显存。
- 比较两个值的差,同时记录生成速度的变化。
注意运行时要用torch.cuda.reset_peak_memory_stats()复位测量点,否则测得的是进程累计峰值,这个小技巧能帮你精准评估重计算在你自己业务场景下的真实收益。
重计算在不同任务下的适用边界
不是所有场景都适合重计算,如果你做的是高频短对话,每次输入输出都只有几十个token,重计算带来的额外延迟会显著影响用户体验,这时候老老实实全量缓存反而最优。
适合重计算的任务画像
- 批量长文本摘要:几千上万字的输入,输出几百字摘要,重算开销占比低。
- RAG检索增强生成:一次性注入多个知识块,上下文很长,但问答之间的等待时间可以容忍。
- 代码补全与仓库分析:需要跨文件上下文,但每次补全后的反馈不需要严格实时。
不适合重计算的任务画像
- 流式语音交互:要求每200-500毫秒返回一个token,重算会卡顿。
- 高频API服务:吞吐量优先,重算会导致TPO(每秒输出token数)明显下降。

显存节省思路的进阶变体
除了基础的整层重算,还有几种变体能进一步优化收益。
选择性重算与启发式策略
记录每一层KV Cache的“命中率”即某层的历史KV在后续计算中被重复使用的频率,命中率高的层保留,命中率低的层丢弃,这种方式比纯重算更聪明,显存节省和速度损失更均衡。
与投机采样(Speculative Decoding)结合
投机采样时,草稿模型会生成多个候选token,目标模型验证时也需要KV,这时候可以让草稿模型全量缓存,目标模型使用重计算,因为目标模型的验证步骤本身就需要重新计算logits,顺带重算KV并不会增加额外前向次数,这个组合的收益是双重的。
动态显存压力感知重算
写一个显存监控线程,当剩余显存低于阈值时,自动清空最古老的那部分KV记录,并标记为“可重算”,当模型再次访问这些位置时,从原始输入重算,这相当于一个智能缓存淘汰系统,比固定策略灵活得多。
常见问题解答
推理痕迹重计算会显著降低生成速度吗?
会降低,但幅度取决于重算的频次,如果每生成一个token都重算全部历史,速度可能下降5到10倍,但如果只重算被淘汰的块,且淘汰块占比很小,速度下降能控制在20%左右,建议用--keep这类参数把最近token固定在缓存里,重算仅作用于早期序列,速度损失更小。
重计算和KV Cache量化哪个更省显存?
量化是纯压缩,KV Cache从FP16变成INT8或INT4,显存直接减少一半或四分之三,但精度略有损失,重计算是删除后重建,显存节省取决于删除比例,理论上可以节省90%以上的KV显存,但引入额外计算,实际项目中,两者可以同时开启,量化降低单块KV的大小,重计算减少KV的块数,叠加效果最明显。
开启重计算后显存还是不够用怎么办?
检查你是否把输入序列的所有token都送入了模型,有些框架会在forward时默认保存所有层的神经元输出(用于训练模式),推理时需要显式切换到torch.no_grad()并关闭grad_mode,检查是否误开梯度检查点(gradient checkpointing),那是训练用的重计算,推理时应该关闭,如果仍不够,把--keep数值调小,让更早的上下文全部重算,或者考虑换用更激进的量化方式。任何重计算策略都无法超越物理显存上限,最终约束依然是模型权重加上最小瞬时缓存的总和。