ARTICLE DETAIL

资讯详情

深耕网站建设与运营推广的一线实战洞察。

大模型推理提速的核心:KV Cache原理、显存占用与优化实战

大模型推理提速的核心:KV Cache原理、显存占用与优化实战 开头先说个现象无论是你本地跑的Ollama还是线上那套并发负载很高的推理服务几乎所有大模型推理框架都在用同一个招数来提速——KV Cache。以前我刚开始接触LLM的时候对它也是一头雾水只知道它是“用来加速推理的缓存”但为什么能加速、缓存了什么、显存为什么经常被它吃光这些事直到自己动手调过几个服务、踩了若干次OOM之后才算真正搞明白。这篇文章就围绕KV Cache展开把原理、显存账本、实际配置和排错经验一次讲透。不管你是刚入门的新手还是已经在用Transformers、vLLM之类框架做推理部署的工程师都能从这里找到可以直接拿去用的东西。1. 先搞懂生成式推理的底层逻辑KV Cache到底在优化哪一步1.1 自回归机制的“每一次只生成一个Token”LLM大语言模型的生成方式和人类写文章很不一样。人类写一段话的时候脑子里大体有整体结构可以“跳着”想但LLM不行它必须一个Token一个Token地往外吐。这就是所谓的自回归Autoregressive机制每次只预测下一个词然后把新的词拼到输入里再继续预测下一个词。举个例子让模型生成“KV Cache是显存大户”这句话。实际过程是这样的第一步输入“KV Cache是”模型预测出“显存”第二步输入变成“KV Cache是显存”模型预测出“大户”第三步输入变成“KV Cache是显存大户”模型预测出“。”第四步输入变成“KV Cache是显存大户。”模型预测出“结束符”停下。每一步输入序列都会变长一点模型要对“当前完整输入”做一次前向计算才能给出下一个Token的预测结果。这就是LLM推理慢的根本原因之一生成N个Token需要做N次完整的前向传播。1.2 注意力计算中那些“重复劳动”要理解KV Cache为什么能提速关键是把一次前向传播里发生的注意力Attention机制看透。Transformer的计算里输入序列的每个Token都会被映射成三个向量QQuery查询代表“我在找什么”KKey键代表“我能被谁匹配”VValue值代表“我实际携带的内容”。注意力计算的本质就像在图书馆里做检索。你手里拿着一个查询词Q挨个跟书架上的书签K做匹配算出一个相似度分数然后按照分数从高到低把对应书籍的内容V加权汇总出来。这个加权汇总的结果就是当前Token在“看完”整个序列之后得到的语义表示。现在最关键的问题来了。你在生成第2个Token时需要计算“KV Cache是”里每个Token的K和V用来跟当前Token的Q做匹配。等你生成第3个Token时序列变成了“KV Cache是显存”如果不做缓存模型就得把“KV Cache是显存”里所有Token的K和V重新算一遍。注意前三个Token“KV”、“Cache”、“是”的K和V跟你在上一步算出来的其实一模一样没有任何变化因为它们的内容没有变过。这就是明显的重复劳动。序列越长重复计算的量就越大。举个例子当上下文长度到4096个Token时你每生成一个新Token都得把之前4096个Token的K和V重新算一遍。这个开销是随序列长度线性增长的生成512个Token时等于白白多算了两百万次Token级别的K/V映射累积计算。KV Cache的出现就是针对这个重复计算做的优化把历史上已经算好的K向量和V向量存下来下次直接用不再重复计算。这就是“缓存”二字的由来。1.3 从一次完整生成过程看KV Cache的收益把自回归循环展开对比一下开缓存和不开缓存的计算差异。假设当前序列长度为N模型要再生成M个Token。不加KV Cache生成第N1个Token时要计算完整的长度为N1的序列的K/V生成第N2个Token时要计算长度为N2的序列的K/V……总计算量几乎是O(N²)级别的增长因为每一步都重新盯着越来越长的输入做全量计算。加了KV Cache生成第N1个Token时只需要计算位置N1这个Token自己的K/V历史K/V直接从缓存里取生成第N2个Token时同样只需要算一个新Token的K/V。每一步的K/V计算都变成了常数级别总计算量只是O(N M)级别而不是O(N×M)。带Transformer结构的大模型K/V这批向量在每一步都会经过对应的权重矩阵做线性变换权重参数量动辄百万级。不做缓存时每一步相当于把几百上千个Token的K/V都重新“过一遍权重”做了缓存后只需算当前这1个Token。实际项目中开启KV Cache后单Token生成延迟能降低一个数量级以上具体数值取决于序列长度和模型规模。2. KV Cache缓存的内容、存储位置与代码视角2.1 为什么缓存K和V却不缓存Q上一节说了K和V要缓存那Q呢为什么没人提“Q Cache”原因在于Q是“当前要预测下一个Token的那个Token”所特有的表示。每一步自回归时当前的新Token都是上一次刚生成的它没有历史——不存在“之前算过”的情况。换句话说Q永远是新的所以没有缓存价值。而且Q只用于和历史的K做匹配算完注意力权重之后它的使命就结束了下一轮生成中不会再被用到。所以KV Cache缓存的是每一层、每一个注意力头、每一个历史Token对应的K向量和V向量。Q每一轮都要新鲜计算只有K和V是“旧货可回收”的。2.2 KV Cache在Transformer里的存储位置很多人以为KV Cache是一块统一的内存区域其实它是按层拆开的。Transformer模型里有L层每一层都有自己的注意力模块所以每一层都需要一份独立的KV Cache。以LLaMA-7B为例它有32层。假设输入长度为1024那么KV Cache里就有32份缓存每份缓存保存了1024个Token在该层产生的K和V。这个结构在实际代码里通常以元组tuple或列表list的形式存在每层一个元素元素里再分出K和V两个张量。在HuggingFace Transformers库中这个缓存对象就是模型generate方法里的past_key_values参数。手动推理时你可以做的事情很直观import torch from transformers import AutoModelForCausalLM, AutoTokenizer model AutoModelForCausalLM.from_pretrained(your-model-path) tokenizer AutoTokenizer.from_pretrained(your-model-path) inputs tokenizer(KV Cache 是, return_tensorspt) # 第一次前向此时还没有缓存 with torch.no_grad(): outputs model(**inputs, use_cacheTrue) past_kv outputs.past_key_values # 用缓存生成下一个Token next_token_id outputs.logits[:, -1, :].argmax(dim-1, keepdimTrue) inputs {input_ids: next_token_id} with torch.no_grad(): outputs model(**inputs, past_key_valuespast_kv, use_cacheTrue)注意第二次前向时input_ids里只有刚生成的这一个Token而模型凭借传入的past_key_values就拥有了对之前所有历史Token的“记忆”。如果没有传入缓存就必须把全部历史Token重新喂给模型。2.3 一个容易忽略的点Cache里的张量形状看模型代码时你会看到past_key_values里的张量形状通常是(batch_size, num_heads, seq_len, head_dim)这跟训练时的标准注意力张量形状略有区别。很多初学时容易疑惑为什么把seq_len放在第三维而不是第二维原因是推理时主要的工作是“用新的Q去跟Cache里的K做点积”把seq_len放在靠后的维度可以让批量矩阵乘法Batch MatMul更高效减少维度转置带来的开销。这个形状设计不是随意的。以LLaMA系列为例head_dim通常为128num_heads因模型而异。形状里batch_size这一维意味着什么意味着同一份推理进程里如果同时处理多个请求这几个请求的KV Cache是拼在一起存储的每个请求占一个batch槽位。这也是后面要说的显存计算的重要基础。3. 显存账本KV Cache占的空间可能比想象中大得多3.1 KV Cache显存计算公式实际部署LLM的时候很多人第一步被卡住的地方不是计算太慢而是显存不够。模型权重都还没算完KV Cache先爆了。KV Cache的显存占用有一个非常明确的公式显存大小 2K和V两个向量× 层数 × 注意力头数 × 每个头的维度 × 序列长度 × batch大小 × 每个参数占的字节数这里“2”代表K和V各有一份“层数”“头数”“头的维度”都是模型结构参数“序列长度”和“batch大小”是推理时的动态参数“每个参数占的字节数”由精度决定FP16通常是2字节FP32是4字节INT8是1字节。这个公式看着简单代入具体数字之后往往很震撼。3.2 用7B/13B/70B模型算一遍感受一下我们分别算三个常见模型的KV Cache占用推理精度都按FP162字节计算序列长度按2048batch设为1。模型层数注意力头数head_dim公式展开KV Cache估算值LLaMA-7B32321282×32×32×128×2048×1×2约1GBLLaMA-13B40401282×40×40×128×2048×1×2约1.6GBLLaMA-70B808GQA1282×80×8×128×2048×1×2约0.64GB这里有个反直觉的点70B的模型KV Cache反而比7B和13B小。原因在于70B使用了GQA分组查询注意力KV头的数量从几十个压缩到了8个。这说明KV Cache的大小不完全取决于模型参数量而更取决于模型结构和注意力头设计。如果不做任何优化用传统MHA多头注意力结构去想象一个70B模型KV Cache大概会到8GB以上训练或推理成本会高得让人崩溃。这也是为什么GQA已经成为近年大模型标配结构的原因。3.3 并发场景下的显存爆炸单个请求的KV Cache看起来还好但在真实服务里并发请求一来显存增长就可怕了。假设一个7B模型上下文长度设为4096batch size为8FP16推理2 × 32 × 32 × 128 × 4096 × 8 × 2 约16.78GB没错一个batch为8、上下文4096的7B模型推理场景KV Cache就吃掉了16GB多显存。而7B模型权重本身在FP16下约14GB。也就是说KV Cache加上权重一张显存24GB的显卡基本到头了这还没算激活值、临时计算缓冲区和框架自身开销。这就是为什么线上推理服务普遍投入高显存设备不是模型本身有多个G的权重需要放而是KV Cache在高并发长上下文的场景下直接成了显存开销的“大头”。从工程角度看这条账本告诉我们三件事长上下文是有代价的上下文长度翻倍KV Cache翻倍高并发是有代价的请求数翻倍KV Cache同步翻倍批处理不是免费的它是用KV Cache的显存换吞吐量。3.4 上下文长度对KV Cache的影响上限很多新入坑的同学在调模型框架时喜欢把max_length或max_model_len调到很大比如32K、128K。模型权重不增加显存却可能悄然涨到难以接受。以7B模型为例按上文参数计算序列长度为4096时KV Cache约2GB序列长度为32768时这个数字变成约16GB到128K时约64GB。即便模型支持长上下文如果没有足够的显存预算推理时照样跑不起来。所以KV Cache直接决定了“模型宣称支持多长上下文”和“实际能跑多长上下文”之间的差距。这部分我后面在优化方案里还会展开。4. 在真实项目里用上KV Cache配置、工具与实测差异4.1 HuggingFace Transformers里的use_cache开关在Transformers库中KV Cache默认就是开启的。generate方法和直接前向时的use_cache参数控制它。# 开启缓存默认 outputs model.generate(**inputs, use_cacheTrue, max_new_tokens256) # 关闭缓存不推荐仅用于对比实验 outputs model.generate(**inputs, use_cacheFalse, max_new_tokens256)我建议你在自己的机器上做一个简单的对照组实验体会一下区别。用同一个模型、同一个输入、同样的最大生成长度分别开启和关闭use_cache记录推理耗时。实测下来关闭缓存后的耗时可能暴涨数倍而且模型越大、生成长度越长差距越离谱。不过有一点要注意use_cache开关只影响自动生成的循环过程不会改变模型权重和输出结果的内容本身。所以你在做推理速度优化时除非在做消融实验否则没有理由关闭它。4.2 vLLM等推理框架中的KV Cache控制在实际部署场景很多人不再直接用Transformers的generate而是用vLLM这类推理引擎。它的主要创新之一就是围绕KV Cache做文章通过PagedAttention把显存利用效率提上去。在vLLM中你主要关心的配置是gpu_memory_utilization和max_model_lenvllm serve your-model-path \ --gpu-memory-utilization 0.9 \ --max-model-len 8192 \ --tensor-parallel-size 1gpu_memory_utilization表示推理引擎最多可以占用多少比例的GPU显存。vLLM会把“模型权重占用的显存”从总显存里扣除之后把剩余可用显存全部规划给KV Cache池。max_model_len决定了单条请求的最大上下文长度它会直接影响KV Cache池的“单个Slot”容量。vLLM启动时通常会打印KV Cache的配置信息。比如一个7B模型在A100 40GB上启动你大概会看到类似这样的输出GPU KV cache size: 25.75 GB Maximum concurrency for 8192 tokens per request: 128这行信息的含义是框架在保证模型权重和运行开销的前提下把剩余的25GB多显存规划为KV Cache池在这个池子大小下如果每条请求最多占用8192个Token的KV Cache理论上最高可以并发处理128条请求。这里有个容易踩的坑max_model_len设得越大每个请求Slot能承载的上下限越高但能同时并发的请求数量就越少。如果你把max_model_len设成128K但实际业务里大部分请求只有几百Token那是纯粹的浪费——每个Slot都按最长上限预留能开出来的并发数直接少一个量级。反过来如果你的业务确实需要长文档问答max_model_len设小了长文档根本塞不进去。4.3 KV Cache与批处理Continuous Batching的关系传统批处理有一个问题一个batch里所有请求的生成长度必须对齐短请求生成完了也得等长请求生成完才能释放显存。这会导致KV Cache出现大量碎片和浪费。vLLM之所以快核心就在于Continuous Batching连续批处理它做到了“某个请求生成完了就立刻从KV Cache池中释放Slot把位置让给新请求”。配合PagedAttention对KV Cache的分块管理显存利用率能比Transformers默认方式高好几倍。所以从工程视角看KV Cache的“管理方式”和“缓存本身是否存在”同等重要。一个会释放、能复用、粒度更细的KV Cache管理机制能直接影响你的服务在同一块显卡上的吞吐天花板。4.4 在不同模块里看到KV Cache的实际痕迹动手写代码的时候可以从几个地方观察到KV Cache的存在generate日志里如果打印了past_key_values那就是缓存对象用Profiler例如PyTorch Profiler查看前向计算时如果看到大量矩阵乘法发生在当前Token的Q与历史K之间说明缓存正在被使用观察显存曲线如果没有KV Cache显存占用几乎随序列长度线性上升有KV Cache时显存占用也会上升但上升曲线更平缓、可控。建议你实际跑一个长文本生成任务用nvidia-smi -l 1实时盯着显存变化。你会发现随着生成Token数增加显存占用在缓慢爬升——这部分增量就是不断变长的KV Cache。这个观察能帮你建立显存直觉后面排查OOM时非常有用。5. 优化KV Cache的几条路线GQA、PagedAttention、量化与投机解码5.1 GQA/MQA从结构上减少KV Cache前面提到GQA能把70B模型的KV Cache压到非常小。它的思路是让多个Query头共享同一组KV头。传统的MHAMulti-Head Attention里32个Query头就有32组K/V头每组各算各的。GQA把这个比例压缩比如70B模型里Query头有好几十个但KV头只有8个多个Query头共用同一份K/V缓存。MQAMulti-Query Attention更进一步所有Query头共用一个KV头。这样做的好处是KV Cache规模直接缩减到原来的1/KV头数/Query头数推理显存压力大幅下降。代价是注意力表达的多样性可能略微下降但大规模预训练后的实践经验表明这个损失通常很小。今天的主流开源模型比如LLaMA-2-70B、LLaMA-3系列、Mistral等基本都采用了GQA或MQA。如果你在选模型做推理部署建议优先选采用GQA/MQA的版本这比在推理时再想各种软优化省事得多。5.2 PagedAttention像虚拟内存一样管理KV CachevLLM提出的PagedAttention灵感来自操作系统里的分页机制。传统KV Cache是一个连续的大块显存按请求预分配容易造成内部碎片和外部碎片。PagedAttention把KV Cache切分成固定大小的块Block以块为单位分配和释放。这个设计有个特别大的好处KV Cache不再要求物理上连续。请求需要的KV块可以被分散存放在显存不同位置通过块表把它们串起来。生成过程中新Token不断到来需要多少块就动态申请多少块用完的块立刻释放其他请求可以接着用。这个机制在长上下文和高并发场景下效果极其明显。vLLM能在相同的显卡资源上比原生Transformers推理高出数倍吞吐主要归功于它。你可以把PagedAttention理解成“KV Cache界的内存管理单元”它让显存利用率从“按最大可能预留”变成“按实际需求分配”。5.3 KV Cache量化用低精度换更多容量前面公式里有一个乘数——“每个参数占的字节数”。FP16是2字节如果把它压到INT81字节或INT40.5字节KV Cache直接减半甚至减到四分之一。实际做法是在模型推理过程中对K和V向量做实时量化。常见的做法是Per-Token或Per-Channel的量化保存缩放因子让K/V在低精度下保持可用精度。很多推理框架已经支持KV Cache量化配置例如vllm serve your-model-path \ --kv-cache-dtype fp8 \ --gpu-memory-utilization 0.9KV Cache量化后显存占用显著下降同一个池子里能塞下的并发请求数就上去了。代价是推理结果可能发生极细微的精度变化但在多数业务场景里几乎不可感知。如果你的服务吞吐压力大KV Cache量化值得认真考虑。5.4 投机解码让KV Cache配合“草稿模型”投机解码Speculative Decoding是这些年LLM推理加速的热门方向。它的思路是先用一个小模型草稿模型快速生成多个候选Token再用大模型一次性验证这些Token。验证通过的直接接收验证失败的再退回纠正。在这个机制里KV Cache扮演的角色很微妙。大模型验证草稿Token时历史和草稿候选Token的K/V会被缓存下来这样即使某个Token被否决、需要回退重新生成缓存的K/V也能节省大量重复计算。投机解码适合那些“小模型快到飞起、大模型慢但质量高”的组合。例如在公司内部做一个代码生成服务用1B级别的草稿模型先写出一段候选代码片段7B主模型负责审核每一行代码整体生成速度可能比直接硬跑7B快2到3倍。不过这个方案对草稿模型的质量要求不低如果草稿模型经常被否决缓存频繁回滚反而会抵消加速收益。6. 实战中绕不开的坑OOM、缓存开关与长文本退化6.1 显存OOM最典型的排查链路KV Cache相关的OOM和模型权重OOM不同。权重OOM一般发生在加载模型时而KV Cache OOM发生在生成过程中让服务跑着跑着突然崩溃。我在实际调试中遇到过一次很典型的场景一个7B模型服务单条请求上下文2000 Token并发一高就报CUDA OOM。排查步骤如下先确认模型权重占用FP16下7B约14GB一张24GB显卡还剩约10GB再扣除运行时开销CUDA context、激活值、临时缓冲区大概吃2GB到3GB查看剩余可用的KV Cache池按公式算2000 Token单请求KV Cache约1GB理论上支持8到10个并发业务方把max_model_len设成了32768导致每个请求Slot按32K Token预留显存单个Slot就能吃掉16GB并发一多立刻爆炸。最终解决方案很简单把max_model_len改回到8192并把vLLM的gpu_memory_utilization从0.92微调到0.88留出缓冲空间服务马上稳定。这个案例说明C端消费者的感知是“突然OOM了”但工程上几乎都是“KV Cache预留策略和实际业务负载不匹配”导致的。遇到OOM别急着换显卡先把max_model_len、batch大小和并发上限捋一遍。6.2 use_cache开关带来的行为差异有些集成LangChain或自研Agent框架的人会自己写推理循环。这时要注意如果你手动实现“拼接历史Token并重新喂给模型”而不使用past_key_values那么即使模型支持KV Cache你的推理方式也是“事实性禁用缓存”。这种情况通常发生在两类代码里第一类把历史消息全部拼成一个大字符串每次调用generate时重新传入。这会让模型每次都重新计算所有历史Token的K/V序列越长越慢第二类自己管理多轮对话但在第二轮调用时忘了把上一轮的past_key_values传进去。排查方法也简单打日志确认generate时的input_ids长度。如果第二轮调用input_ids还是从第一句开始的全量文本说明你白跑了缓存。正确的做法是第二轮只传新消息并带上past_key_values。这个环节有个反直觉的地方Transformers的tokenizer会处理整段对话历史方便你构造prompt但有历史文本在“语言层”合起来和“缓存层”复用是两码事。你仍然可以传全量text给tokenizer但传给模型时要用好past_key_values避免重复计算。6.3 长文本生成时的KV Cache准确性与质量权衡KV Cache不是万能的。当上下文很长尤其是超过模型训练时的常见长度时KV Cache可能会累积精度误差或者模型注意力开始分散生成质量明显下降。这种情况下有人会尝试把KV Cache清空重算——也就是把use_cache部分关闭强制模型重新编码整个历史。这样能解决一部分“长文本中途生成跑偏”的问题但代价是速度骤降。实操中我更推荐用下面的组合策略对长文档做分段处理而不是一股脑塞进上下文在关键转折点例如长文档对话开启新话题时手动重置KV Cache如果必须维持极长上下文优先考虑支持长上下文的结构比如带RoPE外推或YaRN扩展的模型这些模型在长序列下对KV Cache的退化更鲁棒定期监控生成后的连贯性指标出现退化时自动触发“重算”策略。还有一点关于多轮对话的细节随着对话轮次增加KV Cache不断增长。如果某些历史内容已经不在业务关心的范围内但你还把它保留在KV Cache里它就会持续占用显存并可能影响注意力分配。此时需要实现“缓存修剪”或“上下文压缩”把不重要的历史Token从KV Cache里裁剪掉。这件事并不容易因为它涉及“哪些历史Token该被剪掉”的判断。如果你不想做太复杂的逻辑也可以直接设置对话轮次上限到达上限就把历史压缩成摘要然后重建KV Cache。6.4 KV Cache与并发上限算清楚你的服务能扛多少想估算一个推理服务能支撑多少并发手头算清楚KV Cache的账就够了。简易估算公式如下可支持并发数 可用显存 - 模型权重显存 - 运行开销÷ 单请求KV Cache显存以A100 40GB、7B模型FP16为例模型权重约14GBCUDA和运行缓冲约3GB可用显存约23GB假设单请求上下文2048 Token单请求KV Cache约1GB理论并发约23个。如果再上FP8量化或INT8 KV Cache单请求KV Cache降到0.5GB左右理论并发可以提升到46个左右。当然实际并发还会受算力、内存带宽、生成速度等限制但显存维度的上限就是这么算的。这个估算方法在选卡、评估框架配置、规划服务容量时都非常有用。建议你做一个自己的KV Cache显存计算小工具把模型结构参数、精度、上下文长度、batch填入几秒钟就能看出一个配置合不合理。7. 写在最后的个人经验KV Cache这个概念单独拎出来解释并不复杂但它牵涉的工程层面特别广从模型结构GQA、推理框架vLLM、显存管理PagedAttention、数值精度量化到调度策略Continuous Batching全都跟它挂钩。我在多个项目里感觉到真正制约一个LLM服务吞吐上限的往往不是模型权重而是KV Cache的管理方式。模型权重是静态的加载一次就固定了KV Cache是动态的随请求不断增长、释放、复用。把这部分的机制吃透对推理服务的延迟和吞吐优化会有质的帮助。如果你现在正准备部署一个LLM服务我的建议是先把默认配置跑通一版记下显存占用和延迟然后动手改max_model_len、gpu_memory_utilization观察KV Cache变化再试着开启KV Cache量化或换用GQA结构的模型记录对比数据。这样一轮“配置—测量—归因—调整”走下来你对KV Cache的掌控会比看十篇文章都扎实。
返回列表