Transformers KV cache 怎么在显存不足时启用卸载
来源:17golang原创
时间:2026-10-05 13:08:15 499浏览 收藏
Transformers 生成长文本时,如果模型权重已经占去大部分显存,继续增长的 KV cache 很容易成为最后一根稻草。最直接的处理不是关闭缓存,而是在 generate() 中设置 cache_implementation="offloaded":当前注意力层的缓存留在 GPU,其余层的缓存移到 CPU,用更多数据传输换取更低的显存占用。
- 首次就确定显存紧张:直接使用
cache_implementation="offloaded"。 - 希望平时保持默认吞吐、仅在 OOM 时降级:捕获
torch.OutOfMemoryError后清理缓存,再用 offloaded 重试。 - 已经使用固定缓存容量或
torch.compile:评估offloaded_static,同时明确最大缓存长度。
业务负载:为什么生成阶段才突然 OOM
自回归生成每次只预测一个或少量 token,但后续 token 仍要使用此前各层注意力计算产生的 key/value。KV cache 保存这些状态,避免反复计算全部历史。代价是缓存会随上下文长度、输出长度、批量大小和 beam 数增加;模型权重能装入 GPU,不代表完整生成过程一定能装下。
因此,判断是否需要卸载时应看真实生成峰值,而不是只看模型加载后的 nvidia-smi 数字。长提示词能完成预填充、进入解码后才 OOM,或者提高 max_new_tokens、num_beams 后失败,都说明 KV cache 很可能正在挤压剩余显存。
约束条件:卸载不是免费显存
官方当前文档将 DynamicCache 作为默认缓存。启用卸载后,除当前层外的大部分层缓存驻留在 CPU;模型遍历各层时,会异步预取下一层缓存,并把完成计算的当前层缓存送回 CPU。显存压力下降了,但 CPU 内存占用和设备间传输随之增加。

这套方案适合“GPU 显存是硬约束、主机内存仍有余量”的机器。若 CPU 内存也接近上限,或者链路带宽很低,卸载可能只是把 OOM 从 GPU 转移到系统内存,并明显拉低生成吞吐。对延迟敏感的在线服务要先压测,再决定是否默认开启。
方案对比:三种策略怎么选
| 策略 | 显存特点 | 主要代价 | 适用场景 |
|---|---|---|---|
| 默认 DynamicCache | 缓存随生成增长,主要在设备侧 | 长上下文下显存压力较高 | 显存充足,优先吞吐 |
offloaded | 只让当前层缓存驻留 GPU | 增加 CPU/GPU 传输,吞吐可能下降 | 显存紧张,输入和输出长度变化较大 |
offloaded_static | 固定容量缓存并执行卸载 | 要预留固定上限,过大可能浪费内存 | 固定形状、静态缓存或编译优化路径 |

量化缓存也是节省内存的思路,但它改变的是缓存表示精度,并非本文的“层缓存卸载”任务。遇到显存不足时,先用 offloaded 做低侵入回退更容易定位问题;只有 CPU 内存或传输成本也成为瓶颈时,再单独评估缓存量化。
推荐架构:在 generate() 里直接启用卸载
如果已知目标机器显存紧张,可以把卸载作为该部署配置的默认策略。下面沿用官方文档展示的 Phi-3 小模型,关键只在最后一行的 cache_implementation 参数。
import torch
from transformers import AutoModelForCausalLM, AutoTokenizer
model_id = "microsoft/Phi-3-mini-4k-instruct"
tokenizer = AutoTokenizer.from_pretrained(model_id)
model = AutoModelForCausalLM.from_pretrained(
model_id,
dtype=torch.float16,
device_map="auto", # 让 Transformers/Accelerate 选择模型放置位置
)
inputs = tokenizer(
"请用三点解释 KV cache 的作用。",
return_tensors="pt",
).to(model.device)
outputs = model.generate(
**inputs,
do_sample=False,
max_new_tokens=256,
cache_implementation="offloaded", # 将大部分层的 KV cache 放到 CPU
)
print(tokenizer.decode(outputs[0], skip_special_tokens=True))
这不会把模型权重自动全部移到 CPU,也不是把缓存写入磁盘。它只改变 KV cache 的层级驻留与传输策略。若项目使用的是较旧 Transformers 版本,应先以该版本文档和 GenerationConfig 支持项为准;不要只升级一处参数后假设所有模型实现都兼容。
OOM 回退:正常情况走默认缓存,失败再卸载
服务端更常见的需求是保留默认缓存的吞吐,只在极端长请求触发 OOM 时降级。可以把生成入口封装成一次正常尝试和一次卸载重试。重试前清理 PyTorch 缓存分配器中未使用的显存,但仍被活跃张量引用的空间不会因此释放,所以输入和模型对象的生命周期仍要控制好。
import torch
def generate_with_kv_fallback(model, **generation_kwargs):
try:
# 常规请求先使用默认 DynamicCache,保留更好的吞吐表现。
return model.generate(**generation_kwargs)
except torch.OutOfMemoryError:
if not torch.cuda.is_available():
raise
# 只清理未被活跃张量占用的 CUDA 缓存,然后用卸载缓存重试一次。
torch.cuda.empty_cache()
generation_kwargs["cache_implementation"] = "offloaded"
return model.generate(**generation_kwargs)
outputs = generate_with_kv_fallback(
model,
**inputs,
do_sample=False,
max_new_tokens=256,
)
生产环境还应把“是否发生回退”写入指标或日志,并限制重试次数。一次默认失败加一次卸载重试已经足够;无限重试会放大延迟和资源竞争。批量请求场景中,还要避免一个超长请求拖慢同批其他请求。
风险点:显存下降后还要看什么
- CPU 内存:长上下文的 KV cache 仍然存在,只是主要从 GPU 搬到了主机内存。
- 传输带宽:层间预取与回写会增加链路工作量,实际吞吐损失取决于模型、上下文、生成长度和 beam 配置。
- 并发:单请求能跑通不代表多并发安全,多个卸载缓存会共同占用主机内存和传输通道。
- 静态容量:
offloaded_static需要固定容量思维,最大长度设得过小会不够用,设得过大又可能浪费内存。 - 滑动窗口层:直接实例化缓存时可通过
offload_only_non_sliding决定是否卸载滑动窗口或分块注意力层;这些层缓存通常较短,少搬运可能更快。
落地清单
- 用真实最长提示词和最长输出复现显存峰值,确认问题发生在生成而非模型加载。
- 先在单请求下设置
cache_implementation="offloaded",确认 OOM 消失且输出可以正常解码。 - 同步观察 GPU 峰值、CPU 峰值、首 token 延迟和持续生成吞吐,不只记录“能否跑通”。
- 如果正常请求更重视速度,改用 OOM 回退;如果显存始终不足,则把 offloaded 固定为部署配置。
- 只有在固定容量和编译优化确有收益时才切换
offloaded_static,并验证最大缓存长度。 - 按并发上限做压力测试,为 CPU 内存和链路带宽保留安全余量。
常见问题
关闭 use_cache 能解决显存不足吗?
use_cache=False 可以不保留 KV cache,但会让后续 token 重复计算历史,生成通常明显变慢。目标只是降低 GPU 缓存占用时,offloaded 更符合需求。
卸载后生成结果会改变吗?
卸载策略改变的是缓存驻留位置与传输方式,不是采样参数。相同模型、输入和确定性生成配置下,它应作为默认缓存的替代路径;不过不同软件版本和硬件后端仍应在项目中做回归测试。
官方接口依据在哪里?
https://huggingface.co/docs/transformers/main/en/kv_cache 的 Cache offloading 部分给出了 offloaded、offloaded_static 和 OOM 回退示例。main 文档面向源码主线,使用已发布版本时请切换到对应版本文档核对。
-
284 收藏
-
387 收藏
-
328 收藏
-
426 收藏
-
147 收藏
-
164 收藏
-
291 收藏
-
372 收藏
-
314 收藏
-
293 收藏
-
247 收藏
-
194 收藏
-
268 收藏
-
316 收藏
-
150 收藏
-
148 收藏
-
251 收藏
-
- 前端进阶之JavaScript设计模式
- 设计模式是开发人员在软件开发过程中面临一般问题时的解决方案,代表了最佳的实践。本课程的主打内容包括JS常见设计模式以及具体应用场景,打造一站式知识长龙服务,适合有JS基础的同学学习。
- 立即学习 543次学习
-
- GO语言核心编程课程
- 本课程采用真实案例,全面具体可落地,从理论到实践,一步一步将GO核心编程技术、编程思想、底层实现融会贯通,使学习者贴近时代脉搏,做IT互联网时代的弄潮儿。
- 立即学习 516次学习
-
- 简单聊聊mysql8与网络通信
- 如有问题加微信:Le-studyg;在课程中,我们将首先介绍MySQL8的新特性,包括性能优化、安全增强、新数据类型等,帮助学生快速熟悉MySQL8的最新功能。接着,我们将深入解析MySQL的网络通信机制,包括协议、连接管理、数据传输等,让
- 立即学习 500次学习
-
- JavaScript正则表达式基础与实战
- 在任何一门编程语言中,正则表达式,都是一项重要的知识,它提供了高效的字符串匹配与捕获机制,可以极大的简化程序设计。
- 立即学习 487次学习
-
- 从零制作响应式网站—Grid布局
- 本系列教程将展示从零制作一个假想的网络科技公司官网,分为导航,轮播,关于我们,成功案例,服务流程,团队介绍,数据部分,公司动态,底部信息等内容区块。网站整体采用CSSGrid布局,支持响应式,有流畅过渡和展现动画。
- 立即学习 485次学习