LLM 推理优化技术:KV Cache

1. 以 Decode-only 为基础的 LLM

1.1. Transformer 网络

Transformer[1] 网络是大模型的基础,Transformer 网络由多层编码器和解码器堆叠而成,Transformer 的网络结构如下图所示:

在 Transformer 网络结构中,编码器和解码器都是由多头注意力和前馈网络模块组成,在多头注意力机制中,最重要的是注意力的计算,其一般化的计算公式为:

Attention(Q,K,V)=softmax(QKdk)V\mathrm{Attention}(Q,K,V)=\mathrm{softmax}(\frac{QK^\top}{\sqrt{d_k}})V

其中Q,K,VQ,K,V分别是查询、键、值向量。前馈网络通常由两个线性层和非线性激活组成。

在 Transformer 网络结构中共包含了三种类型的 Attention:

  • Self-Attention:在 Transformer 的编码器部分,其中 Q,K,VQ,K,V 均来自于编码器部分的输入;
  • Masked Self-Attention:在 Transformer 的解码器部分,其中 Q,K,VQ,K,V 均来自于解码器部分的输入,与 Self-Attent 不同的是,Masked Self-Attention 并不能看到所有的文本,这个在后面的计算过程部分会详细介绍到;
  • Cross-Attention:这是 Transformer 的编码器和解码器结合的部分,其中 QQ 来自于解码器的输入,而 Q,K,VQ,K,V 来自于编码器的输出;

目前绝大部分的 LLM 使用的是 Decoder-only Transformer,与 Transformer 结构不同的是去除了编码器部分,既然去除了编码器部分,那么 Cross-Attention 也没有必要存在了,LLM 的网络结构如下图所示:

其中 xN 表示的是 N 次重复这个部分。

1.2. 计算过程

大模型的推理过程也被称为自回归过程,模型上一个 Token 的输出会被当作是下一个 Token 的输入。

我们以“大模型的推理”为初始的 Token 序列,完整的序列为“大模型的推理过程也被称为自回归过程”,初始的 Token 序列在 LLM 中也被称为 Prompt。为了方便描述 LLM 的推理过程,我们简化模型,我们取上述图中的 N 为 1,同时,Token 为每一个中文的词,实际的过程中会根据不同的分词方法而不同。此时的计算过程如下图所示:

如图,假设 Embedding 的维度为 6,Prompt 经过 Embedding 后(包括了 Positional Embedding)后,Embedding 的大小为 6×66\times 6,每一个 Token 的向量经过与 WQ,WK,WVW_Q,W_K,W_V 三个矩阵运算后便得到每一个 Token 的 Q,K,VQ,K,V 向量,为了与 Embedding 的向量维度区分开,设置这个的 WQ,WK,WVW_Q,W_K,W_V 的矩阵维度为 6×86\times 8,经过 Masked multi-head Attention 后,便得到上图中的 Attention 矩阵,再经过 FFN 后,得到最终的 Embedding,根据最后一个 Embedding 得到最终的词,核心部分的 Attention 的计算过程如下:

在 LLM 中,使用的是 Masked Multi-Head Attention,简单来说,就是在用“大”这个 Token 去预测 Token “模”时,不能看到后面的 Token,一般都会使用掩码矩阵的方式修改上述的 Attention 计算公式:

Maksed_Attention(Q,K,V)=softmax(QKdk+M)V\mathrm{Maksed\_Attention}(Q,K,V)=\mathrm{softmax}(\frac{QK^\top}{\sqrt{d_k}}+M)V

其中,矩阵 M 称为掩码矩阵,格式为:

M=(000000000000000000000)M=\begin{pmatrix} 0 & -\infty & -\infty & -\infty & -\infty & -\infty \\ 0 & 0 & -\infty & -\infty & -\infty & -\infty \\ 0 & 0 & 0 & -\infty & -\infty & -\infty \\ 0 & 0 & 0 & 0 & -\infty & -\infty \\ 0 & 0 & 0 & 0 & 0 & -\infty \\ 0 & 0 & 0 & 0 & 0 & 0 \\ \end{pmatrix}

掩码矩阵是一个下三角矩阵,这样经过 softmax 计算后,对角线以上的部分都为 0。

预测完“过”这个 Token 后,将其继续放入到输入中,继续预测 token“程”,这里只显示 Attention 部分计算的差异,如下图所示:

Q,K,VQ,K,V 中都增加了对应 Token“过”的向量,依次类推,直到生成结束。

2. KV Cache

KV Cache 在已经成为了 LLM 推理过程中必须要用到的一个优化技术。KV Cache 的想法很简单,即通过把中间结果缓存下来,避免重复计算。

2.1. 删除无效计算

首先,我们需要看一下上述的计算过程中是否存在重复的计算。来看下预测完一个词后的前后对比:

实际上,最终的 Classifier 只会用到 Embedding 的最后一行,而最后一行是 Attention 的最后一行经过 FFN 运算后得到的。因此是不是每次只需要得到上图中 Attention 矩阵的最后一行就可以了。经过简单的数学运算,就能得到如下的图:

上图精简后的图中,虚线框出来的那部分掩码矩阵其实已经没有作用了,留着只是为了上下的对比。从上图,我们发现,其实我们关注的只有三个东西:

  • 新增 Token 的 QQ 向量
  • 完整 Input 的 KK 矩阵
  • 完整 Input 的 VV 矩阵

2.2. 缓存重复计算

上面也提到,每一个 Token 的 Q,K,VQ,K,V 向量是由该 Token 的 Embedding 向量与 WQ,WK,WVW_Q,W_K,W_V 三个矩阵运算后得到的,每一个 Token 会得到该 Token 对应的 Q,K,VQ,K,V,但是上面的计算过程依赖于完整 Input 的 KK 矩阵和 VV 矩阵,因此,这里需要将该新增 Token 之前所有 Token 对应的 K,VK,V 向量缓存起来,这便是 KV Cache 了。也就是下图的红色部分:

3. 总结

通过以上的分析,LLM 在推理阶段,其实用到的是最终 Embedding 矩阵的最后一行,而影响最后一行的最关键的计算是 Attention 的计算,进而在 Attention 中与其相关的是新增 Token 的 QQ 向量,完整 Input 的 KK 矩阵以及完整 Input 的 VV 矩阵,因此,我们可以将该新增 Token 之前所有 Token 对应的 K,VK,V 向量缓存起来,这便形成了 KV Cache。

参考文献

[1] Vaswani A , Shazeer N , Parmar N ,et al.Attention Is All You Need[J].arXiv, 2017.DOI:10.48550/arXiv.1706.03762.