登录
推荐 文章 Go 技术 课程 下载 专题 AI
首页 >  科技周边 >  人工智能

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 内存占用和设备间传输随之增加。

Transformers KV cache 在 GPU 当前层与 CPU 其他层之间的卸载关系图
图1:KV cache 卸载的驻留关系;GPU 只保留当前层缓存,其余层主要放在 CPU,并在层间预取与回写。

这套方案适合“GPU 显存是硬约束、主机内存仍有余量”的机器。若 CPU 内存也接近上限,或者链路带宽很低,卸载可能只是把 OOM 从 GPU 转移到系统内存,并明显拉低生成吞吐。对延迟敏感的在线服务要先压测,再决定是否默认开启。

方案对比:三种策略怎么选

策略显存特点主要代价适用场景
默认 DynamicCache缓存随生成增长,主要在设备侧长上下文下显存压力较高显存充足,优先吞吐
offloaded只让当前层缓存驻留 GPU增加 CPU/GPU 传输,吞吐可能下降显存紧张,输入和输出长度变化较大
offloaded_static固定容量缓存并执行卸载要预留固定上限,过大可能浪费内存固定形状、静态缓存或编译优化路径
DynamicCache offloaded 和 offloaded_static 三种缓存策略约束对照图
图2:三种 KV cache 策略的约束对照;选择时同时检查显存、CPU 内存、吞吐和固定容量需求。

量化缓存也是节省内存的思路,但它改变的是缓存表示精度,并非本文的“层缓存卸载”任务。遇到显存不足时,先用 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 决定是否卸载滑动窗口或分块注意力层;这些层缓存通常较短,少搬运可能更快。

落地清单

  1. 用真实最长提示词和最长输出复现显存峰值,确认问题发生在生成而非模型加载。
  2. 先在单请求下设置 cache_implementation="offloaded",确认 OOM 消失且输出可以正常解码。
  3. 同步观察 GPU 峰值、CPU 峰值、首 token 延迟和持续生成吞吐,不只记录“能否跑通”。
  4. 如果正常请求更重视速度,改用 OOM 回退;如果显存始终不足,则把 offloaded 固定为部署配置。
  5. 只有在固定容量和编译优化确有收益时才切换 offloaded_static,并验证最大缓存长度。
  6. 按并发上限做压力测试,为 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 文档面向源码主线,使用已发布版本时请切换到对应版本文档核对。

声明:本文转载于:17golang原创 如有侵犯,请联系study_golang@163.com删除
相关阅读
更多>
最新阅读
更多>
课程推荐
更多>