返回博客
·算法与原理

KV Cache 到底是什么?面试被问住之后我翻了源码

Transformer 推理优化的基础,但很多人只背了八股文。从源码讲清楚 KV Cache 的原理、显存开销、PagedAttention,以及什么时候不用 KV Cache。

#Transformer#KV Cache#PagedAttention#推理优化

# KV Cache 到底是什么?面试被问住之后我翻了源码

昨天面试一个后辈,问他 Transformer 推理的时候有没有什么优化手段。对方说了个"KV Cache",然后就没下文了。

我问:"能说说它为什么能加速吗?" 他沉默了三秒。

我自己也想过,很多人知道 KV Cache 这三个字,但真要讲清楚,十个人里八个说不明白。今天来把这个东西掰开揉碎讲一讲。

先搞清楚:没有 KV Cache 的时候发生了什么

Transformer 推理是逐个 token 生成的。每生成一个新 token,模型都要重新跑一遍整个序列的 Attention 计算。

假设有 1000 个 token 的输入,已经生成了 50 个输出 token。生成第 51 个 token 时,如果没有 KV Cache,模型会重新计算全部 1050 个 token 的 QKV。生成第 100 个 token 时,要重新算 1100 个。

越往后越贵。这个复杂度是 O(n²) 的——你每多生成一个 token,Attention 的计算量就线性增长。

KV Cache 的本质:别算重复的

KV Cache 的核心思想很简单:**Key 和 Value 矩阵是历史状态的快照,不需要重算。**

每个 Transformer 层维护两组缓存:

  • K cache:历史所有 token 的 Key 向量,形状 `[seq_len, head_num, head_dim]`
  • V cache:历史所有 token 的 Value 向量,同上
  • 生成新 token 时,只需要:

  • 计算当前 token 的 Q、K、V
  • 2. 把新的 K、V 追加到缓存末尾

    3. 用 Q 和完整的 K_cache、V_cache 做 Attention

    这样每次推理的计算量就从 O(n²) 降到 O(n)——因为当前 token 的 Q 只需要和已有缓存做乘法,不需要重新处理整个历史序列。

    代码说话

    用一个最小化的 PyTorch 实现来看:

    import torch

    import torch.nn as nn

    import torch.nn.functional as F

    class SimpleKVCache:

    def __init__(self, num_heads, head_dim, max_len):

    self.k_cache = torch.zeros(1, num_heads, max_len, head_dim)

    self.v_cache = torch.zeros(1, num_heads, max_len, head_dim)

    self.used_len = 0

    def update(self, k, v):

    """k, v shape: (batch, heads, seq_len, head_dim)"""

    seq_len = k.shape[2]

    start = self.used_len

    end = start + seq_len

    self.k_cache[:, :, start:end] = k

    self.v_cache[:, :, start:end] = v

    self.used_len += seq_len

    return self.k_cache[:, :, :self.used_len], self.v_cache[:, :, :self.used_len]

    def attention_with_kv_cache(q, k, v, kv_cache, num_heads, head_dim):

    k_cached, v_cached = kv_cache.update(k, v)

    # Q @ K^T / sqrt(d)

    scores = torch.matmul(q, k_cached.transpose(-2, -1)) / (head_dim ** 0.5)

    attn = F.softmax(scores, dim=-1)

    out = torch.matmul(attn, v_cached)

    return out

    注意 update 方法——这里用的是原地赋值,不是拼接。这是 vLLM 和 llama.cpp 都用的做法:预分配好 buffer,写的时候直接覆盖。拼接的开销比你想的大,尤其是 CUDA 显存操作。

    坑:KV Cache 不是免费的

    很多人以为用了 KV Cache 就万事大吉,实际上有几个要命的细节:

    1. 显存占用

    KV Cache 的显存大小 = num_layers × seq_len × num_heads × head_dim × 2 × 2(float16 占 2 字节,K 和 V 各一份)。

    以 LLaMA-7B 为例:

  • 32 层 × 32 个 head × 128 维 × 2(K+V)× 2 字节 ≈ **512 字节/token/层**
  • 总显存 ≈ 32 层 × 512 字节 × seq_len
  • seq_len=4096 时,光是 KV Cache 就要占 **64MB**。seq_len=32768 时,**512MB**。模型推理占几 GB,KV Cache 也能轻松吃掉几 GB。

    2. PagedAttention:vLLM 的真正杀手锏

    传统 KV Cache 分配方式是连续内存,容易导致显存碎片。vLLM 引入 PagedAttention——把 KV Cache 按 page 分块管理,类似操作系统的虚拟内存分页。

    # 伪代码:vLLM 的 KV Cache 管理

    class PagedAttentionCache:

    def __init__(self, num_layers, num_blocks, block_size, head_num, head_dim):

    self.pages = {} # block_id -> [layer_id] -> page_buffer

    self.free_blocks = list(range(num_blocks))

    self.allocations = {} # req_id -> [block_ids]

    这个设计让 vLLM 的吞吐比 naive 实现高 2-4 倍。不是理论上的,是实测的。

    3. 量化 KV Cache

    INT8 量化 KV Cache 在 llama.cpp 里是标配。从 float16 到 int8,显存直接减半。质量损失呢?实测下来对大多数场景几乎没有感知影响——Attention 的 softmax 本身有容错性,KV 的精度要求没那么高。

    // llama.cpp 量化 KV Cache 的核心思路

    void kv_cache_quantize(struct ggml_context *ctx, struct ggml_tensor *k, struct ggml_tensor *v) {

    // float16 -> int8, 记录 scale 值

    // 计算时反量化回 float16 再做 attention

    }

    什么时候不用 KV Cache?

    说实话,这个场景很少,但存在:

  • 离线 batch 推理:所有输入一次性算完,不需要逐个生成
  • 短序列训练:训练时一般用 Flash Attention 2,不走 KV Cache 路径
  • 超长上下文 + 内存受限:这时候 KV Cache 本身就成了瓶颈,得考虑用 StreamingLLM 这类方案
  • 写在最后

    KV Cache 是 LLM 推理优化的基础,但不是什么银弹。真正的高性能推理服务(vLLM、TGI、llama.cpp)是 KV Cache + PagedAttention + 量化 + CUDA Graph 的一堆组合拳。

    面试的时候如果对方只说了"KV Cache"三个字,可以追问:它的显存开销是多少?PagedAttention 解决了什么?INT8 量化后的精度下降大概多少?

    这三个问题,能筛掉大部分只背过八股的人。


    *我是做 AI 基础设施的,踩过不少坑。如果觉得有用,可以关注后续的推理优化系列。*