当前位置:首页 > 文章列表 > 科技周边 > 人工智能 > Transformers KV cache 怎么在显存不足时启用卸载

Transformers KV cache 怎么在显存不足时启用卸载

来源:17golang原创 2026-10-05 13:08:15 0浏览 收藏

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删除
横风动漫为什么会请求通知权限?应用商店字段与更新提醒边界说明横风动漫为什么会请求通知权限?应用商店字段与更新提醒边界说明
上一篇
横风动漫为什么会请求通知权限?应用商店字段与更新提醒边界说明
Go SIMD 代码怎么保留标量回退路径
下一篇
Go SIMD 代码怎么保留标量回退路径
查看更多
最新文章
查看更多
课程推荐
  • 前端进阶之JavaScript设计模式
    前端进阶之JavaScript设计模式
    设计模式是开发人员在软件开发过程中面临一般问题时的解决方案,代表了最佳的实践。本课程的主打内容包括JS常见设计模式以及具体应用场景,打造一站式知识长龙服务,适合有JS基础的同学学习。
    543次学习
  • GO语言核心编程课程
    GO语言核心编程课程
    本课程采用真实案例,全面具体可落地,从理论到实践,一步一步将GO核心编程技术、编程思想、底层实现融会贯通,使学习者贴近时代脉搏,做IT互联网时代的弄潮儿。
    516次学习
  • 简单聊聊mysql8与网络通信
    简单聊聊mysql8与网络通信
    如有问题加微信:Le-studyg;在课程中,我们将首先介绍MySQL8的新特性,包括性能优化、安全增强、新数据类型等,帮助学生快速熟悉MySQL8的最新功能。接着,我们将深入解析MySQL的网络通信机制,包括协议、连接管理、数据传输等,让
    500次学习
  • JavaScript正则表达式基础与实战
    JavaScript正则表达式基础与实战
    在任何一门编程语言中,正则表达式,都是一项重要的知识,它提供了高效的字符串匹配与捕获机制,可以极大的简化程序设计。
    487次学习
  • 从零制作响应式网站—Grid布局
    从零制作响应式网站—Grid布局
    本系列教程将展示从零制作一个假想的网络科技公司官网,分为导航,轮播,关于我们,成功案例,服务流程,团队介绍,数据部分,公司动态,底部信息等内容区块。网站整体采用CSSGrid布局,支持响应式,有流畅过渡和展现动画。
    485次学习
查看更多
AI推荐
  • PubMedQA数据集详解:生物医学问答基准、功能与应用指南
    PubMedQA
    深入了解PubMedQA生物医学问答数据集,涵盖其核心功能、使用方法及在临床决策、药物研发等场景的应用,助力提升NLP模型性能。
    342次使用
  • H2O EvalGPT:开源LLM大模型评估与排行榜工具
    H2O EvalGPT
    H2O EvalGPT是H2O.ai推出的开源LLM评估平台,提供详细的大模型性能排行榜、行业特定基准测试及A/B测试功能,助您快速选择最适合项目的高性能大语言模型。
    398次使用
  • LMArena是什么?伯克利AI模型评估平台使用指南与功能解析
    LMArena
    LMArena是加州大学伯克利分校推出的AI模型匿名评测平台。通过盲测投票机制,用户可对比不同大模型回答并生成实时排行榜,助力开发者优化模型及用户选择最佳AI工具。
    392次使用
  • 斯坦福HELM:大语言模型Holistic Evaluation整体评估框架详解
    HELM
    深入了解斯坦福推出的HELM(Holistic Evaluation of Language Models)大模型评测体系。本文解析其核心功能、安装配置步骤及应用场景,涵盖准确性、公平性、鲁棒性等多维度指标,助力开发者全面优化语言模型性能。
    355次使用
  • MMBench详解:多模态大模型基准测试、功能特点与使用指南
    MMBench
    MMBench是由上海人工智能实验室等机构联合推出的多模态基准测试平台,提供细粒度能力评估、大规模数据集及VLMEvalKit工具。本文详细介绍其核心功能、安装使用方法及应用场景,助力开发者全面评估多模态模型性能。
    181次使用
微信登录更方便
  • 密码登录
  • 注册账号
登录即同意 用户协议 和 隐私政策
返回登录
  • 重置密码