跳转至

KV Cache 深度解析:从原理到多轮对话优化

什么是 KV Cache?它为什么能够加速自回归生成?

在标准的自回归生成(如GPT)中,每次预测下一个token都需要将整个序列输入模型,重新计算所有位置的Key和Value。对于长度为 LL 的序列,第 L+1 步需要重新计算前面 L 个token的注意力,这导致了 O(L2) 的重复计算。KV Cache 正是为了消除这种冗余而提出:在生成过程中,将之前所有时间步的Key和Value张量缓存起来,每个新时间步只需计算当前token的Query、Key、Value,然后Query与缓存中的历史Key计算注意力分数,新的Key和Value追加到缓存中供后续步骤使用。

加速的核心在于:避免了历史token的重复计算。每步生成的计算复杂度从 O(L2⋅d)O(L2⋅d) 降低到 O(L⋅d)O(Ld),因为自注意力的计算量主要消耗在Query与所有Key的点积上。在长序列生成时,KV Cache带来的速度提升极为显著。

推导 KV Cache 的显存占用公式

对于一层Transformer,缓存需要存储每一层的Key和Value。假设:

  • 批次大小:B

  • 序列长度(生成过程中的当前长度):L(注意KV Cache大小随生成步数线性增加)

  • 层数:N

  • 注意力头数:H

image.png

在长序列生成中,KV Cache 成为显存瓶颈,分析其根本原因

从公式可以看出,KV Cache的显存与序列长度 L 成正比。对于长文本生成或大量并发请求,L 可达数万乃至数十万,显存占用线性增长。更关键的是,在自回归解码中,每个时间步都要将完整的KV Cache驻留在显存中,并且通常是连续且无法被换出的。当批处理多个不同长度的请求时,碎片化更加严重,大量显存被KV缓存占据,限制了最大批大小和序列长度。这就是为什么即使有80GB显存的GPU也经常不够用的根本原因。

Multi-Query Attention (MQA) 是如何减少 KV Cache 大小的?它共享了什么?

MQA 让所有注意力头共享同一组 Key 和 Value,而Query仍然保持多头。在标准MHA中,每个头有独立的K和V,因此KV缓存大小与头数 HH 成正比。MQA则将所有头的K和V合并为一组,从而将KV缓存的大小降低到原来的 1/H1/H。同时,参数量也相应减少。例如,LLaMA-7B H=32H=32,改用MQA后KV缓存仅为原来的1/32。这显著降低了显存占用和内存带宽需求,但可能会损失一些生成质量,因为不同的头被迫使用相同的注意力模式,导致表达能力受限。

Grouped-Query Attention (GQA) 是如何在 MQA 和 MHA 之间折中的?分组数如何选择?

GQA 将注意力头分成 GG 个组,每个组共享一组K和V,而Query仍是每个头独立。这样,KV缓存大小变为原来的 G/HG/H(假设每组头数相同)。当 G=1G=1 时即MQA,当 G=HG=H 时即MHA。通过调整 GG,可以在速度和模型质量之间取得平衡。

分组数的选择取决于对推理效率和生成质量的权衡。LLaMA 2 70B 使用 G=8G=8(H=64H=64),这样每个KV组服务于8个头,在几乎不损害下游任务性能的情况下,将KV缓存减少为原来的1/8。通常,模型越大,适当的分组数(如8)带来的性能损失越小,而推理加速明显。实验表明,GQA在多数任务上与MHA质量相当,但推理速度显著提升。

比较 MQA、GQA 和 MHA 在参数量、推理速度和生成质量上的差异

  • 参数量:MHA > GQA > MQA。MHA的K、V投影权重最大,MQA最小,GQA居中。但参数总量的差异通常不大(相对于FFN和Q投影)。

  • 推理速度:MQA最快,GQA次之,MHA最慢。因为KV缓存读取量不同,MQA内存带宽需求最小,可以支持更大的批大小或更长的序列。

  • 生成质量:通常MHA最高,GQA非常接近,MQA在某些任务(如长文本生成)中质量下降较明显。但通过扩大其他维度或加强训练,GQA的质量可几乎无损。因此,目前大型生产模型普遍采用GQA(如LLaMA 2、Mistral、Gemini),以在效率和效果间取得最佳平衡。

Multi-head Latent Attention (MLA) 的原理是什么?它如何通过低秩压缩进一步减少 KV 缓存?

MLA 出自DeepSeek-V2,旨在进一步压缩KV缓存。传统方法中,K和V的投影矩阵将 dmodeldmodel 映射到 H×dkH×dk。MLA提出:对Key和Value分别使用低秩分解,即先将输入投影到一个低维潜在空间(压缩),再从这个潜在空间展开到各个头的K和V。同时,在推理时,只需要缓存潜在空间的压缩表示(而非每个头的完整K和V),然后在计算注意力时动态解压。这相当于用计算换取存储:额外增加一次小矩阵乘法,但缓存大小由与头数相关的量减少为仅与潜在维度相关的量(通常远小于 H×dk)。MLA能够将KV缓存大小减少到原来的1/10甚至更多,对极长上下文推理极为有利。

什么是 PagedAttention?vLLM 如何借鉴操作系统的分页管理来优化 KV 缓存?

PagedAttention 由 vLLM 提出,它将 KV 缓存从传统的连续张量存储改为分页(block)存储。类似操作系统中的虚拟内存分页,PagedAttention 将每个序列的 KV 缓存划分为固定大小的块(例如每个块包含16个token的K和V),这些块不必在物理显存中连续。每个序列维护一个块表,逻辑位置到物理块的映射。这样:

  • 解决碎片化:不再需要为每个序列预分配最大长度的连续显存,按需分配块,显存利用率显著提高。

  • 共享:不同序列可以共享相同的KV块(例如,多个请求具有相同的系统提示前缀时,可共享前缀的KV块)。

  • 高效调度:块可以动态分配和回收,支持更大的并发批处理。

PagedAttention 如何处理 KV 缓存碎片化问题?它的块大小如何影响效率?

传统KV缓存要求每个序列预留连续显存,当多个序列长度不同、动态分配和释放时,会产生严重的显存碎片。PagedAttention 通过固定大小的块消除了连续分配需求,每个块可以独立分配,没有外部碎片。当序列变长时,只需分配新的块并追加到块表中,无需整体搬迁。序列结束后,整个块表对应的块被释放,不会遗留碎片。

块大小的影响:块越大,块表更小,逻辑映射开销低,但可能浪费显存(因为块内部分未使用)。块越小,显存利用更精细,但块表变大,查询块表开销增加,且可能增加内存访问延迟。通常取16或32作为折中。

Prefix Caching 的原理是什么?在什么场景下可以显著减少计算?

Prefix Caching 利用共享前缀的特性。在很多应用中,多个请求可能具有相同的系统提示或上下文开头。通过将前缀的KV缓存以块或序列形式缓存起来,当新请求到来时,可以直接复用这些已计算的KV缓存,而无需重新进行前向计算。PagedAttention 通过共享物理块实现这一功能:对于相同前缀,多个序列的块表指向同一组物理块,并采用引用计数管理。当没有序列再使用这些块时,才释放显存。

显著有效的场景:多轮对话中用户不同的提问但共享系统提示;批量推理中多个请求以相同长文档作为背景知识;以及其他存在共同前缀的并发请求,可节省大量计算和显存。

能否对 KV Cache 进行量化?如果可以,会面临哪些挑战?

可以对 KV Cache 进行量化,将FP16的K和V压缩为INT8甚至更低精度。例如,LLM推理中常用的INT8 KV Cache。主要挑战:

  • 精度损失:注意力计算对量化噪声敏感,尤其是当缓存值范围较大时,直接均匀量化可能导致注意力分布严重失真。需要逐token或逐通道的量化策略,并可能需要在模型训练后加入量化校准。

  • 硬件支持:现有GPU对低精度矩阵乘加的支持有限,量化后需要反量化到FP16再进行注意力计算,可能带来额外开销,甚至抵消带宽节省。需要定制CUDA kernel来直接计算量化注意力。

  • 动态范围:生成过程中,新token的K/V分布可能与历史不同,需要动态调整量化参数。

尽管如此,W8A8(FP8)等量化方案已取得不错效果,将KV缓存减半,而精度损失在可接受范围内。

KV Cache 在多轮对话中是如何复用的?历史轮次缓存的管理策略有哪些?

在多轮对话中,每一轮都会有新的用户输入和助手回复,序列不断增长。如果简单地将所有历史轮次作为一个长序列,每次新对话都要重新编码整个历史,计算量会不断增加。复用策略:

  • 缓存整个历史序列的KV:最直接的方法,一直保留,生成新回复时只需追加新token的KV。这样历史信息完全不丢失,但缓存占用随对话轮次线性增长。

  • 滑动窗口保留:只保留最近 W 个token的KV缓存,旧token丢弃。节省显存,但可能丢失早期重要信息。

  • 重要token保留 + 压缩:识别重要token(如用户指令、关键实体)并将其KV保留,其余丢弃或压缩为摘要向量。

  • 前缀复用:不同对话轮次共享系统提示的KV缓存,只缓存变化部分。

现代框架(如vLLM)通过块管理可以自动处理多轮对话的KV缓存,根据显存限制决定是否释放旧块。

不使用 KV Cache 的自回归推理计算量有多大?请以 L 序列为例计算。

image.png

滑动窗口注意力如何配合 KV Cache 实现长度限制?

滑动窗口注意力限制每个token只能关注其前后固定窗口大小 W 内的token。配合KV Cache时,缓存只需保留最近 WW 个token的K和V,超出窗口的历史token可以被丢弃,从而将KV Cache的显存占用上界固定在 W,与总序列长度 L 无关。这从根本上避免了长序列下的线性增长问题。

实现上,每步生成时,将当前token的K和V存入缓存;如果缓存长度超过 W,则删除最旧的token的KV。因为注意力掩码会确保当前token只计算与缓存中这些token的注意力,所以计算效率和显存都得到控制。典型的窗口大小 W 为4096或8192,足以捕获足够的局部上下文,同时保持显存可预测。