什么是 LLM 里的 KV Cache?

这篇文章里,我们要弄懂 KV Cache——K 代表 Key(键),V 代表 Value(值)——以及它为什么用在大语言模型(LLM)里、用来加速文本生成。

我们会先看 LLM 怎么一个 token 一个 token 地生成文本,搞清楚 Key、Value、Query 在模型内部各自扮演什么角色,再通过一个例子看到重复计算的问题,最后一步步走一遍 KV Cache 怎么通过存下并复用过去的结果来解决它。

我们开始吧。

LLM 是怎么生成文本的

在理解 KV Cache 之前,我们得先搞懂 LLM 是怎么生成文本的。

LLM——大语言模型——是一种在海量文本数据上训练出来的模型。它能理解并生成人类语言。我们给它一句话,它就会预测接下来是什么。

LLM 生成文本是一个 token 一个 token 来的。token 是一小段文本——可以是一个词、一个词的一部分,甚至单个字符。为了简单起见,我们就把每个词当作一个 token。

假设我们给模型这样一段输入:

“I love”

模型看着 “I”“love”,预测出下一个 token:“teaching”

现在完整的序列变成:

“I love teaching”

模型看着 “I”“love”“teaching”,预测出下一个 token:“AI”

现在完整的序列变成:

“I love teaching AI”

这个过程一直继续,一个 token 一个 token 地往下走,直到模型决定停下来。

这里要注意的关键一点是:模型每预测一个新 token,都需要回头看前面所有的 token,才能决定接下来是什么。

模型内部发生了什么

接下来,我们看看模型生成每个 token 时,内部到底发生了什么。

模型里有一个组件叫注意力层(attention layer)。注意力层的工作,是帮模型判断前面哪些 token 对预测下一个 token 重要。前面的 token 并非个个同等有用,有些比其他的更重要。

在注意力层内部,每个 token 都被转换成三样东西。我们用一个简单的类比来理解它们。

想象一间教室,来了个新同学,他想搞清楚谁能帮自己弄懂某个特定主题:

  • Query(Q):新同学的提问——“这儿谁懂这个主题?”每个待预测的 token 都有一个 Query。它代表这个 token 在找什么。
  • Key(K):每个已有同学戴着的名牌,上面写着他们懂什么——“我懂数学”或“我懂科学”。每个之前的 token 都有一个 Key。它描述这个 token 装着什么信息。
  • Value(V):每个同学手里真正的笔记。新同学先靠名牌(Key)找到对的人,真正要用的是那份笔记(Value)。每个之前的 token 都有一个 Value。它承载着真正的信息。

所以,当前这个 token 用自己的 Query 去和前面所有 token 的 Key 做比对。这次比对会产生注意力分数——一组数字,告诉模型该给前面每个 token 多少关注。分数越高,说明那个 token 越相关;分数越低,越不相关。然后模型用这些分数,从相关的 token 那里把 Value 收集起来。

注意力层就是这么工作的。

接下来,我们看看问题出在哪。

问题所在:重复计算

模型每预测一个下一个 token,都会为序列里的所有 token 计算 Key、Value 和 Query——不只是那个新 token。

我们用例子一步步看看会发生什么:

第 1 步: 输入是 “I love”

模型为下面这些计算 Key、Value 和 Query:

  • “I”
  • “love”

用它们预测出下一个 token:“teaching”

第 2 步: 输入是 “I love teaching”

模型为下面这些计算 Key、Value 和 Query:

  • “I”(第 1 步已经算过,但又算了一遍)
  • “love”(第 1 步已经算过,但又算了一遍)
  • “teaching”(新的)

用它们预测出下一个 token:“AI”

第 3 步: 输入是 “I love teaching AI”

模型为下面这些计算 Key、Value 和 Query:

  • “I”(第 1 步和第 2 步都算过,但又算了一遍)
  • “love”(第 1 步和第 2 步都算过,但又算了一遍)
  • “teaching”(第 2 步算过,但又算了一遍)
  • “AI”(新的)

看出问题了吗?

“I” 的 Key 和 Value 在第 1 步就算过了。但模型在第 2 步又算了一遍,第 3 步再算一遍。“love”“teaching” 也是同样的情况。

模型在为那些已经见过的 token 重复做同样的活儿。 这是被白白浪费的计算。

随着序列越来越长,这个问题会越来越糟。如果模型已经生成了 100 个 token,那么下一步,它要把前面全部 100 个 token 的 Key 和 Value 重新算一遍,就为了预测一个新 token。这让文本生成变得非常慢。

解决办法:KV Cache

KV Cache 背后的想法很简单:每个 token 的 Key 和 Value 只算一次,存下来,在之后每一步里复用。

把它想象成记笔记。设想你在开会。每次有新的人发言,你不用让前面所有发言过的人把说过的话再重复一遍,只要看自己的笔记,然后专心听新发言的人就行。笔记就是你的缓存。

同样地,KV Cache 是一块内存,我们把每个已经处理过的 token 的 Key 和 Value 存在里面。这样一来,模型就不用再重新算它们了。

我们看看同一个例子用 KV Cache 是怎么走的:

第 1 步: 输入是 “I love”

模型为下面这些计算 Key、Value 和 Query:

  • “I”
  • “love”

它把 “I”“love” 的 Key 和 Value 存进 KV Cache

用它们预测出下一个 token:“teaching”

此时 KV Cache 里有: “I”、“love” 的 Key 和 Value

第 2 步: 输入只有 “teaching”(只有那个新 token)

模型从 KV Cache 里取出 “I”“love” 的 Key 和 Value。不需要重新计算。

它只为新 token “teaching” 计算 Key、Value 和 Query。

它把 “teaching” 的 Key 和 Value 存进 KV Cache。

它把这些一起用上,预测出下一个 token:“AI”

此时 KV Cache 里有: “I”、“love”、“teaching” 的 Key 和 Value

第 3 步: 输入只有 “AI”(只有那个新 token)

模型从 KV Cache 里取出 “I”“love”“teaching” 的 Key 和 Value。不需要重新计算。

它只为新 token “AI” 计算 Key、Value 和 Query。

它把 “AI” 的 Key 和 Value 存进 KV Cache。

此时 KV Cache 里有: “I”、“love”、“teaching”、“AI” 的 Key 和 Value

所以,模型不再在每一步为每个 token 重算 Key 和 Value,而是只为新 token 计算,前面所有 token 都复用缓存里的值。

KV Cache 就是这样避免重复计算的。

为什么只缓存 Key 和 Value,不缓存 Query

你自然会问:为什么我们只缓存 Key 和 Value,不缓存 Query 呢?

Query 只对当前这个 token 有用——也就是此刻正在生成的那个。当前 token 用自己的 Query 去和前面所有 token 的 Key 比对,找出哪些相关。一旦预测做完,这个 Query 就不再需要了。

但每个过去 token 的 Key 和 Value,在之后的每一步都用得着,因为每个新 token 都必须回头看前面所有 token,才能做出自己的预测。

所以,我们只需要存 Key 和 Value。这就是它叫 KV Cache 的原因——它缓存的是 Key 和 Value。

到底能快多少

我们把两种做法并排比一比:

不用 KV Cache:

Step 1: Compute K, V, Q for 2 tokens
Step 2: Compute K, V, Q for 3 tokens  (2 recomputed)
Step 3: Compute K, V, Q for 4 tokens  (3 recomputed)
Step 4: Compute K, V, Q for 5 tokens  (4 recomputed)
...
Step N: Compute K, V, Q for (N+1) tokens  (N recomputed)

计算量在每一步都在往上涨。

用了 KV Cache:

Step 1: Compute K, V, Q for 2 tokens → Save K, V for 2 tokens in cache
Step 2: Compute K, V, Q for 1 token  → Reuse K, V for 2 tokens from cache
Step 3: Compute K, V, Q for 1 token  → Reuse K, V for 3 tokens from cache
Step 4: Compute K, V, Q for 1 token  → Reuse K, V for 4 tokens from cache
...
Step N: Compute K, V, Q for 1 token  → Reuse K, V for N tokens from cache

第一步之后,模型在每一步只为一个新 token 计算,而不是整个序列。

具体感受一下:如果模型要生成一段 100 个 token 的序列,不用 KV Cache,所有步骤加起来的 Key 和 Value 计算总数会是 2 + 3 + 4 + ... + 100 = 5,049 次计算。用了 KV Cache,则是 2 + 1 + 1 + ... + 1 = 101 次计算。大约少了 50 倍。序列越长,省得越多。

这就是 KV Cache 在加速文本生成上如此有效的原因。

取舍:速度换内存

KV Cache 让生成变快了,但它带来一个取舍:它要用额外的内存,去存下到目前为止生成的每个 token 的全部 Key 和 Value 信息。

序列越长,缓存就越大。对于动辄成千上万个 token 的超长序列,缓存会吃掉相当可观的内存。

所以,KV Cache 是一个取舍:我们用更多的内存,换计算时间的节省。 对大多数场景来说,这个取舍很划算,因为速度的提升相当显著。

到这里,我们就把 LLM 里的 KV Cache 弄懂了。

下一篇博客里,我们会讲 Paged Attention,它解决的是 KV Cache 的内存问题。

更新:

https://x.com/i/status/2038190072668025297

今天就到这里。

谢谢

Amit Shekhar

Outcome School 创始人