ARTICLE DETAIL

资讯详情

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

KV Cache 显存优化实战:从注意力原理到 PagedAttention 与 KV_Catch 观测

KV Cache 显存优化实战:从注意力原理到 PagedAttention 与 KV_Catch 观测 KV Cache 这个话题只要你在做 LLM 推理相关的工作早晚都得正面碰上。我第一次认真盯它是因为线上一个 7B 模型的服务在并发上来之后显存直接爆掉日志里全是 OOM但 GPU 利用率却低得可怜。排查了一圈才发现问题不在模型本身而在 KV Cache 的分配和回收策略上——请求排队时预分配的缓存块没被及时释放长序列请求又把缓存池撑满短请求只能干等。这件事之后我花了大概两周时间把 KV Cache 从注意力机制的数学原理到 vLLM 的 PagedAttention 实现完整捋了一遍顺手写了个小工具叫 KV_Catch专门用来观测和抓取 KV Cache 在运行时的分配、复用、碎片情况。这篇就把我踩过的坑、验证过的结论、以及 KV_Catch 的设计思路完整摊开讲一遍。1. 先把 KV Cache 到底缓存了什么说清楚1.1 从自回归解码的重复计算说起Transformer 解码器在生成 token 时是自回归的每生成一个新 token都要拿当前 token 的表示去和前面所有 token 做注意力计算。如果每一轮都把前面所有 token 的 Key 和 Value 重新算一遍那计算量会随序列长度平方级增长这在工程上完全不可接受。KV Cache 的核心思路非常朴素已经算过的 Key 和 Value 矩阵直接存下来复用不重复计算。第 t 步解码时只需要计算当前这个新 token 的 Q、K、V然后把新的 K、V 追加到缓存里用当前 Q 去和缓存中全部的 K 做点积再对全部的 V 加权求和。这样每一步的计算量从 O(t²) 降到 O(t)代价是显存里多存了一份历史 K、V。这里有个容易被忽略的点KV Cache 缓存的是每一层、每一个注意力头的 K 和 V不是只缓存最后一层。一个 L 层的模型每层都有自己的 KV Cache所以显存占用是随层数线性叠加的。很多人第一次估算显存时只算了一层结果实际占用差了十几倍。1.2 显存占用的精确计算公式KV Cache 的显存占用可以用一个很直接的公式算出来KV Cache 字节数 2 × L × H × S × D × P × B其中各符号含义如下符号含义典型值以 7B 模型为例2K 和 V 两份固定L层数32HKV 头数注意不是 Q 头数32MHA或 8GQAS序列长度2048D每头维度128P精度字节数2FP16B批大小1按这个公式7B 模型在 MHA、FP16、序列长度 2048、批大小 1 的情况下KV Cache 占用约为 2 × 32 × 32 × 2048 × 128 × 2 1.07 GB。注意这只是单条请求如果并发 16 条就是 17 GB 以上再加上模型权重本身约 14 GB一张 24 GB 的卡基本就满了。提示GQA分组查询注意力之所以能大幅降低显存就是因为把 H 从 32 降到了 8KV Cache 直接缩小到四分之一。这也是为什么现在主流模型几乎都默认用 GQA。1.3 为什么 KV Cache 是推理服务的头号显存杀手模型权重是静态的加载完就固定了。但 KV Cache 是动态的它随并发数、序列长度实时变化。一个推理服务的显存瓶颈十有八九不是权重而是 KV Cache。更麻烦的是KV Cache 的分配是按请求的而请求的长度事先并不知道。传统做法是给每个请求预分配一个最大长度的连续显存块这就导致两个问题一是短请求浪费了大量预留空间内部碎片二是长请求可能因为找不到足够大的连续块而无法调度外部碎片。vLLM 的 PagedAttention 就是为了解决这个问题把 KV Cache 切成固定大小的 block像操作系统管理虚拟内存一样按需分配碎片率能从 60% 以上降到 4% 以下。2. KV_Catch 想解决的问题和它的观测维度2.1 为什么现成的工具不够用我一开始用的是 nvidia-smi 和 PyTorch 的显存统计但这两个粒度都太粗。nvidia-smi 只能看到整卡显存分不清哪部分是权重、哪部分是 KV Cache、哪部分是临时激活。PyTorch 的torch.cuda.memory_allocated能看到分配量但看不到 KV Cache 内部的 block 使用情况、复用命中率、碎片分布。vLLM 本身有 metrics 接口能暴露gpu_cache_usage_perc这类指标但它是聚合值看不到单个请求的缓存生命周期。我想知道的是一个请求从进入到结束它的 KV Cache 是怎么被分配的、中间有没有被抢占、block 有没有被复用、释放后有没有产生碎片。这些信息对于调优调度策略和排查 OOM 至关重要但现成工具给不了。KV_Catch 就是在这个背景下写的。它的定位很明确在运行时抓取 KV Cache 的分配、复用、释放全链路事件并做可视化聚合不替代 vLLM 的调度器而是作为旁路观测层存在。2.2 抓取哪几类关键事件KV_Catch 目前抓取的事件分四类每一类对应一个具体的排查场景allocate 事件记录请求 ID、请求的 block 数量、分配到的物理 block 编号、分配耗时。用于分析分配延迟和 block 分布。append 事件记录每次解码步新增的 token 数、对应的 block 是否跨块。用于分析序列增长模式和跨块频率。reuse 事件记录命中的前缀缓存prefix cacheblock 数、复用来源请求 ID。用于评估前缀缓存的收益。free 事件记录释放的 block 数、释放后是否合并到空闲池、空闲池碎片状态。用于分析碎片产生和回收效率。这四类事件串起来就是一个请求完整的 KV Cache 生命周期。我在实际排查中发现大部分 OOM 不是真的显存不够而是 free 事件没有及时触发或者 reuse 命中率太低导致重复分配。2.3 事件采集的开销控制旁路观测最大的风险是拖慢主流程。KV_Catch 在采集层做了三件事来控制开销第一事件写入走无锁环形缓冲区采集线程只做 memcpy不做任何格式化或 IO。第二采样率可配置高并发场景下可以只采集 1% 的请求做全链路追踪其余请求只记聚合计数。第三聚合在独立线程完成主推理线程完全不感知。实测下来在 100 QPS 的场景下全量采集带来的额外延迟在 0.3ms 以内对首 token 延迟TTFT的影响可以忽略。如果采样率降到 1%开销基本测不出来。3. 从注意力数学到 PagedAttention 的实现链路3.1 标准注意力的 KV 计算过程要理解 KV Cache 的存储布局得先回到注意力的计算本身。标准多头注意力里输入 X 经过三个线性投影得到 Q、K、V# 简化版忽略 batch 和 head 维度 Q X W_q # [seq_len, d_model] K X W_k # [seq_len, d_model] V X W_v # [seq_len, d_model] # 拆成多头后 # Q, K, V: [num_heads, seq_len, head_dim] scores Q K.transpose(-2, -1) / sqrt(head_dim) attn softmax(scores) output attn VKV Cache 缓存的就是这里的 K 和 V。在解码阶段X 只有当前一个 token所以新算出的 K、V 形状是[num_heads, 1, head_dim]需要和缓存里的历史 K、V 拼接。3.2 连续缓存布局的致命缺陷最直观的缓存布局是给每个请求分配一块连续的显存形状为[num_layers, 2, num_heads, max_seq_len, head_dim]。这种布局实现简单但有两个硬伤。一是预留浪费。max_seq_len 通常按模型上限设比如 8192但实际请求平均长度可能只有 500浪费率超过 90%。二是无法共享。多个请求如果有相同的前缀比如相同的 system prompt它们的 KV 是完全一样的但连续布局下每个请求各存一份无法复用。这两个问题在并发一高就暴露无遗。我做过一个测试16 并发、平均长度 800、system prompt 长度 200 的场景下连续布局的实际有效缓存利用率只有 11%其余全是预留和重复。3.3 PagedAttention 的 block 化管理PagedAttention 的思路借鉴了操作系统的虚拟内存分页。它把 KV Cache 切成固定大小的 block每个 block 存固定数量 token 的 K、VvLLM 默认 block_size 是 16。每个请求维护一张 block table记录逻辑 block 到物理 block 的映射。这样做的好处很直接按需分配请求增长到需要新 block 时才分配不预留。前缀共享相同前缀的请求可以指向同一批物理 block通过引用计数管理写时复制。碎片可控block 大小固定空闲池管理简单碎片率极低。block table 本质上是一个页表注意力计算时通过它把逻辑上连续的 KV 映射到物理上可能离散的 block。这也是为什么 vLLM 的 attention kernel 和标准实现不一样它需要先做一次 block 索引的 gather。3.4 block 大小对性能的实际影响block_size 是个需要权衡的参数。太小block table 变长索引开销上升太大内部碎片增加前缀共享的粒度变粗。我实测过 block_size 从 8 到 64 的表现结论是16 在大多数场景下是甜点。block_size8 时block table 索引开销让解码吞吐下降约 4%block_size64 时短请求的内部碎片让有效缓存利用率下降约 7%。16 在两者之间取得了比较好的平衡。当然如果你的请求长度分布特别集中比如都是 4096 左右那调大 block_size 反而更划算。4. 用 KV_Catch 定位三类典型线上问题4.1 问题一并发上不去显存却先满了这是我最常遇到的场景。表现是 QPS 卡在某个值上不去nvidia-smi 显示显存接近 100%但 GPU 利用率只有 30% 左右。用 KV_Catch 抓一段时间的 allocate 和 free 事件画成时间线就能看出问题。我遇到的一次是free 事件的平均延迟达到了 800ms也就是说请求结束后它的 block 要等 800ms 才真正回到空闲池。这 800ms 里新请求无法使用这些 block只能排队或触发抢占。根因是释放逻辑里有一个同步操作在释放前要等一个统计上报完成。把上报改成异步之后free 延迟降到 5ms 以内同样的显存下并发能力提升了近 3 倍。排查这类问题的关键是不要只看显存总量要看 block 的周转率。显存满不代表 block 都在用可能是大量 block 卡在已释放但未回收的中间态。4.2 问题二前缀缓存命中率低得离谱前缀缓存是省显存的大杀器但前提是命中率要够高。我见过一个服务明明所有请求都带同一个 500 token 的 system prompt但前缀缓存命中率只有 3%。用 KV_Catch 的 reuse 事件一查就明白了请求的 system prompt 虽然文本相同但 tokenize 之后有细微差异因为不同请求在 prompt 末尾多了一个空格或换行导致 token 序列不完全一致前缀匹配在第一个不同 token 处就断了。这类问题的排查思路是把 reuse 事件里匹配长度的分布画出来。如果大量请求的匹配长度是 0 或个位数那基本就是前缀不一致如果匹配长度集中在某个值附近那可能是 block 对齐的问题。修复方式也很简单在 prompt 拼接层做规范化去掉尾部空白统一换行符。改完之后命中率从 3% 涨到 87%显存占用直接降了四成。4.3 问题三长序列请求把短请求饿死这个问题的表现是短请求的 TTFT 忽高忽低长请求一来短请求就卡住。用 KV_Catch 看 block 分配的时间线能看到长请求在持续 append 新 block而空闲池被逐渐耗尽短请求的 allocate 事件开始出现等待。根因是调度策略对长请求没有做限制。vLLM 的调度器有抢占机制但抢占本身有开销频繁抢占会让整体吞吐下降。我的做法是在 KV_Catch 里加了一个长请求占比的实时指标当长请求占用的 block 超过总容量的 60% 时触发告警并临时降低长请求的调度优先级。这里有个经验长请求和短请求混部时最好给它们设置不同的优先级或配额不要让它们在同一池子里自由竞争。我试过按序列长度分池长请求单独一个池子短请求的 TTFT 稳定性提升了 5 倍以上。5. 部署和调优中那些文档不会写的细节5.1 显存预留比例不是越大越好vLLM 有个gpu_memory_utilization参数默认 0.9意思是拿 90% 的显存来做 KV Cache 池。很多人为了保险把它调到 0.95 甚至更高结果反而更容易 OOM。原因是KV Cache 池之外还需要显存来做临时激活、CUDA graph、通信缓冲。这些开销在推理过程中是动态的如果预留太少一旦某个时刻激活峰值上来就会和 KV Cache 抢显存直接 OOM。我的经验值是0.85 到 0.9 之间具体取决于模型大小和 batch 配置。模型越大激活占比越高预留要越多。5.2 量化对 KV Cache 的影响要单独评估现在很多人用 FP8 或 INT8 量化模型权重但权重量化不等于 KV Cache 量化。KV Cache 的量化需要单独开启而且对精度的影响比权重量化更敏感。我实测过 KV Cache 用 FP8 的效果显存占用减半吞吐提升约 30%但在长序列任务上输出质量有可感知的下降尤其是需要精确回忆早期上下文的场景。所以我的建议是短序列、高并发场景可以上 KV Cache 量化长序列、高精度要求的场景慎用。5.3 监控指标要盯住有效缓存利用率很多人监控只看gpu_cache_usage_perc但这个指标高不代表健康。真正该盯的是有效缓存利用率也就是实际存储有效 token 的 block 数 / 总 block 数。这个指标低说明大量 block 被预留但没存满或者被前缀缓存的引用计数占着但实际没被读。我在 KV_Catch 里把这个指标做成了实时曲线配合 block 生命周期时间线一起看基本上一眼就能定位问题。5.4 压测时要用真实长度分布用固定长度压测是最容易骗自己的做法。真实请求的长度分布通常是长尾的少数长请求会占用大量 block对调度策略的考验和固定长度完全不同。我的做法是从线上采样真实的请求长度分布然后在压测时按这个分布生成请求。这样压出来的并发上限才有参考价值。用固定长度压测得到的最大并发在真实分布下往往要打七折。6. 几个我反复验证过的结论关于 KV Cache有几个结论是我在多个项目里反复验证过的写在这里供参考。第一KV Cache 的瓶颈往往不在容量而在周转。显存够不够是一回事block 能不能快速回收复用是另一回事。后者对并发的影响通常更大。第二前缀缓存的收益高度依赖请求的相似度。如果你的请求之间前缀差异很大前缀缓存基本没用这时候不如把精力放在 block 管理和调度上。第三GQA 和 MQA 是降低 KV Cache 最有效的手段比任何量化都直接。选模型时如果对显存敏感优先选 GQA 的模型。第四观测粒度决定了排查效率。聚合指标能告诉你有问题但只有事件级的全链路追踪才能告诉你问题在哪。KV_Catch 的价值就在于此。最后分享一个我在调试时常用的小技巧把 KV_Catch 抓到的 block 分配时间线导出成 CSV用 pandas 按请求 ID 分组算每个请求的 block 持有时间和平均利用率。那些持有时间长但利用率低的请求往往就是拖慢整体调度的元凶。这个方法帮我定位过好几次隐蔽的缓存泄漏问题比看任何聚合指标都管用。
返回列表