如果用一句话描述大模型的推理流程,那就是:

输入文本先变成 token,模型先一次性理解整段输入(Prefill),然后再一个 token 一个 token 地往后生成(Decode),直到遇到结束条件。

然后,我们具体的讲解一下这个流程:

1.将用户的输入转换成Token IDs,然后再转换成向量

比如用户输入“什么是PagedAttention”,模型系统里有一个叫Tokenizer的组件,它会将用户输入的字符串(即“什么是PagedAttention”)切分成token,然后再将每一个token转换成对应的token id。

最后的输出可能类似于[id1, id2, id3, …]

注意:Tokenizer不是固定的。不同模型,不同厂家,甚至同一个厂家的不同模型都可以采用不同的Tokenizer。这就导致token的切分方法,以及对应的token id都不是通用的。

2.根据Embedding表格,将Token IDs转换成对应的向量

token id只是一个编号,一个对于token的编号。相邻编号的token,其语义不一定有关联性。但是大模型需要的不是token的编号,而是token的语义,所以在输入给大模型之前,我们还需要使用embedding将token id转换为更有语义表达的向量,同时加上token的位置信息,方便大模型知道这个token在输入中的位置。

可能有人奇怪,我们转换出来的向量不是按照顺序排列了,为什么还需要显性提供每一个token对应的位置信息?

这是因为大模型所使用的Transformer架构需要位置参数。

3.将token的向量和位置信息传递给Transformer,模型对输入的token做一遍prefill

prefill阶段可以理解为模型一次性将输入的prompt读完并理解,将后面decode需要的KV Cache准备好。注意这里的“一次性”,是因为在prefill阶段可以并行处理所有的prompt token。

讲prefill之前,先要介绍一下Transformer:

3.1Transformer的组成

一个完整的Transformer,通常是由多个Transformer Block堆叠而成。一个简化版的Transformer Block的结构如下所示:

输入 ↓ Self-Attention ↓ Residual + Norm ↓ MLP / FFN ↓ Residual + Norm ↓ 输出给下一层

3.2Prompt token在Transformer层做了哪些工作

我们输入的prompt token会在Self-Attention中分别计算Q (Query), K (Key), V (Value),其中K和V就是常说的KV Cache,这个K和V值会被保留下来作为后面的缓存。

假设,我们输入的prompt token中有5个token,则会计算:

Token 1 → K1 V1 Token 2 → K2 V2 Token 3 → K3 V3 … Token 5 → K5 V5

其中,对于像Qwen/GPT这类的模型,每一个token会关注之前的token,因为 GPT/Qwen这类模型是 Causal Transformer,所以不能偷看未来。所以,Attention 关系大致是:

能看到 t1 t1 t2 t1 t2 t3 t1 t2 t3 t4 t1 t2 t3 t4 t5 t1 t2 t3 t4 t5

不能这样:

t1 →t5

因为 t1 位置不能看到未来。

所以 Attention Matrix 是一个下三角:

​ K1 K2 K3 K4 K5 Q1 ✓ Q2 ✓ ✓ Q3 ✓ ✓ ✓ Q4 ✓ ✓ ✓ ✓ Q5 ✓ ✓ ✓ ✓ ✓

这就是 Causal Mask

3.3Prefill结束时,第一个token即将产生

我们还是假设输入的prompt有5个token,则prefill跑完后,最后一个token (t5)在最后一层Transformer中获得hidden state h5。

然后,经过LM Head这一步(即把 Transformer 最后输出的 hidden vector,转换成“词表里每个 token 的分数”),获得词表中每一个token的分数,叫做logits。这些logits还不是token的概率,还需要继续执行softmax,将分数转换为概率。

再根据sampling策略从这些token中选出下一个输出的token。

4.进入Decode阶段,直到结束

在prefill阶段结束后,输出了第一个token t6后,重新进入下一轮循环:

先计算t6对应的Q, K, V的值,即q6, k6, v6。再通过q6和前面缓存的前5个token的KV值计算Attention值。然后把k6和v6的值放到KV Cache中。最后生成新的token t7。

每一轮生成新的token之后,都会检查:

是不是 EOS(End Of Sequence)? 是不是达到 max_tokens? 是不是遇到 stop token? 是不是遇到 stop string?

如果满足,则终止Decode过程,将输出的token通过tokenizer的反向过程最终输出成需要的结果。