原文地址:Understanding and Coding the KV Cache in LLMs from Scratch,by Sebastian Raschka, on 2025-07-17
从零开始理解并实现大语言模型中的KV缓存
KV缓存是生产环境中大语言模型实现高效推理的核心技术之一,也是提升大语言模型推理计算效率的重要组件。本文将从概念原理和代码实现两个层面,通过一套从零编写、可读性强的实现方案,讲解KV缓存的工作机制。
距离我上次分享讲解大语言模型基础概念的技术教程已经有一段时间了。目前我正处于伤病恢复期,同时在撰写一篇更重磅的大语言模型研究主题文章,因此打算分享一篇读者多次询问的主题教程(这个主题并未收录在我的《从零构建大语言模型》一书中)。
祝阅读愉快!
概述
简而言之,KV缓存会存储推理阶段(训练结束后)中间的键(K)和值(V)计算结果以供复用,从而大幅提升文本生成的速度。KV缓存的缺点在于会增加代码复杂度、提升内存需求(这也是我最初没有把它写入书中的主要原因),并且无法在训练过程中使用。但在生产环境部署大语言模型时,推理速度的提升通常足以覆盖代码复杂度和内存开销上的代价。
什么是KV缓存?
想象一下大语言模型正在生成文本的场景。具体来说,假设给模型输入的提示词是:“Time”。你可能已经知道,大语言模型每次生成一个单词(或词元),后续的两步文本生成过程可以用下图说明:

该图展示了大语言模型如何逐词元生成文本。从提示词“Time”开始,模型生成下一个词元“flies.”;下一步,再基于完整序列“Time flies”继续生成词元“fast”。
值得注意的是,大语言模型的文本生成过程存在大量重复计算,正如下图所突出显示的:

这张图标出了每一步生成时都需要重复处理的上下文(“Time flies”)。由于大语言模型没有缓存中间的键/值状态,每生成一个新词元(比如“fast”),它都要对完整序列重新编码一次。
我们实现大语言模型文本生成函数时,通常只会用到每一步生成的最后一个词元。但上图从概念层面揭示了其中最主要的低效问题之一。如果我们把视角聚焦到注意力机制本身,这种低效(或者说冗余)会更加清晰。(如果你对注意力机制感兴趣,可以阅读我的书https://amzn.to/4fqvn0D第3章,或者我发布的文章https://magazine.sebastianraschka.com/p/understanding-and-coding-self-attention了解更多内容。)
下图是大语言模型核心的注意力机制计算片段。图中输入词元(“Time”和“flies”)被编码为三维向量(实际应用中这些向量的维度要大得多,只是为了适配图片尺寸做了简化)。矩阵W是注意力机制的权重矩阵,负责将输入转换为键向量、值向量和查询向量。
下图展示了注意力分数计算的底层片段,其中高亮标出了键向量和值向量:

该图说明了大语言模型在注意力计算过程中,如何从词元嵌入中得到键(k)向量和值(v)向量。每个输入词元(比如“Time”和“flies”)都会通过训练得到的矩阵$W_k$和$W_v$进行投影,得到对应的键向量和值向量。
如前所述,大语言模型每次生成一个单词(或词元)。假设模型生成了单词“fast”,那么下一轮的提示词就变成了“Time flies fast”,如下图所示:

这张图展示了在每一步生成过程中,大语言模型如何对已经见过的词元(“Time”和“flies”)重复计算键向量和值向量。生成第三个词元(“fast”)时,模型会重新计算一遍完全相同的k(1)/v(1)和k(2)/v(2)向量,而不是直接复用它们。这种重复计算凸显了自回归解码过程中不使用KV缓存的低效性。
通过对比前两张图可以发现,前两个词元的键向量和值向量是完全相同的,如果每生成一个新词元都重新计算一遍,会造成大量的计算浪费。
而KV缓存的核心思路,就是实现一套缓存机制,存储之前生成的键向量和值向量以供复用,从而避免这些不必要的重复计算。
大语言模型生成文本的过程:无缓存 vs 有KV缓存
上一节我们讲完了基础概念,在看具体的代码实现之前,我们再进一步细化讲解。如果用不带KV缓存的方式生成“Time flies fast”这段文本,过程大致如下:

注意其中的冗余:词元“Time”和“flies”在每一步生成时都要被重新计算。KV缓存通过存储并复用之前计算好的键向量和值向量,解决了这种低效问题:
- 初始阶段,模型计算输入词元的键向量和值向量,并缓存起来。
- 每生成一个新词元时,模型只计算该词元对应的键向量和值向量。
- 之前计算好的向量直接从缓存中读取,避免重复计算。
下表总结了计算和缓存的步骤与状态:

这样做的好处是,“Time”只计算1次、复用2次,“flies”只计算1次、复用1次。(这里为了简化用了短文本举例,不难理解:文本越长,已计算的键和值能被复用的次数就越多,生成速度的提升也就越明显。)
下图并排对比了第3步生成时,使用和不使用KV缓存的差异。

文本生成有无KV缓存的对比。上图(无缓存)中,每一步词元生成都要重新计算所有键向量和值向量,造成了冗余运算。下图(有缓存)中,之前计算好的键和值直接从KV缓存中读取,避免了重复计算,提升了生成速度。
所以,如果要在代码中实现KV缓存,我们只需要像往常一样计算键和值,然后把它们存储起来,供下一轮调用即可。下一节我们将通过具体的代码示例进行说明。
从零实现KV缓存
KV缓存的实现方式有很多,核心思想都是每一步生成时,只计算新生成词元对应的键张量和值张量。
我选择了一种偏重代码可读性的简单实现方案。我认为直接浏览代码改动,是最容易理解实现方式的途径。
我在GitHub上分享了两个文件,都是独立的Python脚本,分别从零实现了带KV缓存和不带KV缓存的大语言模型:
- https://github.com/rasbt/LLMs-from-scratch/blob/main/ch04/03_kv-cache/gpt_ch04.py :摘自我的《从零构建大语言模型》一书第3、4章的完整代码,实现了大语言模型和基础的文本生成功能
- https://github.com/rasbt/LLMs-from-scratch/blob/main/ch04/03_kv-cache/gpt_with_kv_cache.py :在上述代码基础上,加入了实现KV缓存所需的改动。
要查看KV缓存相关的代码修改,你可以选择以下任意一种方式:
a. 打开gpt_with_kv_cache.py文件,查找标注了# NEW的部分,这些就是新增的改动:

b. 用你习惯的文件对比工具,对比两个代码文件的差异:

此外,为了总结实现细节,下面的小节会做简要的步骤讲解。
1. 注册缓存缓冲区
在MultiHeadAttention(多头注意力)的构造函数中,我们添加两个缓冲区cache_k和cache_v,用于存储多步累积拼接起来的键和值:
self.register_buffer("cache_k", None)
self.register_buffer("cache_v", None)
(如果你想了解更多关于缓冲区的知识,我制作过一期YouTube视频:https://youtu.be/PetlIokI9Ao )
2. 带use_cache标志的前向传播
接下来,我们扩展MultiHeadAttention类的forward方法,增加一个use_cache参数:
def forward(self, x, use_cache=False):
b, num_tokens, d_in = x.shape
keys_new = self.W_key(x) # Shape: (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节只是处理一些实现上的细节问题。
3. 清空缓存
生成文本时,我们必须记得在两次独立的文本生成调用之间重置键缓冲区和值缓冲区。否则,新提示词的查询会注意力到上一条序列遗留下来的旧键值,导致模型依赖无关上下文、输出内容混乱。为了避免这种情况,我们给MultiHeadAttention类添加一个reset_cache方法,方便后续在两次文本生成调用之间调用:
def reset_cache(self):
self.cache_k, self.cache_v = None, None
4. 在完整模型中传递use_cache参数
完成MultiHeadAttention类的修改后,我们现在来修改GPTModel类。首先,我们在构造函数中添加一个词元索引的位置追踪器:
self.current_pos = 0
这是一个简单的计数器,用于记录在增量生成过程中,模型已经缓存了多少个词元。
然后,我们把单行的块调用替换为显式循环,将use_cache参数传递给每个Transformer块:
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开始,模型就会认为新词元和之前的词元是重叠的。(另一种实现方式是通过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
5. 在生成过程中使用缓存
完成GPTModel、TransformerBlock和MultiHeadAttention的修改后,最终我们可以这样在简单的文本生成函数中使用KV缓存:
def generate_text_simple_cached(
model, idx, max_new_tokens, use_cache=True
):
model.eval()
ctx_len = model.pos_emb.num_embeddings # max sup. len., e.g. 1024
if use_cache:
# Init cache with full prompt
model.reset_kv_cache()
with torch.no_grad():
logits = model(idx[:, -ctx_len:], use_cache=True)
for _ in range(max_new_tokens):
# a) pick the token with the highest log-probability
next_idx = logits[:, -1].argmax(dim=-1, keepdim=True)
# b) append it to the running sequence
idx = torch.cat([idx, next_idx], dim=1)
# c) feed model only the new 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)。而在无缓存的情况下,我们需要把完整的输入都传给模型:logits = model(idx[:, -ctx_len:], use_cache=False),因为没有存储好的键和值可以复用。
简单的性能对比
从概念层面讲完KV缓存后,更重要的问题是:在实际的小例子中,它的性能到底怎么样?我们可以把上面提到的两个代码文件作为Python脚本运行,用一个1.24亿参数的小型大语言模型,基于4个词元的提示词“Hello, I am”生成200个新词元,来测试效果:
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个词元的短序列,也已经获得了约5倍的速度提升。(注意:本实现的优化目标是代码可读性,而非CUDA或MPS的运行速度;如果要优化硬件运行速度,需要预分配张量,而不是反复重建和拼接张量。)
注意:两种情况下模型生成的都是“乱码”,也就是类似这样的文本:
输出文本:Hello, I am Featureiman Byeswickattribute argue logger Normandy Compton analogous bore ITVEGIN ministriesysics Kle functional recountrictionchangingVirgin embarrassedgl …
这是因为我们还没有训练模型。下一章会训练模型,之后你就可以在训练好的模型上用KV缓存(注意:KV缓存仅用于推理阶段)生成通顺的文本。这里用未训练的模型只是为了简化代码演示。
但更重要的是,gpt_ch04.py和gpt_with_kv_cache.py两种实现生成的文本完全一致。这说明KV缓存的实现是正确的——索引问题很容易出错,一旦出错就会导致输出结果不一致。
感谢阅读《Ahead of AI》!免费订阅即可接收新文章,支持我的创作。
KV缓存的优缺点
随着序列长度增加,KV缓存的优势和劣势都会更加凸显,具体表现为:
- 【优点】计算效率提升:没有缓存时,第t步的注意力必须将新查询与前t个键做对比,累计计算量呈平方级增长,复杂度为$O(n^2)$。使用缓存后,每个键和值只计算一次、后续复用,每一步的总复杂度降为线性的$(O(n))$。
- 【缺点】内存占用线性增长:每个新词元都会追加到KV缓存中。对于长序列和更大的大语言模型,累计的KV缓存会变得非常大,可能会占用大量(GPU)内存,甚至高到无法承受。作为折中方案,我们可以对KV缓存做截断,但这会进一步增加复杂度(不过在部署大语言模型时,这些代价通常是值得的)。
优化KV缓存的实现
上面这套KV缓存的概念实现侧重清晰易懂,主要服务于代码可读性和教学目的。如果要在实际场景中部署(尤其是更大的模型、更长的序列长度),还需要更细致的优化。
缓存扩展时的常见问题
- 内存碎片与重复分配:如前文所示,通过
torch.cat不断拼接张量,会导致频繁的内存分配与重分配,形成性能瓶颈。 - 内存占用线性增长:如果处理不当,序列非常长时,KV缓存的体积会大到无法使用。
技巧1:预分配内存
与其反复拼接张量,我们可以根据预期的最大序列长度,预分配一个足够大的张量。这样能保证内存使用稳定,减少开销。伪代码大致如下:
# Example pre-allocation for keys and values
max_seq_len = 1024 # maximum expected sequence length
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
)
推理时,我们只需要往这些预分配张量的对应切片里写入数据即可。
技巧2:通过滑动窗口截断缓存
为了避免GPU内存爆炸,我们可以实现带动态截断的滑动窗口方案。通过滑动窗口,缓存中只保留最近的window_size个词元:
# Sliding window cache implementation window_size = 512 cache_k = cache_k[:, :, -window_size:, :] cache_v = cache_v[:, :, -window_size:, :]
实际应用中的优化
你可以在这个文件中找到这些优化实现:
https://github.com/rasbt/LLMs-from-scratch/blob/main/ch04/03_kv-cache/gpt_with_kv_cache_optimized.py
在搭载M4芯片的Mac Mini(CPU运行)上,生成200个词元、且窗口大小等于大语言模型的上下文长度(以保证结果一致,实现公平对比)时,代码的运行时间对比如下:

遗憾的是,在CUDA设备上,这种小模型的速度优势会消失。因为模型太小,设备间的数据传输和通信开销超过了KV缓存带来的收益。
结论
虽然缓存会带来额外的复杂度和内存开销,但效率上的显著提升通常足以覆盖这些代价,在生产环境中尤其如此。
请记住,本文的实现优先考虑代码的清晰性和可读性,而非运行效率;但核心结论是:工程落地的实现往往需要经过深思熟虑的优化,比如预分配内存、或者使用滑动窗口缓存来有效管理内存增长。从这个角度来说,希望这篇文章能对你有所帮助。
欢迎大家动手尝试这些技术,编码愉快!
补充:Qwen3和Llama 3中的KV缓存
在我从零实现的Qwen3(6亿参数)和Llama 3(10亿参数)中加入KV缓存后,我又做了额外的实验,对比了使用和不使用KV缓存时的模型运行时间。注意:我采用的是上文提到的torch.cat拼接方案,而不是“优化KV缓存实现”一节中说的预分配KV缓存张量。因为Llama 3和Qwen3支持的上下文长度非常大(分别为13.1万和4.1万个词元),预分配张量会额外消耗约8GB内存,开销很高。
此外,由于我使用更省内存的torch.cat方式动态生成张量,我把KV缓存移到了模型外部,这样就可以用torch.compile编译模型,进一步提升计算效率。
代码可以在这里查看:
性能表现如下所示。


可以看到,在CPU上,KV缓存带来的速度提升最为显著,编译优化还能进一步提速。但在GPU上,常规的编译后模型就能达到最佳性能——这很可能是因为我们没有在GPU上预分配张量,而且模型本身规模也比较小。
这本杂志是我的个人热爱项目。如果你想支持我这位独立研究者,可以考虑购买我的书(https://amzn.to/4fqvn0D),或者订阅我的杂志(https://magazine.sebastianraschka.com/subscribe)。

《从零构建大语言模型》购买链接:https://amzn.to/4fqvn0D
如果你已经读过这本书,并且能抽出几分钟时间,我非常希望你能留下评价:https://www.amazon.com/Build-Large-Language-Model-Scratch/dp/1633437167 。这对我们作者帮助很大!
非常感谢你的支持!