推理痕迹重计算不是把已算过的结果扔掉重来,而是主动释放推理过程中”用后即弃”的中间激活值,在需要时用算力换显存。这套思路在长序列推理和超大模型部署场景下,能把峰值显存压缩一个量级,代价仅仅是多花20%左右的GPU时间,目前主流推理框架和国产化适配方案中,这已经是与INT8量化、KV Cache量化并列的三大显存优化手段之一。
推理显存不够怎么办:先看清你的显存被谁吃掉了
运行一个大模型推理任务,显存占用主要分四大块,先把敌人认清,才能明白重计算为什么省显存。
- 权重参数:模型文件本身,FP16格式下一个7B模型约14GB,13B约26GB,这部分通常用模型并行切片分摊到多卡。
- KV Cache:生成每个token时,注意力机制要缓存历史token的Key和Value矩阵,这个变量随序列长度线性增长,在长上下文的场景下,它才是显存大头,比如8K上下文长度的13B模型,KV Cache峰值轻松超过15GB。
- 中间激活值:前向传播过程中,每一层LayerNorm、MLP、Attention算出来的临时张量。这部分是纯计算中转站,用完立刻失效,但这一刻它占用4GB到8GB不等。
- 框架运行时开销:CUDA context、算子调度、碎片化存储,约1GB到2GB。
不用拆机检测,直接在代码里加三行PyTorch日志就能看见配方:
import torch
print(f"allocated: {torch.cuda.memory_allocated() / 10243:.2f} GB")
print(f"reserved: {torch.cuda.memory_reserved() / 10243:.2f} GB")
print(f"max_allocated: {torch.cuda.memory_max_allocated() / 10243:.2f} GB")
跑一个batch=1、上下文4K的推理任务,你通常会看到:权重占固定值,KV Cache逐步爬升,但激活值的瞬时峰值才是那个”冲顶”的元凶,行业共识认为,长序列推理场景下,激活值加KV Cache的临时变量合计占整体显存的40%到60%,这个比例超过很多人的直觉。
大模型推理显存优化方案:重计算的底层逻辑与落地姿势
如果把显存比作一个精打细算的房间,权重参数像大件家具不能扔,KV Cache像随时要翻的旧账本得留着,那激活值就是装修用的脚手架拼完一面墙就该拆掉,腾出空间干下一件事,重计算的做法是:完事儿就拆,用到的时候再搭。
具体到技术实现,推理阶段的重计算不是把整个前向过程全部推倒重来,而是按层按粒度做策略性丢弃,Transformer的每一层包含Attention子层和MLP子层,两者中间有残差连接和LayerNorm,执行顺序是:
- 前向传播时,每算完一个Block,把该Block内部的全部中间激活值(比如Q、K、V矩阵,MLP中间层结果)从显存里释放掉。
- 梯度或后续计算需要用到某个Block的输出时,
在那个时刻重新跑一遍该Block的前向计算
,重新生成激活值。 - 只保留每个Block的最终输出和残差连接,用于下游层计算。
这个方案的省显存效果不是线性的,而是按层数累积的,以13B模型、42层Transformer为例,释放全部中间激活后,峰值显存可以直接下降30%到45%,代价是每一个被重计算的层都要多跑一遍前向计算,GPU利用率增加,整体延迟变长。
一半重计算是收益和代价的黄金平衡点每隔两层丢弃一次,也就是50%的重计算比例,这个配置下显存节省可以达到全面重计算的70%左右,而性能损耗只有全面重计算的大约三分之一,在vLLM、SGLang框架中通过参数直接开启即可,不需要改模型代码。
重计算与流水线并行显存对比:不同的取舍路径
重计算和流水线并行常常被混为一谈,但两者的节省方向和代价完全不同,流水线并行的思路是把层切成几段,放几张卡上各算各的,它解决的是单卡放不下模型整体的问题,节省的是权重参数那一部分显存,重计算解决的是单卡上有富余算力、但显存峰值顶不住的问题,节省的是临时变量那一部分。
两张方案可以叠加使用,但要注意生效顺序:
- 先做模型并行把权重分到多张卡
- 再做张量并行削减单层激活值
- 最后开重计算压缩剩余的中间变量
对于一张3090或4090级别的消费级显卡,不适合做流水线并行,因为单卡显存根本承载不了模型的输入序列和中间计算结果,这种情况,首选就是重计算加INT8量化,对于A100或H100级别的数据中心卡,如果是多卡机群,优先考虑张量并行加切片,同时开启重计算来缓冲显存抖动。
实测下来,用H20卡部署7B模型、上下文16K时的表现:
| 方案组合 | 峰值显存占用 | 生成速度变化 |
|---|---|---|
| 不优化基线 | 38GB,频繁OOM | 基准 |
| INT8量化 | 21GB,稳定运行 | 下降约8% |
| INT8 + 全量重计算 | 14GB,稳定运行 | 下降约22% |
| INT8 + 半量重计算 | 17GB,稳定运行 | 下降约12% |
这张表想表达的核心是:显存从38GB压到14GB,大部分功劳来自重计算,但对于追求低延迟的生产环境,过高的重计算比例反而破坏了体验。
推理显存优化里的重计算实操手册
按具体场景梳理,重计算的落地路径因框架而异,但整个思路完全一致。
- HuggingFace Transformers直接推理:在
generate()函数中传入use_cache=True、output_attentions=False、output_hidden_states=False,这三个参数能拦截大量中间激活值的保留,是最轻量的一层重计算变体。 - DeepSpeed-Inference:配置文件中启用
replace_with_kernel_inject和enable_recompute,并设置recompute_granularity = "selective",表示只对Attention层做重计算,MLP层保留因为Attention层的激活值更大,重算代价也更可控。 - vLLM部署场景:vLLM的
--max-num-seqs参数决定了并发序列数量,每个序列的激活值都独立占据显存,开启--enable-prefix-caching的同时,把--gpu-memory-utilization调低至0.85,给激活值的临时波动留出缓冲带,这不完全等同于传统重计算,但思路是相同的通过控制并发数来限制激活值的总量。 - SGLang的RadixAttention:原生支持Prefix Cache,在长对话和批量推理场景下,能跳过重复部分的重计算,只对增量部分走一遍前向,间接减少了激活值的重新计算量。
如果习惯手动控制,还有一个粒度更细的PyTorch原生技巧用torch.utils.checkpoint包装自定义的Transformer Block:
from torch.utils.checkpoint import checkpoint
def forward_with_recompute(block, x):
return checkpoint(block, x, use_reentrant=False)
use_reentrant=False是PyTorch 2.x后的推荐写法,显存管理更精确,但要求Block内部不能有原地操作(比如torch.add_(...))。
重计算的聪明用法:只对长上下文场景开启
重计算不是免费的午餐,在短上下文、低并发的场景下,重计算是纯纯的负优化,试想上下文1K的短文生成,激活值总共占不到2GB,再来回重算一遍,只增加了延迟,显存盈余根本没变化。
哪些场景值得开重计算,判断标准很简单:
- 上下文长度超过4K:此时KV Cache和激活值开始线性膨胀,重计算收益显著放大
- 并发数超过8:多路同时推理时,激活值拼在一起形成显存峰值,重计算可以让峰值腰斩
- 显存吃紧但不想降精度:不想从FP16降到INT8损失一点精度,重计算是保精度的唯一大门
- 做7B以上本地部署、跑在消费级显卡上:3090玩13B模型时,重计算通常是活下来的关键
有一个相关参数经常被忽略max_length,很多人的服务OOM不是模型权重爆了,而是生成的输出长度上限设得太高,把max_length从8192调到4096后再开半量重计算,显存直接砍半,速度损失可以忽略不计,这种组合拳比单纯开重计算更实用。
产出到这一步,有一个SGLang原生支持的特性值得提一句:当两个请求共享同一个前缀时,重计算可以完全跳过该前缀,这意味着在多用户RAG场景下,大量请求会引用同一份知识库内容,前向计算只跑在差异化的尾部,显存和算力同步被省下。
重计算与KV Cache量化的联合使用要点
推理显存优化方案不只有重计算一个选项,KV Cache量化是它的天然搭档,重计算徒手压缩激活值,KV Cache量化负责把Key和Value矩阵从FP16降到INT8甚至INT4,两套方案作用在不同层级的显存对象上,互不冲突一个针对Attention计算中间产物,一个针对Attention的历史状态容器。
联合使用时有一个顺序要求:先量化KV Cache,再配置重计算比例,原因是KV Cache量化后,显存峰值从KV Cache转移到了激活值,此时重计算的比例可以适当降低,用更小的性能损耗换回更多的显存交易。
部署7B模型时,合理的推演路径是:默认权重加载约占14GB,8K上下文KV量化到INT8占约4GB,激活值保留在全量重计算模式下的2GB内,最终总占用约20GB,刚好塞进一张4090的24GB显存,同样的模型不加以优化,32GB模块也未必装得下,近期国产芯片适配大模型推理时,多数情况下也是将重计算作为开启项,配合少量KV量化来换取平台可承载的并发路数,对于广大开发者来说,vLLM提供的--kv-cache-dtype配合--recompute-proportion参数,可在一台A800上同时扛起4路长对话请求。
关于推理痕迹重计算的常见疑问解答
重计算会影响最终生成文本的质量吗?
不会,推理阶段的重计算只是重新算一遍相同的数学操作,FP16精度下浮点运算是确定性的,重新计算出来的激活值和原来完全一致,它与训练时的激活重计算(会省内存但梯度有近似)有本质区别,推理重计算是无损的,这一点可以从PyTorch的checkpoint接口设计里得到佐证它不改变任何张量的数值路径,只是多了几次前向耗时。
为什么有时候开启重计算后显存反而没降多少?
大概率是内存碎片问题,PyTorch的CUDA缓存分配器不立即把释放的显存还给驱动,而是留在本地缓存池中重复使用,这会让nvidia-smi观察到的显存占用没有明显下降,实际能用的显存上限已经提高,需要看torch.cuda.max_memory_reserved()这个指标,另一个高发原因是残差连接把各层激活值串在了一起,导致释放失败,可以通过检查代码中是否有对中间张量的全局引用来排查。
重计算和DeepSpeed的ZeRO-Offload哪个更适合单卡跑大模型?
单卡场景下两者可以并列使用,但目标不同,ZeRO-Offload是把权重参数和优化器状态搬到CPU内存,解决的是权重放不下单卡的问题,它的瓶颈在PCIe带宽,重计算只负责处理GPU显存里临时的激活值,不碰权重存储,实测中,7B模型配单张消费级显卡时,先做INT8量化权重,再开启半量重计算,代价和收益最均衡,如果上下文要求达到32K以上,则需要量化KV Cache配合,仅靠重计算已经无力回天。
首发原创文章,作者:王坚,如若转载,请注明出处:https://idctop.com/article/623261.html





