从大模型执行过程理解 KV Cache 和 Prompt Cache

背景

之前看到 KV Cache,只知道它缓存历史 token 的 Key 和 Value。

但这句话后面接着冒出来的词更多:token、embedding、Transformer、attention head、Q/K/V、prefill、decode、Prompt Cache……感觉像是在用一堆新名词解释另一个新名词。

这次直接从根源理解,看一下「一段文字进入大模型之后,到底经历了什么?」。

使用 distilbert/distilgpt2(DistilGPT-2)跑实验。它是 GPT-2 的蒸馏版:有 6 层 Transformer、每层 12 个 attention head,规模远小于今天常见的大模型,但完整保留了自回归生成、attention 和 KV Cache 的执行过程。

大模型是一个 token 一个 token 地生成文本

大模型不会一次把整段回答都写好。给它一段已有文本,它只做一件事:为下一个 token给出整张词表上的分数,然后按某种策略选出一个 token,接到原文本后面,再重复同样的过程。

1
2
3
4
5
已有 token:'The' / ' cat' / ' sat' / ' on'
→ 模型预测第 5 个 token:' the'
→ 已有 token:'The' / ' cat' / ' sat' / ' on' / ' the'
→ 模型预测第 6 个 token
→ ……

这叫自回归生成。在这个例子里,The / cat / sat / on 是已经给定的提示词;模型先处理这 4 个 token,得到第 5 个 token 的预测。选出第 5 个 token 后,它又成了新的上下文的一部分,模型据此继续预测第 6 个。

KV Cache 就出现在这个循环里:第 5 个 token 到来时,前 4 个 token 已经被处理过;第 6 个 token 到来时,前 5 个 token 又已经被处理过。后面会看到,模型怎样避免每一步都从头计算这段不断变长的上下文。

要理解这个循环,先从最开始的一次处理看起:文本怎样变成 token,token 又怎样进入模型。

文本进模型之前,模型其实还什么都没做

先加载 tokenizer 和模型:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
from transformers import AutoModelForCausalLM, AutoTokenizer

tokenizer = AutoTokenizer.from_pretrained("distilbert/distilgpt2")
model = AutoModelForCausalLM.from_pretrained(
"distilbert/distilgpt2",
).eval()

embedding_table = model.get_input_embeddings().weight

print("词表大小:", tokenizer.vocab_size)
print("embedding 表形状:", tuple(embedding_table.shape))

# 输出:
# 词表大小: 50257
# embedding 表形状: (50257, 768)

这里有两样东西,容易混在一起。

tokenizer 可以理解成一套翻译规则:它把文本翻译成模型能处理的编号。

model 才是训练出来的大模型本体。它里面有很多参数矩阵,包括 embedding 表、Transformer 的参数、最后预测下一个 token 的参数。

词表有 50,257 个位置。每个位置对应一个 token ID;embedding 表有 50,257 行,每个 ID 都能查到一行浮点数。

接着让 tokenizer 处理文本:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
prompt = "The cat sat on"
input_ids = tokenizer(prompt, return_tensors="pt").input_ids

tokens = [
repr(tokenizer.decode([token_id]))
for token_id in input_ids[0].tolist()
]

print("文本:", repr(prompt))
print("token IDs:", input_ids[0].tolist())
print("tokens:", tokens)

# 输出:
# 文本: 'The cat sat on'
# token IDs: [464, 3797, 3332, 319]
# tokens: ["'The'", "' cat'", "' sat'", "' on'"]

token 不是固定的一个字或一个词,而是 tokenizer 按自己的词表切出的一段文本。比如这里的 ' cat' 连前面的空格也带上了,中文里「今天天气怎么样」可能被切成 ['今天', '天气', '怎么样', '?'],也可能切得更细,具体取决于模型的 tokenizer。

回到例子中,此时模型实际看到的是:

1
[464, 3797, 3332, 319]

不是:

1
The cat sat on

再看其中第一个 ID 怎样变成向量:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
first_token_id = input_ids[0, 0].item() # 模型要支持一次并行处理多条输入,所以这里是二维数组
first_embedding = embedding_table[first_token_id]

print(
"第一个 token 的 embedding 形状:", tuple(first_embedding.shape),
"\\n前 8 个数:", first_embedding[:8].tolist(),
)

# 输出:
# 第一个 token 的 embedding 形状: (768,)
# 前 8 个数: [-0.06264858692884445, -0.04490645229816437,
# 0.0558876097202301, -0.05465700104832649,
# -0.1171262264251709, -0.07286953926086426,
# -0.22325637936592102, -0.0032198030967265368]

这条链路到这里才真正进入模型内部:

1
2
3
4
5
6
文本
→ tokenizer 切 token
→ token ID
→ 查 token embedding 表
→ 查位置 embedding 表
→ 两条向量逐元素相加

"The" 这个 token 的 ID 是 464,模型从 token embedding 表第 464 行取出一个 768 维向量,记作 E_The。但 E_The 只说明「这是 The」;它还没有说明 The 在这段输入中的第几个位置。

1
2
3
[-0.06264858692884445, -0.04490645229816437, 0.0558876097202301, -0.05465700104832649,
-0.1171262264251709, -0.07286953926086426,
-0.22325637936592102, -0.0032198030967265368... 共 768 个浮点数]

DistilGPT-2 还有一张训练好的位置 embedding 表,形状是 (1024, 768):最多 1024 个位置,每个位置也用一个 768 维向量表示。当前这 4 个 token 的位置编号是 [0, 1, 2, 3],模型从这张表中取出 P_0P_1P_2P_3,再和对应 token 向量逐元素相加:

1
2
3
4
h_The = E_The + P_0
h_cat = E_cat + P_1
h_sat = E_sat + P_2
h_on = E_on + P_3

这里不是把两个向量拼接起来;每一对都是 768 个位置一一相加,所以结果仍是 768 维。这样,同一个 token 即使出现在不同位置,进入模型的向量也会不同。

1
2
3
4
5
6
7
8
9
10
11
12
transformer = model.transformer  # 取出模型中的 Transformer 主体,里面有 wte 和 wpe 两张表
position_ids = torch.arange(input_ids.shape[1]).unsqueeze(0) # 4 个 token 的位置编号:[[0, 1, 2, 3]]

token_embeddings = transformer.wte(input_ids) # 用 4 个 token ID 查 token embedding 表
position_embeddings = transformer.wpe(position_ids) # 用 4 个位置 ID 查位置 embedding 表
h0 = token_embeddings + position_embeddings # 每个 token 向量和自己对应的位置向量逐元素相加

print(position_ids.tolist()) # [[0, 1, 2, 3]]:本次输入实际用到的 4 个位置
print(tuple(token_embeddings.shape)) # (1, 4, 768):1 条输入、4 个 token、每条 768 维
print(tuple(position_embeddings.shape)) # (1, 4, 768):同样查出了 4 条位置向量
print(tuple(transformer.wpe.weight.shape)) # (1024, 768):位置表共有 1024 行,每行 768 维
print(tuple(h0.shape)) # (1, 4, 768):相加不改变形状

这里的 1024 不是这次输入有 1024 个 token,而是这张位置 embedding 表预先准备了位置 01023 的 1024 行。当前输入只有 4 个 token,所以只查到其中的第 0123 行;它也意味着这个版本的 DistilGPT-2 最多只能处理 1024 个位置。

后面所有 Transformer 计算,处理的是 h_The / h_cat / h_sat / h_on 这样的向量,不再是字符。

先看 Transformer 是干什么的

此时模型已经拿到四个带有位置信息的 token 向量:

1
2
3
4
The 的向量
cat 的向量
sat 的向量
on 的向量

但它们还只是每个 token 最初的表示。模型需要做的是:让每个位置结合前面的上下文,变成一个更有上下文信息的新向量。

例如,单独看:

1
on

信息很少。但放在:

1
The cat sat on

里,它前面有主体、有动作,模型需要把这些关系带进当前位置的向量里。

完成这件事的一整套模块,就叫一个 Transformer 层

一层 Transformer 中最重要的两步可以先粗略理解成:

1
2
3
4
5
Attention:
当前 token 去前面 token 里取自己需要的信息。

MLP:
拿到这些信息后,再对当前 token 做一轮自己的加工。

所以一层的作用不是「生成一个词」,而是:

1
2
3
所有 token 的当前向量
→ 彼此交换上下文信息
→ 每个 token 得到更丰富的新向量

模型会把这种层一层一层叠起来:

1
2
3
4
5
6
token embedding + position embedding
→ 第 0 层 Transformer
→ 第 1 层 Transformer
→ …
→ 最后一层的向量
→ 预测下一个 token

靠后一层收到的不是最初的 token 向量,而是已经混入前文信息、被前面层加工过的向量。

这就是「层」的作用:把一次还不够的上下文加工,连续做很多轮。

Attention Head 又是干什么的

前面说的 Attention,可以先想成一条通道:当前 token 用自己的 Query 去查询前面 token 的 Key 和 Value,得到一份上下文信息。

Transformer 实际不会只保留一条通道,而是同时运行多条独立的 Attention 通道;每条通道就叫一个 attention head。不同 head 有各自的 Q/K/V 参数,最后再把它们得到的信息合并。

可以先把它理解成:

1
2
3
4
5
一层 Transformer
├─ head 0:用一套自己的参数看上下文
├─ head 1:用另一套自己的参数看上下文
├─ …
└─ 把这些结果合并,再继续加工

每个 head 都有独立的参数,所以它们不一定会看同一种关系。

有的 head 可能更偏向相邻位置,有的可能更在意某些长距离关联。真实模型里,这些模式未必总能被清楚命名;不需要说「某个 head 专门负责语法、某个专门负责事实」,那样有点过度解释。

这个实验模型的配置是:

1
2
3
4
5
6
7
8
print(
"Transformer 层数:", model.config.n_layer,
"attention head 数:", model.config.n_head,
"hidden size:", model.config.n_embd, # 叫 hidden 是因为这串向量只在模型内部各层之间传递,用户看不到。
)

# 输出:
# Transformer 层数: 6 attention head 数: 12 hidden size: 768

现在这三个数字就能对上了:

1
2
3
4
5
6
7
8
6 层 Transformer:
输入会连续经过六轮上下文加工。

12 个 attention head:
每一层里有 12 套并行的 attention。

hidden size = 768
每个 token 当前的完整向量有 768 个数字。

因为 hidden size 是 768,同时又有 12 个 head,所以每个 head 分到的向量宽度是:

1
768 ÷ 12 = 64

这也是后面缓存形状里最后那个 64 的来源。

真实大模型通常有几十到上百层 Transformer、几十到上百个 attention head,hidden size 往往是数千到上万维;DistilGPT‑2 仍然小得多,但已经能让我们看到多个 head 和 64 维 K/V 的实际形状。

Attention 里面的 Q、K、V 是怎么来的

前面已经拿到了 The / cat / sat / on 的 token ID 和 embedding。下面沿着 DistilGPT-2 第 0 层 attention 的实际执行顺序往下看:先得到层输入,再做 QKV 投影,接着拆成 12 个 head,最后计算 ' on' 对前面位置的读取结果。

先把 Q、K、V 在这件事里的分工说清楚。它们都不是单个能直接翻译成中文的数,而是模型为每个 token 算出的三条向量;名字描述的是它们在 attention 中扮演的角色:

1
2
3
Q(Query,查询):当前 token 想从上下文中找什么
K(Key,索引):每个可见 token 提供什么线索,供 Q 判断是否匹配
V(Value,内容):匹配后真正被取回、再混合进当前 token 的信息

例如更新 ' on' 时,Q_on 会和 The / cat / sat / on 各自的 K 比较,先得到「该读谁、各读多少」的权重;再用同一组权重混合四条 V,得到 ' on' 在这个 head 中的更新结果。Q 和 K 负责决定读取比例,V 才是被读取的内容。QKV 只是把 Q、K、V 这三组向量合在一起的简称;下面再看它们的数值是怎样从模型参数中算出来的。

1. 进入第 0 层时,手里有什么

四个 token 先查 embedding 表,再加上各自的位置 embedding。于是第 0 层看到的是一个矩阵,而不是四次互不相干的调用:

1
2
3
4
5
H₀ =
[h_The
h_cat
h_sat
h_on] shape: (1, 4, 768)

三个维度依次是:batch=14 个 token每个 token 的完整向量宽度 768。到这里,H₀ 只是第 0 层的输入;attention 还没有开始。

2. attention 之前,先做一次 LayerNorm

DistilGPT-2 的每一层都会先对每个 token 当前的 768 个数做 LayerNorm,再送入 attention。它的作用是把这一条向量的数值范围拉回比较稳定的尺度,让后续层更容易训练和计算。

对某一个 token 的向量 h,LayerNorm 只在它自己的 768 个维度内计算均值 μ 和方差 σ²

1
2
3
μ   = mean(h₁, h₂, ..., h₇₆₈)
σ² = mean((h - μ)²)
xᵢ = γᵢ × (hᵢ - μ) / √(σ² + ε) + βᵢ

γβ 也是训练好的 768 维参数。这里没有 token 之间的信息交换:The 的 LayerNorm 不会读取 cat,它只调整 The 自己那条向量的尺度。因此 LayerNorm 前后形状完全不变:

1
2
3
4
5
6
7
8
9
transformer = model.transformer
layer0 = transformer.h[0]

position_ids = torch.arange(input_ids.shape[1]).unsqueeze(0)
h0 = transformer.wte(input_ids) + transformer.wpe(position_ids)
x = layer0.ln_1(h0)

print(tuple(h0.shape)) # (1, 4, 768)
print(tuple(x.shape)) # (1, 4, 768)

下面把 LayerNorm 输出记作 X。其中 x[0, 3]' on' 在进入 attention 前的当前向量:它有 768 个浮点数。打印前 8 个数看起来会是 [-0.12, 0.35, ...] 这样的普通浮点数;单独某一维没有可直接翻译成人类词义的含义,重要的是它们作为一个 768 维整体参与后续计算。

3. W_QKV 从哪里来

它不是运行时临时算出来的,而是训练好的模型参数。训练时它从随机数开始,随着「预测下一个 token」的误差不断经反向传播调整;训练完成后,数值被保存在模型权重文件(如 model.safetensors)里。

理论图里常把三套参数写成 W_QW_KW_V。DistilGPT-2 为了计算方便,把它们拼成一张大矩阵:

1
2
W_QKV = [W_Q | W_K | W_V]        shape: (768, 2304)
b_QKV shape: (2304,)

其中 2304 = 3 × 768。所以一次矩阵乘法就能同时产生 Q、K、V。下面把这次计算得到的结果统一记作小写 qkv

1
qkv = X × W_QKV + b_QKV          shape: (1, 4, 2304)

这一次「乘权重矩阵、再加偏置」的线性变换称作投影:把每个 token 当前的 768 个数,按训练好的规则算成 2304 个新数。不是把向量简单截断,也不是几何课里把东西投到平面上。

到这一步所有 token 的 Q、K、V 其实已经算出来了,后续只是切割。

4. 先切成 Q、K、V 三份,再排成 12 个 head

投影输出 qkv 的形状是 (1, 4, 2304)。先只看其中一个位置,比如 ' on':它对应的一行有 2304 个数,按顺序存成下面三段:

1
2
3
4
' on' 的 2304 个数
├─ 前 768 个数:Q_on(完整 Q)
├─ 中间 768 个数:K_on(完整 K)
└─ 后 768 个数:V_on(完整 V)

Thecatsat 也各有这样一行;所以沿最后一维每 768 个数切一次,就得到三块形状相同的张量:

下面的 qkv 只是 Python 变量名,分别装着 Q、K、V;大小写不表示两次不同的计算,也不表示两种不同的数据。

1
2
3
4
5
q, k, v = qkv.split(768, dim=-1)  # Python 中的小写 q/k/v,分别装 Q/K/V 三段

print(tuple(q.shape)) # (1, 4, 768):4 个 token 各有一条完整的 q(Query)
print(tuple(k.shape)) # (1, 4, 768):4 个 token 各有一条完整的 k(Key)
print(tuple(v.shape)) # (1, 4, 768):4 个 token 各有一条完整的 v(Value)

例如 k[0, 0]The 的完整 768 维 Key,k[0, 3]' on' 的完整 768 维 Key。它们都在这一次批量投影中产生;不是等到 ' on' 发起查询时才去生成前面三个位置的 Key。

这时的每条 qkv 仍是 768 维的总结果。DistilGPT-2 在配置中把这 768 维安排给 12 个 attention head 使用,因此每个 head 分到 768 ÷ 12 = 64 维:

1
2
3
一条 768 维 Q
→ [head 0 的 64 维 | head 1 的 64 维 | ... | head 11 的 64 维]
→ 形状从 (768) 变为 (12, 64)

这里只是把已经算好的 768 个数重新标出「哪 64 个属于哪个 head」,没有再做一次矩阵乘法。完整的 768 维输入 X 已经先参与了上面的 X × W_QKV 投影;64 不是从原始输入向量中预先切出来单独计算的。

代码把这三个张量都做同样的重排:

1
2
3
4
5
6
7
8
9
10
num_heads = 12                  # DistilGPT-2 的 head 数
head_dim = 768 // num_heads # 每个 head 的宽度:64

def split_heads(t):
batch_size, token_count, _ = t.shape # 这里是 1 条输入、4 个 token
t = t.view(batch_size, token_count, num_heads, head_dim) # (1, 4, 768) → (1, 4, 12, 64)
return t.permute(0, 2, 1, 3) # 调整为 (batch, head, token, 维度)

q, k, v = map(split_heads, (q, k, v)) # 三个张量都得到 12 个 head 的表示
print(tuple(q.shape)) # (1, 12, 4, 64)

现在维度顺序是 (batch, head, token 位置, head 内向量维度)。因此在第 0 层、第 0 个 head里:

1
2
3
4
5
Q_on  = q[0, 0, 3, :]    shape: (64,)
K_The = k[0, 0, 0, :] shape: (64,)
K_cat = k[0, 0, 1, :] shape: (64,)
K_sat = k[0, 0, 2, :] shape: (64,)
K_on = k[0, 0, 3, :] shape: (64,)

同理,V_TheV_on 也是四条 64 维向量。

5. ' on' 用自己的 Q,读取四个位置的 K/V

on 的 Q 与历史 Key 匹配,再使用权重混合 Value

现在 K_TheK_catK_satK_on 的来源已经清楚了:它们是四个位置各自的 64 维 Key。更新 ' on' 时,模型只取它自己的 Q_on,分别和这四条 Key 做点积:

1
2
3
scores = Q_on · [K_The, K_cat, K_sat, K_on]ᵀ / √64
α = softmax(scores)
h′_on = α_The × V_The + α_cat × V_cat + α_sat × V_sat + α_on × V_on

第一行的输出是 4 个标量分数;除以 √64 = 8 后再做 softmax,

第二行得到 4 个非负、总和为 1 的权重;

第三行才混合 4 条 64 维 Value,输出仍然是 64 维向量。

Q/K 只负责产生 α,V 不参加点积匹配。

对当前下载的 DistilGPT-2,脚本重建出的第 0 层、第 0 个 head、' on' 的权重是:

1
2
α = [0.4397, 0.2247, 0.2872, 0.0484]
The cat sat on

所以这一个 head 的实际输出是:

1
2
3
4
h′_on = 0.4397 × V_The
+ 0.2247 × V_cat
+ 0.2872 × V_sat
+ 0.0484 × V_on

这里的百分比不是下一个 token 的概率,也不是「The 比 cat 更重要」的全局结论。它只说明:在第 0 层、第 0 个 head、更新 ' on' 的这一刻,这个 head 主要从前三个位置取回 Value,很少读取 ' on' 自己的 Value。别的 head 和别的层会有不同的权重。

6. 明明四个位置的 Q/K/V 都算出来了,为什么 The 看不到后边的 token

这里要把「能并行生成 Q/K/V」和「能读取谁」分开。prefill 时四个位置的 Q/K/V 会一次性计算完;但在 scores 进入 softmax 前,causal mask 会把未来位置的分数遮住。于是:

1
2
3
4
The  的 Q 只能匹配 K_The
cat 的 Q 只能匹配 K_The、K_cat
sat 的 Q 只能匹配 K_The、K_cat、K_sat
on 的 Q 只能匹配 K_The、K_cat、K_sat、K_on

也就是说,未来 token 的 K/V 可以已经算好,却不能参与当前 token 的权重和加权和。这个因果限制还有一个直接后果:生成新 token 不会反过来改变旧位置已经得到的 K/V;后面 KV Cache 能复用历史 K/V,就建立在这里。

一个 Transformer 层结束后,怎样得到下一个 token 的预测

前面展开的是一个 head 如何更新 ' on'。一层里其实有 12 个 head;它们各自为每个 token 产出一条 64 维结果。对同一个 token,把 12 条结果并排放回去,就又回到了 768 维:

1
2
12 个 head 的输出: (1, 12, 4, 64)
调整维度并拼接后: (1, 4, 768)

接着有一个输出投影 W_O,形状是 (768, 768)。它把 12 个 head 拼接后的信息再混合一次,输出仍是每个 token 一条 768 维向量。这里的「混合」让一个 head 的结果不必永远只待在自己原来的 64 维槽位里。

attention 输出不会直接覆盖旧向量,而是先加回进入 attention 前的向量,这叫残差连接

1
h_after_attention = h_before_attention + attention_output

随后每个 token 再经过一次 LayerNorm 和 MLP。DistilGPT-2 的 MLP 先把宽度从 768 扩大到 3072,经过 GELU 非线性激活,再投影回 768。3072 = 4 × 768:这里的 4 是 GPT-2 架构预先定下的 MLP 扩展倍率,不是根据 head 数或当前输入临时算出来的;换一种模型,这个倍率也可能不同。

1
2
3
4
5
(1, 4, 768)
→ Linear:768 → 3072
→ GELU
→ Linear:3072 → 768
→ 加回 MLP 前的向量(第二次残差连接)

MLP 不在 token 之间读取信息;它和前面的 LayerNorm 一样,逐个加工每条 token 向量。attention 负责从上下文取信息,输出投影负责混合 head,MLP 负责在每个位置内部进一步变换。完成这三部分后,才得到这一层的输出,并作为下一层的输入。

DistilGPT-2 有 6 层,所以这套过程会连续重复 6 次。最后一层之后还有一次 LayerNorm,模型再用词表投影把每个 768 维向量映射成 50,257 个 logits:

1
2
最后一层输出:       (1, 4, 768)
词表投影后的 logits: (1, 4, 50257)

这一步只在第 6 层 Transformer 完成之后做一次,不会每经过一层就做一次词表投影。前面 6 层一直在把每个位置的 768 维向量加工得更适合预测;最后才统一把它们变成词表分数。

这次输入有 4 个位置,所以最后一次投影会同时得到 4 行 logits:

1
2
3
4
'The'  这一行 → 对应「下一个应是 ' cat'」的预测
' cat' 这一行 → 对应「下一个应是 ' sat'」的预测
' sat' 这一行 → 对应「下一个应是 ' on'」的预测
' on' 这一行 → 用来选择第 5 个 token

训练时,这 4 行都会分别和真实的下一个 token 对照来计算误差;但现在是在生成文本,所以只需要最后一行 ' on' 的结果。

50,257 正是前面打印的 tokenizer.vocab_size:DistilGPT-2 词表中一共有 50,257 个 token ID。词表投影的最后一维因此不是「50,257 个词」,而是 ID 从 050,256 的 50,257 个候选 token;每一个候选项都有一个 logits 分数。

这里每个 logits 是一个 token 的原始分数,还不是概率;需要 softmax 才会变成 50,257 个概率。对于 The / cat / sat / on,第 4 个位置 ' on' 的那一行 logits 才用于从这 50,257 个候选项中选择第 5 个 token。

把这一行按概率从大到小取前 5 个,就能直接看到模型此刻认为最可能接在后面的 token 是什么:

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
18
19
import torch  # softmax、topk 等张量运算

with torch.no_grad(): # 这里只做推理,不需要为反向传播保存中间结果
prefill = model(input_ids, use_cache=True) # 一次处理 4 个输入 token,得到每个位置的 logits 和 KV Cache

last_logits = prefill.logits[0, -1] # 取第 1 条输入的最后位置(' on'):形状是 (50257,)
probabilities = torch.softmax(last_logits, dim=-1) # 把 50,257 个原始分数变成总和为 1 的概率
top_probs, top_ids = torch.topk(probabilities, k=5) # 找出概率最高的 5 个概率,以及它们各自的 token ID

for token_id, probability in zip(top_ids.tolist(), top_probs.tolist()): # 一次取一个候选 ID 和它的概率
token = tokenizer.decode([token_id]) # 把 ID 还原成人能读到的 token 文本
print(f"ID={token_id:>5}, token={token!r}, probability={probability:.4f}") # 打印这一个候选

next_id = top_ids[0].view(1, 1) # 贪心策略:5 个候选中直接取概率最高的第 1 名,并恢复成模型输入所需的 (1, 1)
next_text = tokenizer.decode(next_id[0]) # 将选中的 ID 解码为文本
print("贪心选择:", next_id.item(), repr(next_text)) # 显示最终追加到句子后的 token

# 本次运行排第 1 的是:
# 贪心选择: 262 ' the'

所以「可能的 token」一开始其实是整张词表中的 50,257 个 ID;上面的循环会打印概率最高的 5 个。这里用的是贪心策略,直接选择其中概率最高的 ID 262,再经 tokenizer 解码成文本 ' the'。原始文本于是变成 The cat sat on the。如果改用 temperature、top-p 或随机采样,仍然从同一行概率中选,但不一定每次都选第 1 名。

第一次处理完整输入时,KV Cache 出现了

先把完整提示词送进模型的阶段,叫 prefill。这里输入是 4 个 token,模型会并行处理整段 The / cat / sat / on;但 causal mask 仍保证每个位置只能读取左侧和自己。

prefill 建立 KV Cache,decode 只计算新 token 并追加 K V

prefill 的副产品就是 past_key_values。把它打印出来:

1
2
3
4
5
6
7
8
9
10
cache = prefill.past_key_values

for layer, (key, value) in enumerate(cache):
print(f"layer {layer}: K{tuple(key.shape)}, V{tuple(value.shape)}")

# 输出:6 层 transformer
# layer 0: K(1, 12, 4, 64), V(1, 12, 4, 64)
# layer 1: K(1, 12, 4, 64), V(1, 12, 4, 64)
# ...
# layer 5: K(1, 12, 4, 64), V(1, 12, 4, 64)

每层的形状 (1, 12, 4, 64) 分别表示:

1
2
3
4
1   :batch size,一次只处理 1 条输入
12 :12 个 attention head
4 :已经处理过的 4 个 token
64 :每个 head 的 K/V 向量宽度

因此,KV Cache 不是一句对话的摘要,也不是只存一份数据。它在 6 层 × 12 个 head 中分别保存历史 token 的 Key 和 Value。之后的 token 会继续拿自己的 Query 去匹配这些 Key,并按权重读取这些 Value。

这里为什么只缓存 K 和 V,不缓存 Q?因为新的 token 只会使用自己的 Q发起一次新查询;历史 token 的 Q 已经完成使命,未来不会再拿它来查询。

新 Token 到来时,KV Cache 到底省掉了什么

先把时序固定住。prefill 已经处理完 The / cat / sat / on,并从最后一行 logits 选出了第 5 个 token:' the'。这时文本看起来已经是:

1
The / cat / sat / on / the

但缓存里还只有前 4 个位置的 K/V;' the' 刚被选出来,还没有经过 6 层 Transformer。因此还要把它送回模型一次。这次调用不是为了再预测 ' the',而是处理 ' the',并预测第 6 个 token。

1
2
3
4
5
6
7
8
9
10
11
12
13
14
15
16
17
cached = model(
next_id, # 这次只输入新选出的 ' the',形状是 (1, 1)
past_key_values=cache, # 同时交给模型前 4 个位置、每一层的 K/V
use_cache=True, # 让模型把 ' the' 新算出的 K/V 也保存下来
output_attentions=True, # 额外返回 attention 权重,便于观察;正常推理不必打开
)

print("本次模型输入:", tuple(next_id.shape))
print("新缓存长度:", cached.past_key_values[0][0].shape[2])
print("第 0 层权重形状:", tuple(cached.attentions[0].shape))
print("本次 logits 形状:", tuple(cached.logits.shape))

# 输出:
# 本次模型输入: (1, 1)
# 新缓存长度: 5
# 第 0 层权重形状: (1, 12, 1, 5)
# 本次 logits 形状: (1, 1, 50257)

(1, 1) 表示「1 条输入、这次只有 1 个新 token」,不是模型只知道一个 token。

前 4 个 token 的 K/V 已经藏在 cache 里;缓存长度从 4 变成 5,表示 ' the' 的 K/V 已经追加进去。

cached.logits 只有 1 行,是因为本次只处理了 ' the';这一行的 50,257 个分数用来选择第 6 个 token。

以第 0 层、一个 head 为例,这次真正发生的是:

1
2
3
4
5
' the' 的当前向量
→ 只为 ' the' 生成新的 Q、K、V 每条:(1, 12, 1, 64)
→ 新 K/V 接到缓存的 4 个旧 K/V 后面 K/V:(1, 12, 5, 64)
→ 新 Q 与 5 个 Key 匹配,再混合 5 个 Value 权重:(1, 12, 1, 5)
→ 得到 ' the' 在这一层的新向量,继续进入下一层

(1, 12, 1, 5) 在这里的作用只有一个:它说明 ' the' 虽然是单独送进来的 1 个新 token,但每个 head 仍要读取 5 个位置(前 4 个 token 加上 ' the' 自己)。这些权重怎样从 Q、K、V 算出,前面已经展开过;这里继续只看缓存省掉的计算。

现在再看 KV Cache 到底省了什么。没有缓存时,为了处理 ' the' 并预测第 6 个 token,模型要把 The / cat / sat / on / the 五个 token 一起重新送过 6 层。前 4 个 token 明明上一步已经处理过,却还得再走一遍。

有缓存时,模型只让新来的 ' the' 走这 6 层;前 4 个 token 不再重新计算。模型直接拿出它们早已留下的 K/V,给 ' the' 读取。换句话说:KV Cache 省掉的是「重做历史 token」,不是「处理新 token」。

提示词第一次进入模型时,长文本本身就要花更久的 prefill;进入生成阶段后,每生成一个新 token,它都还要和全部历史 K/V匹配并读取。历史从 100 个 token 增加到 10,000 个 token,新 token 要看的资料就从 100 份增加到 10,000 份。这就是长上下文仍然会慢的原因。

Prompt Cache:把不同请求的共同前缀复用掉

上一节的 KV Cache 只服务于同一次连续生成:回答每多一个 token,缓存就多一个位置。Prompt Cache 把同样的复用扩大到不同请求之间,由服务端保存已经处理过的公共提示词前缀。

假设应用每次都带上相同的系统提示词、工具定义和固定文档:

1
2
请求 A = [固定前缀] + [用户问题 A]
请求 B = [固定前缀] + [用户问题 B]

服务端处理请求 A 时,会为 [固定前缀] 完成 prefill,并留下每层的 K/V。处理请求 B 时,如果开头切出来的 token 序列和这个前缀完全相同,服务端就直接复用这些 K/V,只从 [用户问题 B] 开始计算。它缓存的不是回答,也不是「这段话的大意」,而是这段前缀已经算好的 K/V。

Prompt Cache 复用完全相同的 token 前缀

复用会在第一个不同 token 停下。例如改写一个词、插入动态日期、调整工具顺序,甚至某些空格变化,都可能从那个位置起不再命中;前面相同的部分仍可复用,后面则重新 prefill。因此稳定且较长的内容适合放在提示词前面,用户问题、时间和临时上下文放在后面。

这和用户修改提示词后的普通请求并不冲突:新请求不会回头修改上一轮生成的 KV Cache;服务端只是在新请求开始时,尝试从自己保存的 Prompt Cache 中找到共同前缀。

一句话区分:KV Cache 是一轮回答内部、随着生成不断增长的缓存;Prompt Cache 是服务端在多次请求间复用相同开头的缓存。

服务商如果提供「缓存读取输入」这一档,命中的固定前缀会按缓存读取计价,单价通常低于普通输入。网上也有人分享因为 Prompt Cache 吃大亏的:

某团队的客服 Agent 每天处理 10 万次对话,原本一切正常。某天工程师为了让 Agent “知道”当前时间,在系统提示词里加了一行 Current time: ,把时间戳实时注入进去。第二天监控告警:所有对话的首 token 延迟从 0.5 秒涨到 3-5 秒,月度推理账单几乎翻了一倍。

Current time 值每次在变,导致它之后的内容每次能无法命中 Prompt Cache。 所以设计 Agent 提示词时,应尽量把不变的内容放在前面,把频繁变化的信息放在后面。

回到一次真实的大模型 API 调用,用户在聊天框里输入一句话后,应用通常会把系统提示词、工具定义、历史消息和这句新消息一起组织成 API 请求,通过网络发给服务端。客户端一般发送的是文本或 messages 结构,不需要自己管理 token ID、K/V 或 past_key_values

1
2
3
4
5
6
7
8
用户输入
→ 客户端拼出 prompt / messages,并调用 API
→ 服务端应用模型模板,tokenizer 切成 token
→ 检查开头是否命中 Prompt Cache
→ prefill:处理未命中的提示词部分,建立本次请求的 KV Cache
→ decode:每次生成 1 个 token,并把它的 K/V 追加到本次 KV Cache
→ 将 token 解码为文本,持续或一次性返回给客户端
→ 客户端显示给用户

若 API 开启流式返回,用户看到的文字会一小段一小段出现;底层仍是服务端不断生成 token,再把已生成的文本片段推送回来。每一步选择哪个 token,还会受 temperature、top-p 等采样参数影响,但不影响 KV Cache 的工作方式。

本次回复结束后,这一轮连续生成使用的 KV Cache 通常随请求结束而释放;Prompt Cache 是否还能被下一次请求命中,则取决于共同前缀和服务商自己的缓存策略。

到这里可以把三层东西分开看:

层次 里面有什么 通常谁负责
可推理的模型包 训练好的浮点参数、模型结构配置、tokenizer 模型发布者
推理工程 加载权重、执行 attention、管理 KV Cache、批处理、流式 API 服务商或部署者
应用代码 组织 messages、调用模型、处理回复、接工具或业务数据 应用开发者

所谓「训练出来的模型」,核心是一大批被训练调整过的数字:embedding 表、位置 embedding、LayerNorm 的 γ/β、每层的 W_QKV、输出投影、MLP 参数,以及最后把隐藏向量映射到词表分数的参数。它们一般保存在 .safetensors 等权重文件中。模型配置(层数、hidden size、head 数等)和 tokenizer 虽然不是训练得到的参数,但让这些参数能被正确解释和执行;三者合在一起,才是一份可推理的模型包。

「开源模型」这个说法容易太笼统,至少要分三层看:

1
2
3
开放权重:可以下载训练好的参数、配置和 tokenizer,在本地推理。
开放代码:可以查看或使用模型、推理或训练相关的实现代码。
开放训练过程:还公开训练数据、数据处理、训练配方和日志。

一个模型可能只开放权重,不等于数据和完整训练过程也开放;实际使用前还需要查看它的许可证和使用限制。

拿到开放权重后,部署并不是重新实现 Q/K/V。一般要做的是:选择一个推理运行时,把权重和 tokenizer 加载到足够的 CPU/GPU 内存中,启动一个能接收请求的服务。运行时负责把矩阵计算真正跑起来,并通常负责批处理、调度、KV Cache 显存管理、上下文长度限制、采样和流式输出。是否支持 Prompt Cache、怎样淘汰缓存,则属于这层运行时或服务端的实现。

如果只是调用云端 API,应用开发者通常不需要写 KV Cache,也不需要自己部署模型;需要写的是业务侧代码:怎样拼系统提示词和历史消息、怎样把用户输入发给 API、怎样流式显示回复、怎样处理工具调用、检索结果、权限和异常。如果自己部署开放模型,除了这些业务代码,还要负责模型服务的硬件、启动、扩缩容、监控和安全等工程工作。

windliang wechat