跳到正文

KV-Cache 空间复杂度怎么算?

自回归推理中 KV-Cache 显存占用分析及优化策略

原题:在大语言模型的自回归推理过程中,KV-Cache(键值缓存)被用来加速解码。请分析其空间复杂度随序列长度和模型层数的变化关系,并讨论其对显存占用的影响及可能的优化策略。

推理优化 · 字节真题

回答与解析

空间复杂度

对decoder self-attention,若K/V缓存形状按 [B, n_kv_heads, L, d_head] 计算,总字节数近似为:

2 × B × L × n_layers × n_kv_heads × d_head × bytes_per_element

前面的2代表K和V。因而空间对batch、已缓存序列长度、层数、KV头数和头维都是线性增长。必须使用 n_kv_heads,不能一律代入query头数;MQA/GQA正是通过减少KV头数降低缓存常数。

可核验示例

Llama-2-7B有32层、32个KV头、头维128。batch=1、长度4096、FP16时,缓存约为 2×1×4096×32×32×128×2 bytes = 2 GiB。这并未超过约7B参数的FP16权重,后者约14 GB十进制量级。Llama-2论文中7B/13B使用MHA,70B使用GQA,不能说全系列都是GQA。

优化边界

  • MQA/GQA:减少KV头数;质量、带宽和并行效果要按模型验证。
  • KV量化:降低每元素字节数,但需评估长上下文误差和反量化开销。
  • PagedAttention:按块管理缓存,主要减少动态请求的内存碎片并提高可调度性,不改变单条请求所需逻辑KV数量。
  • 滑动窗口/稀疏注意力:给缓存设上限,但会改变可见上下文。
  • 前缀共享与offload:分别减少重复前缀副本、以传输带宽换显存。

O(L/k)在固定k下仍是O(L);上述手段多是降低常数、碎片或可见窗口,不能随意宣称改变大O。

口语版讲法(约4分钟)

  • 从缓存张量形状推导公式
  • 用Llama-2-7B做可复算示例
  • 澄清Llama 2不同规格的注意力类型
  • 分类讨论五种优化
  • 区分常数优化与复杂度变化

我会先从张量形状推导,而不是背一个固定显存数字。自回归解码时,每一层都要保存历史token的K和V。单个缓存通常可以看成B、KV头数、序列长度L、头维Dh,所以总字节数近似是二乘B乘L乘层数乘KV头数乘Dh,再乘每个元素的字节数。前面的二代表K和V两份。

这里最容易错的是把query头数一律代入公式。普通多头注意力里,KV头数等于query头数;MQA只有很少的KV头,GQA介于两者之间。正因为缓存使用的是KV头数,MQA和GQA才能明显降低缓存常数。空间对batch、已缓存长度和层数都是线性增长,所以长上下文和并发请求会一起放大显存压力。

用Llama-2-7B举一个可复算的例子。它有32层、32个KV头、头维128。batch等于1,缓存长度4096,FP16每元素2字节,代入公式得到2 GiB。这个数字本身没问题,但不能说它超过模型权重。大约70亿参数的FP16权重在14 GB十进制量级,仍明显更大。还要注意,Llama-2论文里7B和13B使用MHA,70B才使用GQA,不能把全系列都写成GQA。

优化可以分几类。第一类是模型架构层的MQA或GQA,直接减少KV头数。第二类是KV量化,用更少bit存储,但要评估反量化开销和长上下文质量。第三类是PagedAttention,它按块管理动态请求,主要解决预留和碎片问题,提高调度弹性;它不会凭空减少一条请求逻辑上需要保存的全部KV。

第四类是滑动窗口或稀疏注意力,只保留允许访问的一部分历史,从而给缓存设上限,但代价是模型不再看到被丢弃的远程token。第五类是前缀共享和offload:共享系统提示等公共前缀能避免重复副本,offload则用CPU或其他层级的容量换取传输带宽和延迟。

工程上我会按模型配置和实际shape计算,而不是引用固定的百分比。PagedAttention把利用率从多少提升到多少、batch一定翻几倍,都必须来自同一硬件和工作负载的测量。最后还要区分大O和常数:固定分组数下O(L除以k)仍然是O(L)。大多数优化是在降低KV头数、位宽、碎片或可见窗口,而不是把线性复杂度变成另一种阶。

估算服务总容量时还要区分prompt长度、已生成长度和最大预留长度。静态预留会产生浪费,连续批处理又会让不同请求不断进入和退出。因而容量模型最好直接读取模型配置中的层数、KV头数和头维,再结合请求长度分布计算分位数,而不是只按最大上下文乘并发数。

关键一句:PagedAttention减少的是内存碎片和预留浪费,不是单请求的逻辑KV条目。

核验来源

  1. Llama 2: Open Foundation and Fine-Tuned Chat Models
  2. Efficient Memory Management for Large Language Model Serving with PagedAttention

面试官还可能这样问

  1. 问法 1 · 场景切入

    假设你在做客服机器人的流式推理,用户说了很长一段话,模型需要逐个token生成回复。如果每次生成都重新计算前面所有token的注意力,延迟会很高。你怎么利用KV-Cache来加速?

  2. 问法 2 · 层层追问

    自回归解码时,你会怎么缓存中间结果来避免重复计算?……缓存了K和V之后,序列越长显存占用怎么变化?……那模型层数多了呢?这个开销有多大?

  3. 问法 3 · 直球架构

    分析一下KV-Cache的空间复杂度,跟序列长度、模型层数、batch size的关系是怎样的?它对显存有什么影响?有哪些常见的优化手段?

同模块相关题目