KV缓存是生产环境中实现LLM高效推理最关键的技术之一。KV缓存是生产环境中实现计算高效LLM推理的重要组成部分。本文将通过概念讲解和从零开始、可读性强的代码实现,介绍其工作原理。

距离我上次分享讲解LLM基础概念的技术教程已经有一段时间了。由于我目前正在从伤病中恢复,同时也在撰写一篇更大型的LLM研究文章,我想借此机会分享一篇关于读者多次询问主题的教程文章(因为我的《从零开始构建大型语言模型》一书中没有包含这个主题)。

祝阅读愉快!

简而言之,KV缓存会存储中间键(K)和值(V)的计算结果,以便在推理(训练后)过程中复用,从而在生成文本时实现显著的加速。KV缓存的缺点在于:增加了代码的复杂性,提高了内存需求(这也是我最初未将其纳入书中的主要原因),并且无法在训练过程中使用。然而,在生产环境中使用LLM时,推理速度的提升往往足以弥补代码复杂性和内存方面的权衡。

假设LLM正在生成一些文本。具体来说,假设给LLM提供了以下提示词:“Time”。你可能已经知道,LLM一次只生成一个词(或token),接下来的两个文本生成步骤可能如下图所示:

该图展示了LLM如何一次生成一个token。从提示词"Time"开始,模型生成下一个token"flies"。在下一步中,整个序列"Time flies"被重新处理以生成token"fast"。

请注意,生成的LLM文本输出中存在一些冗余,如下图所示:

此图突出显示了在每个生成步骤中LLM必须重新处理的重复上下文("Time flies")。由于LLM没有缓存中间键/值状态,每次生成新token(例如"fast")时,它都会重新编码整个序列。

当我们实现LLM文本生成函数时,通常只使用每一步最后生成的token。然而,上面的可视化从概念层面揭示了主要低效之一。如果我们放大注意力机制本身,这种低效(或冗余)会更加明显。(如果你对注意力机制感到好奇,可以阅读我的《从零开始构建大型语言模型》一书第3章,或我的文章《理解并编码LLM中的自注意力、多头注意力、因果注意力和交叉注意力》。)

[

https://substackcdn.com/image/fetch/$s_!3NS4!,w_140,h_140,c_fill,f_webp,q_auto:good,fl_progressive:steep,g_auto/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F69bfee26-ea3b-42a6-8a1a-6b8187852082_738x564.png">![理解并编码LLM中的自注意力、多头注意力、因果注意力和交叉注意力](https://substackcdn.com/image/fetch/$s_!3NS4!,w_140,h_140,c_fill,f_auto,q_auto:good,fl_progressive:steep,g_auto/https%3A%2F%2Fsubstack-post-media.s3.amazonaws.com%2Fpublic%2Fimages%2F69bfee26-ea3b-42a6-8a1a-6b8187852082_738x564.png)

](https://magazine.sebastianraschka.com/p/understanding-and-coding-self-attention)

下图展示了大语言模型(LLM)核心注意力机制计算的一个片段。在此图中,输入词元(“Time”和“flies”)被编码为三维向量(实际应用中这些向量的维度要大得多,但若按真实维度绘制将难以在小型示意图中呈现)。矩阵 W 是注意力机制的权重矩阵,负责将这些输入转换为键向量、值向量和查询向量。

下图展示了底层注意力分数计算的片段,其中键向量和值向量已被高亮标注:

此图展示了LLM在注意力计算过程中如何从词元嵌入中推导出键向量(k)和值向量(v)。每个输入词元(例如"Time"和"flies")都通过已学习的矩阵W_kW_v进行投影,从而获得对应的键向量和值向量。

如前所述,LLM每次生成一个词(或词元)。假设LLM生成了单词“fast”,那么下一轮的提示词就变成了“Time flies fast”。如下图所示:

此图展示了LLM在每一步生成过程中如何为之前见过的词元("Time"和"flies")重新计算键向量和值向量。当生成第三个词元("fast")时,模型会再次重新计算相同的k(1)/v(1)k(2)/v(2)向量,而不是重复使用它们。这种重复计算突显了在自回归解码过程中不使用KV缓存的低效性。

通过对比前两幅图可以看出,前两个词元的键向量和值向量完全相同,在每一轮下一个词元的文本生成中重新计算它们是一种浪费。

因此,KV缓存的思想是实现一种缓存机制,存储之前生成的键向量和值向量以供重复使用,这有助于我们避免这些不必要的重新计算。

在上一节中我们了解了基本概念之后,现在让我们在查看具体代码实现之前,先深入探讨一些细节。如果我们有一个没有 KV缓存的文本生成过程,用于生成“Time flies fast”,可以这样理解:

请注意其中的冗余:词元“Time”和“flies”在每一步新的生成过程中都被重新计算。KV缓存通过存储和重复使用之前计算过的键向量和值向量来解决这一低效问题:

  1. 最初,模型计算并缓存输入词元的键向量和值向量。

  2. 对于每个新生成的词元,模型仅计算该特定词元的键向量和值向量。

  3. 之前计算过的向量会从缓存中检索,以避免重复计算。

下表总结了计算和缓存的步骤及状态:

这样做的好处是,“Time”只计算一次并被重复使用两次,“flies”只计算一次并被重复使用一次。(为了简单起见,这里用了较短的文本示例,但直观上不难理解,文本越长,我们就能越多地重复使用已计算好的键和值,从而提升生成速度。)

下图并排展示了使用和不使用 KV 缓存的生成步骤 3。

对比使用和不使用 KV 缓存的文本生成。在上方面板(无缓存)中,每个 token 步骤都会重新计算键和值向量,导致冗余操作。在下方面板(有缓存)中,之前计算好的键和值会从 KV 缓存中检索,避免重复计算,从而实现更快的生成。

所以,如果我们想在代码中实现 KV 缓存,所要做的就是像往常一样计算键和值,然后将它们存储起来,以便在下一轮中检索。下一节将通过一个具体的代码示例来说明这一点。

实现 KV 缓存的方法有很多种,核心思想是:在每个生成步骤中,只计算新生成 token 的键张量和值张量。

我选择了一种简单的方法,侧重于代码的可读性。我认为最直接的方式就是浏览代码的改动,看看它是如何实现的。

我在 GitHub 上分享了两份文件,它们都是自包含的 Python 脚本,分别实现了不带和带 KV 缓存的 LLM:

  1. gpt_ch04.py:取自我的《从头构建大型语言模型》一书第 3 章和第 4 章的自包含代码,用于实现 LLM 并运行简单的文本生成函数。

  2. gpt_with_kv_cache.py:与上述相同,但进行了必要的修改以实现 KV 缓存。

要阅读与 KV 缓存相关的代码修改,你可以:

a. 打开 gpt_with_kv_cache.py 文件,查找标记新改动的 # NEW 部分:

b. 使用你选择的文件差异比较工具,查看这两个代码文件以对比改动:

此外,为了总结实现细节,以下小节提供了一个简短的说明。

MultiHeadAttention 的构造函数中,我们添加了两个缓冲区 cache_kcache_v,用于跨步骤存储拼接后的键和值:

self.register_buffer("cache_k", None)
self.register_buffer("cache_v", None)

(如果你想了解更多关于缓冲区的信息,我制作了一个 YouTube 视频:理解 PyTorch 缓冲区。)

接下来,我们扩展 MultiHeadAttention 类的 forward 方法,使其接受一个 use_cache 参数:

def forward(self, x, use_cache=False):
    b, num_tokens, d_in = x.shape

    keys_new = self.W_key(x)  # 形状: (b, num_tokens, d_out)
    values_new = self.W_value(x)
    queries = self.W_query(x)
    #...

    if use_cache:
        if self.cache_k is None:
            self.cache_k, self.cache_v = keys_new, values_new
        else:
            self.cache_k = torch.cat([self.cache_k, keys_new], dim=1)
            self.cache_v = torch.cat([self.cache_v, values_new], dim=1)
        keys, values = self.cache_k, self.cache_v
    else:
        keys, values = keys_new, values_new

这里的键和值存储与检索实现了 KV 缓存的核心思想。

存储

具体来说,在通过 if self.cache_k is None: ... 初始化缓存后,我们分别通过 self.cache_k = torch.cat(...)self.cache_v = torch.cat(...) 将新生成的键和值添加到缓存中。

检索

然后,keys, values = self.cache_k, self.cache_v 从缓存中检索存储的值和键。

基本上就是这样:KV 缓存的核心存储与检索机制。接下来的第 3 节和第 4 节只处理一些次要的实现细节。

在生成文本时,我们必须记住在两次独立的文本生成调用之间重置键和值缓冲区。否则,新提示的查询会关注前一个序列中残留的旧键,导致模型依赖无关的上下文并产生不连贯的输出。为了防止这种情况,我们在 MultiHeadAttention 类中添加了一个 reset_kv_cache 方法,以便在后续的文本生成调用之间使用:

def reset_cache(self):
    self.cache_k, self.cache_v = None, None

在完成对 MultiHeadAttention 类的修改后,我们现在修改 GPTModel 类。首先,我们在构造函数中添加一个用于跟踪 token 索引的位置计数器:

self.current_pos = 0

这是一个简单的计数器,用于记录模型在增量生成会话期间已经缓存了多少个 token。

然后,我们将单行的块调用替换为一个显式循环,通过每个 transformer 块传递 use_cache

def forward(self, in_idx, use_cache=False):
    # ...

if use_cache:
        pos_ids = torch.arange(
            self.current_pos, self.current_pos + seq_len,
            device=in_idx.device, dtype=torch.long
        )
        self.current_pos += seq_len
    else:
        pos_ids = torch.arange(
            0, seq_len, device=in_idx.device, dtype=torch.long
        )

pos_embeds = self.pos_emb(pos_ids).unsqueeze(0)
    x = tok_embeds + pos_embeds
    # ...
    for blk in self.trf_blocks:
        x = blk(x, use_cache=use_cache)

上面设置 use_cache=True 时,我们从 self.current_pos 开始计数 seq_len 步。然后,增加计数器,以便下一次解码调用从上次中断的地方继续。

self.current_pos 追踪的原因在于,新的查询必须紧跟在已存储的键和值之后。如果不使用计数器,每一步新操作都会从位置 0 重新开始,导致模型将新 token 视为与之前的 token 重叠。(另一种方式是通过 offset = block.att.cache_k.shape[1] 来追踪偏移量。)

上述改动还需要对 TransformerBlock 类进行小幅修改,以接受 use_cache 参数:

def forward(self, x, use_cache=False):
    # ...
    self.att(x, use_cache=use_cache)

最后,我们在 GPTModel 中添加一个模型级别的重置方法,以便一次性清除所有块的缓存:

def reset_kv_cache(self):
    for blk in self.trf_blocks:
        blk.att.reset_cache()
    self.current_pos = 0

在对 GPTModelTransformerBlockMultiHeadAttention 进行修改后,下面是在简单文本生成函数中使用 KV 缓存的方式:

def generate_text_simple_cached(
        model, idx, max_new_tokens, use_cache=True
    ):
    model.eval()

    ctx_len = model.pos_emb.num_embeddings  # 最大支持长度,例如 1024
    if use_cache:
        # 用完整提示初始化缓存
        model.reset_kv_cache()
        with torch.no_grad():
            logits = model(idx[:, -ctx_len:], use_cache=True)

        for _ in range(max_new_tokens):
            # a) 选择对数概率最高的 token
            next_idx = logits[:, -1].argmax(dim=-1, keepdim=True)
            # b) 将其追加到当前序列中
            idx = torch.cat([idx, next_idx], dim=1)
            # c) 仅将新 token 输入模型
            with torch.no_grad():
                logits = model(next_idx, use_cache=True)
    else:
        for _ in range(max_new_tokens):
            with torch.no_grad():
                logits = model(idx[:, -ctx_len:], use_cache=False)
            next_idx = logits[:, -1].argmax(dim=-1, keepdim=True)
            idx = torch.cat([idx, next_idx], dim=1)

    return idx

注意,在步骤 c) 中,我们通过 logits = model(next_idx, use_cache=True) 仅将新 token 输入模型。而不使用缓存时,由于模型没有存储的键和值可复用,我们需要将整个输入 logits = model(idx[:, -ctx_len:], use_cache=False) 输入模型。

在概念层面了解 KV 缓存后,关键问题是在实际小规模示例中它的性能如何。为了测试实现效果,我们可以将上述两个代码文件作为 Python 脚本运行,它们将驱动一个 1.24 亿参数的小型 LLM 生成 200 个新 token(以 4 个 token 的提示 “Hello, I am” 开头):

pip install -r https://raw.githubusercontent.com/rasbt/LLMs-from-scratch/refs/heads/main/requirements.txt

python gpt_ch04.py

python gpt_with_kv_cache.py

在配备 M4 芯片的 Mac Mini(CPU)上,结果如下:

因此,我们可以看到,即使使用 1.24 亿参数的小模型和 200 token 的短序列长度,也能获得约 5 倍的加速。(注意,此实现以代码可读性为优先,并未针对 CUDA 或 MPS 运行时速度进行优化——若需优化,应预分配张量而非反复创建和拼接。)

注意: 两种情况下模型都会生成“乱码”文本,例如:

输出文本:Hello, I am Featureiman Byeswickattribute argue logger Normandy Compton analogous bore ITVEGIN ministriesysics Kle functional recountrictionchangingVirgin embarrassedgl …

这是因为我们还没有训练模型。下一章会训练模型,届时你可以在训练好的模型上使用KV缓存(但KV缓存仅用于推理阶段)来生成连贯文本。这里我们使用未训练模型是为了保持代码简洁。

不过更重要的是,gpt_ch04.pygpt_with_kv_cache.py 两种实现生成的文本完全相同。这说明KV缓存的实现是正确的——索引错误很容易导致结果不一致。

随着序列长度增加,KV缓存的优缺点会变得更加明显:

  • [优点] 计算效率提升:不使用缓存时,第t步的注意力机制需要将新查询与之前t个键进行比较,累计计算量呈二次方增长O(n²)。使用缓存后,每个键和值只计算一次并重复使用,将每步复杂度降低为线性O(n)。

  • [缺点] 内存占用线性增长:每个新token都会追加到KV缓存中。对于长序列和大型LLM,累积的KV缓存会变得很大,可能消耗大量甚至不可接受的(GPU)内存。作为变通方案,我们可以截断KV缓存,但这会增加更多复杂性(不过在部署LLM时这可能是值得的)。

虽然上述KV缓存的概念性实现有助于清晰理解,主要面向代码可读性和教学目的,但在实际场景中部署(尤其是处理更大模型和更长序列时)需要更精细的优化。

  • 内存碎片化和重复分配:如前所示,通过torch.cat持续拼接张量会导致频繁的内存分配和重新分配,造成性能瓶颈。

  • 内存占用的线性增长:若处理不当,KV缓存的大小对于超长序列会变得不切实际。

与其重复拼接张量,我们可以根据预期的最大序列长度预先分配足够大的张量。这能确保内存使用的一致性并减少开销。伪代码如下:

# 键和值的预分配示例
max_seq_len = 1024  # 最大预期序列长度
cache_k = torch.zeros(
    (batch_size, num_heads, max_seq_len, head_dim), device=device
)
cache_v = torch.zeros(
    (batch_size, num_heads, max_seq_len, head_dim), device=device
)

在推理过程中,我们只需将这些预分配张量的切片写入即可。

为了避免GPU内存爆炸,我们可以实现带动态截断的滑动窗口方法。通过滑动窗口,我们只保留缓存中最后window_size个token:

# 滑动窗口缓存实现
window_size = 512
cache_k = cache_k[:, :, -window_size:, :]
cache_v = cache_v[:, :, -window_size:, :]

你可以在 gpt_with_kv_cache_optimized.py 文件中找到这些优化。

在配备M4芯片(CPU)的Mac Mini上,生成200个token且窗口大小等于LLM上下文长度(以确保结果相同从而实现公平比较)时,代码运行时间对比如下:

不幸的是,在 CUDA 设备上,由于这是一个小型模型,设备传输和通信开销超过了 KV 缓存带来的收益,因此速度优势消失了。

尽管缓存引入了额外的复杂性和内存考量,但效率上的显著提升通常能抵消这些权衡,尤其是在生产环境中。

请记住,虽然我在此优先考虑了代码的清晰性和可读性而非效率,但关键点在于,实际实现通常需要深思熟虑的优化,例如预分配内存或应用滑动窗口缓存来有效管理内存增长。从这个意义上说,我希望本文能为您提供有价值的信息。

欢迎尝试这些技术,祝编码愉快!

在向我的 Qwen3(0.6B)和 Llama 3(1B)从头实现版本添加 KV 缓存后,我运行了额外实验,比较了有无 KV 缓存时的模型运行时间。请注意,我选择了上述提到的 torch.cat 方法,而不是像“优化 KV 缓存实现”部分描述的那样预分配 KV 缓存张量。由于 Llama 3 和 Qwen3 支持非常大的上下文长度(分别为 131k 和 41k 个 token),预分配的张量会消耗约 8 GB 的额外内存,代价相当高昂。

此外,由于我使用了更节省内存的 torch.cat 方法来动态创建张量,我将 KV 缓存移到了模型外部,以便使用 torch.compile 编译模型,从而提升计算效率。

相关代码可在此处找到:

性能表现如下所示。

正如我们所见,在 CPU 上,KV 缓存带来了最显著的加速。而编译进一步提升了这一性能。然而,在 GPU 上,常规编译模型可以达到最佳性能,这可能是因为我们没有在 GPU 上预分配张量,并且模型相对较小。

本杂志是一个个人热情项目。为了支持我作为独立研究员,请考虑购买我的书《从零开始构建大语言模型》,或订阅付费版。

如果您读过这本书并且有几分钟时间,我将非常感激您能写一篇简短的书评。这对我们作者帮助很大!

您的支持意义重大!谢谢!

关于此帖的讨论

准备好了解更多了吗?