预备知识:Attention 与 Transformer

在标准的 Transformer 中,对于输入序列 XRN×dX \in \mathbb{R} ^ {N \times d},有计算公式:

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

整个计算过程可以这样理解:

  1. 输入序列 XX 经过三组线形投影,用三个可学习的矩阵计算出查询矩阵 Q=XWQQ=XW^Q,键矩阵 K=XWKK=XW^K,值矩阵 V=XWKV=XW^K

  2. 每一次做查询时,用 QQKK 做点乘来查自己想“问”的索引,得到 Query 对各个 Key 的分数,用 SoftmaxSoftmax 做归一之后得到对各个 Key 的概率;

  3. 对概率进行采样。有四种常见的采样策略:

    • Greedy:永远选 logits 最大的 token;
    • Temperature:温度越低越保守,越高越随机;
    • Top-k:保留分数最高的 k 个候选;
    • Top-p:只保留概率达到 p 的集合。
  4. 采样后的 Key 与 Value 矩阵做点乘得到查询的输出。

对 Attention 头进行拼装就得到 Tranformer:

左侧部分是 Encoder。Prompt 经过 NN 层特征提取层,输出一系列含完整语义的特征向量。这些向量经过变换产生 KKVV,传递给 Decoder 以供后续查询,不需要再重复计算。

右侧部分是 Decoder。一开始右下角的 “Outputs” 只有一个代表开始的起始标记,这个标记把整个句子向后推一格(Shifted Right)。第一次输入这个起始符,经过嵌入层后作为当前位置的 Query,与之前 Encoder 编码出来的 Key 做交叉注意力,通过查阅之前 Key 的信息得到注意力矩阵,计算出下一个 token 的预测。

假设用 Transformer 来做翻译任务。Inputs (Prompt)是 “一个苹果”,它先用 Encoder 将这句话拆成两个 token [i_t0, i_t1]并映射成 KKVV,分别对应 “一个”,“苹果”。一开始的 “Outputs” 由于做了 Shifted Right,里面只有一个初始 token [<START>]。经过一轮计算,输出了一个o_t0(如此处应是一个 “An”),现在 Outputs 变成 [<START>, o_t0],进入 Decoder 生成下一个预测的 Query,以此类推完成接下来的翻译。

KV Cache

计算流程

KV Cache 的核心是缓存历史中的 Key 和 Value,来避免重复计算。以序列长度为 3 的输入为例来拆解原始计算步骤。另外,注意力分数需要加 Mask,确保先输出的 token 不会看到后输出的 token(相当于只计算或只保留下三角部分)。

在没有引入 KV Cache 时:

  • Step1:O=[o0]O=\begin{bmatrix}o_0\end{bmatrix} 处理后产生 q0,k0,v0q_0, k_0, v_0

    计算 QK=[q0][k0]=[s00]QK^\top=\begin{bmatrix}q_0\end{bmatrix}\begin{bmatrix}k_0\end{bmatrix}=\begin{bmatrix}s_{00}\end{bmatrix},做 Softmax 得到 [p00]\begin{bmatrix}p_{00}\end{bmatrix},采样之后乘 [v0]\begin{bmatrix}v_0\end{bmatrix},得到输出 o1o_1,追加到 OO 后面;

  • Step2:O=[o0o1]O=\begin{bmatrix}o_0\\o_1\end{bmatrix} 重新处理产生完整的 [q0q1],[k0k1],[v0v1]\begin{bmatrix}q_0\\q_1\end{bmatrix}, \begin{bmatrix}k_0\\k_1\end{bmatrix}, \begin{bmatrix}v_0\\v_1\end{bmatrix}

    计算 QK=[q0q1][k0k1]=[s00s01s10s11]QK^\top=\begin{bmatrix}q_0\\q_1\end{bmatrix}\begin{bmatrix}k_0&k_1\end{bmatrix}=\begin{bmatrix}s_{00}&s_{01}\\s_{10}&s_{11}\end{bmatrix},做 Softmax 得到 [p00p01p10p11]\begin{bmatrix}p_{00}&p_{01}\\p_{10}&p_{11}\end{bmatrix},采样之后乘 [v0v1]\begin{bmatrix}v_0\\v_1\end{bmatrix},得到输出结果的最末项 o2o_2,追加到 OO 后面;

  • Step3:O=[o0o1o2]O=\begin{bmatrix}o_0\\o_1\\o_2\end{bmatrix} 重新处理产生完整的 [q0q1q2],[k0k1k2],[v0v1v2]\begin{bmatrix}q_0\\q_1\\q_2\end{bmatrix}, \begin{bmatrix}k_0\\k_1\\k_2\end{bmatrix}, \begin{bmatrix}v_0\\v_1\\v_2\end{bmatrix}

    计算 QK=[q0q1q2][k0k1k2]=[s00s01s02s10s11s12s20s21s22]QK^\top=\begin{bmatrix}q_0\\q_1\\q_2\end{bmatrix}\begin{bmatrix}k_0&k_1&k_2\end{bmatrix}=\begin{bmatrix}s_{00}&s_{01}&s_{02}\\s_{10}&s_{11}&s_{12}\\s_{20}&s_{21}&s_{22}\end{bmatrix},做 Softmax 得到 [p00p01p02p10p11p12p20p21p22]\begin{bmatrix}p_{00}&p_{01}&p_{02}\\p_{10}&p_{11}&p_{12}\\p_{20}&p_{21}&p_{22}\end{bmatrix},采样之后乘 [v0v1v2]\begin{bmatrix}v_0\\v_1\\v_2\end{bmatrix},得到输出结果的最末项 o3o_3,追加到 OO 后面。

在第 tt 步计算中,前面已经计算出的 {v0,,vt1}\{v_0, \dots, v_{t-1}\}{k0,,kt1}\{k_0, \dots, k_{t-1}\},以及 QKTQK^T 分数矩阵左上角的 (t1)×(t1)(t-1) \times (t-1) 方阵,都不需要重复计算。不使用 Cache 时,每一次预测都把提示词与前缀输出序列全部送入 Decoder 做 forward,对于长度为 NN 的序列,总体计算开销会显著增加。

引入 KV Cache 后:

  • Step1: 仅输入 O=[o0]O=\begin{bmatrix}o_0\end{bmatrix} 处理产生 q0,k0,v0q_0, k_0, v_0。将 k0,v0k_0, v_0 写入缓存,此时K=[k0],V=[v0]K^\top=\begin{bmatrix}k_0\end{bmatrix}, V=\begin{bmatrix}v_0\end{bmatrix}
    计算 QK=[q0][k0]=[s00]QK^\top=\begin{bmatrix}q_0\end{bmatrix}\begin{bmatrix}k_0\end{bmatrix}=\begin{bmatrix}s_{00}\end{bmatrix},做 Softmax 得到 [p00]\begin{bmatrix}p_{00}\end{bmatrix},采样之后乘 [v0]\begin{bmatrix}v_0\end{bmatrix},得到输出 o1o_1,追加到 OO 后面;

  • Step2: 仅输入最新的 o1o_1 处理产生 q1,k1,v1q_1, k_1, v_1。追加更新缓存,此时K=[k0k1],V=[v0v1]K^\top=\begin{bmatrix}k_0&k_1\end{bmatrix}, V=\begin{bmatrix}v_0\\v_1\end{bmatrix}
    计算 QK=[q1][k0k1]=[s10s11]QK^\top=\begin{bmatrix}q_1\end{bmatrix}\begin{bmatrix}k_0&k_1\end{bmatrix}=\begin{bmatrix}s_{10}&s_{11}\end{bmatrix}注:这刚好是分数矩阵第二行的有效下三角部分,天然避开了原本需要 Mask 的 s01s_{01}),做 Softmax 得到 [p10p11]\begin{bmatrix}p_{10}&p_{11}\end{bmatrix},采样之后乘 [v0v1]\begin{bmatrix}v_0\\v_1\end{bmatrix},直接得到输出 o2o_2,追加到 OO 后面;

  • Step3: 仅输入最新的 o2o_2 处理产生 q2,k2,v2q_2, k_2, v_2。追加更新缓存,此时K=[k0k1k2],V=[v0v1v2]K^\top=\begin{bmatrix}k_0&k_1&k_2\end{bmatrix}, V=\begin{bmatrix}v_0\\v_1\\v_2\end{bmatrix}
    计算 QK=[q2][k0k1k2]=[s20s21s22]QK^\top=\begin{bmatrix}q_2\end{bmatrix}\begin{bmatrix}k_0&k_1&k_2\end{bmatrix}=\begin{bmatrix}s_{20}&s_{21}&s_{22}\end{bmatrix}注:这刚好是分数矩阵第三行的有效下三角部分),做 Softmax 得到 [p20p21p22]\begin{bmatrix}p_{20}&p_{21}&p_{22}\end{bmatrix},采样之后乘 [v0v1v2]\begin{bmatrix}v_0\\v_1\\v_2\end{bmatrix},直接得到输出 o3o_3,追加到 OO 后面。

这使得 Decode 阶段对于长度为 NN 的序列,单步计算复杂度从的 O(N2)O(N^2) (计算 N×NN \times N 矩阵)降至线性的 O(N)O(N) (计算大小为 NN 的一行),是典型的用空间换时间。

Prompt Prefill

Transformer 在实际工作时还会有提示词输入,这种情况下怎么将提示词的信息呈现给模型呢?回到之前的 Transformer 架构,左侧的 Decoder 会有一个 Input 入口,输入的内容会被传递给 Decoder 处理生成对应的 Key 和 Value。

那么在有提示词的情况下如何让 Decoder 访问提示词的信息呢?直接将提示词塞给 Eecoder 处理后传给 Decoder 就可以了。生成的 Key 和 Value 预填充进 KV Cache 中,后面直接参与 Attention 计算,就可以让每一次输入的 oio_i 都可以看到提示词的信息了。

仍然走一遍计算流程,在之前的基础上,我们加上长度同样为 3 的提示词序列:

  • Step 1:预填充阶段,完整输入 Prompt 序列 X=[x0x1x2]X=\begin{bmatrix}x_0\\x_1\\x_2\end{bmatrix},一次性并行处理产生完整的查询、键、值矩阵:
    Q=[q0q1q2],K=[k0k1k2],V=[v0v1v2]Q=\begin{bmatrix}q_0\\q_1\\q_2\end{bmatrix}, \quad K^\top=\begin{bmatrix}k_0&k_1&k_2\end{bmatrix}, \quad V=\begin{bmatrix}v_0\\v_1\\v_2\end{bmatrix}

    建立初始缓存,将全量的 KKVV 写入缓存。

    Kcache=[k0k1k2],Vcache=[v0v1v2]K^\top_{cache}=\begin{bmatrix}k_0&k_1&k_2\end{bmatrix}, \quad V_{cache}=\begin{bmatrix}v_0\\v_1\\v_2\end{bmatrix}

    计算 QKcache=[q0q1q2][k0k1k2]=[s00s10s11s20s21s22]QK^\top_{cache}=\begin{bmatrix}q_0\\q_1\\q_2\end{bmatrix}\begin{bmatrix}k_0&k_1&k_2\end{bmatrix}=\begin{bmatrix}s_{00}&-\infty&-\infty\\s_{10}&s_{11}&-\infty\\s_{20}&s_{21}&s_{22}\end{bmatrix}

    做 Softmax 得到 [p0000p10p110p20p21p22]\begin{bmatrix}p_{00}&0&0\\p_{10}&p_{11}&0\\p_{20}&p_{21}&p_{22}\end{bmatrix},乘 Vcache[v0v1v2]V_{cache}\begin{bmatrix}v_0\\v_1\\v_2\end{bmatrix},得到隐藏状态 H=[h0h1h2]H=\begin{bmatrix}h_0\\h_1\\h_2\end{bmatrix}

    提取最后一行包含全量上下文的 h2h_2,送入输出层,得到生成的第 1 个回答 Token o3o_3,追加到序列末尾。

  • Step 2:完成 Prefill 后,切换到 Decoder 递推模式。此时不需要再重新计算 x0,x1,x2x_0, x_1, x_2 对应的历史 Key/Value,只需要提取缓存;由于未来的词还未生成,当前步对历史 token 的注意力不再需要显式的完整 Mask 矩阵。
    仅输入最新的 o3o_3 处理产生 q3,k3,v3q_3, k_3, v_3。追加缓存,此时:

    Kcache=[k0k1k2k3],Vcache=[v0v1v2v3]K^\top_{cache}=\begin{bmatrix}k_0&k_1&k_2&k_3\end{bmatrix}, \quad V_{cache}=\begin{bmatrix}v_0\\v_1\\v_2\\v_3\end{bmatrix}

    计算 QKcache=[q3][k0k1k2k3]=[s30s31s32s33]QK^\top_{cache}=\begin{bmatrix}q_3\end{bmatrix}\begin{bmatrix}k_0&k_1&k_2&k_3\end{bmatrix}=\begin{bmatrix}s_{30}&s_{31}&s_{32}&s_{33}\end{bmatrix}

    做 Softmax 得到 [p30p31p32p33]\begin{bmatrix}p_{30}&p_{31}&p_{32}&p_{33}\end{bmatrix},乘 Vcache[v0v1v2v3]V_{cache}\begin{bmatrix}v_0\\v_1\\v_2\\v_3\end{bmatrix},直接得到隐藏状态 h3h_3,送入输出层,得到生成的第 2 个回答 Token:o4o_4,追加到序列末尾。

  • Step 3:仅输入最新的 o4o_4 处理产生 q4,k4,v4q_4, k_4, v_4。追加缓存,此时:

    Kcache=[k0k1k2k3k4],Vcache=[v0v1v2v3v4]K^\top_{cache}=\begin{bmatrix}k_0&k_1&k_2&k_3&k_4\end{bmatrix}, \quad V_{cache}=\begin{bmatrix}v_0\\v_1\\v_2\\v_3\\v_4\end{bmatrix}

    计算 QKcache=[q4][k0k1k2k3k4]=[s40s41s42s43s44]QK^\top_{cache}=\begin{bmatrix}q_4\end{bmatrix}\begin{bmatrix}k_0&k_1&k_2&k_3&k_4\end{bmatrix}=\begin{bmatrix}s_{40}&s_{41}&s_{42}&s_{43}&s_{44}\end{bmatrix}

    做 Softmax 得到 [p40p41p42p43p44]\begin{bmatrix}p_{40}&p_{41}&p_{42}&p_{43}&p_{44}\end{bmatrix},乘 Vcache[v0v1v2v3v4]V_{cache}\begin{bmatrix}v_0\\v_1\\v_2\\v_3\\v_4\end{bmatrix},直接得到隐藏状态 h4h_4,送入输出层,得到生成的第 3 个回答 Token:o5o_5,追加到序列末尾。

Prefill 过程偏计算密集,其一次性计算大量输入 token 的矩阵乘,尤其在使用 Full Attention 的时候有 O(N2)O(N^2) 复杂度;Decode 过程则偏访存密集,每一步都需要读取模型权重和更新 KV Cache。

Linear Attention

当使用 KV Cache 时,缓存会随着序列输出的不断追加线形增长。为了避免重复计算,所有的历史 KVKV 都会被完整保存在显存中。在有限的显存下,不可能让 Cache 无限增长,因此要找到一种方法对历史进行压缩。

标准的 Full Attention 中,完整的输出公式是这样的:

oi=j=1Nexp(qikj)mexp(qikm)vjo_i = \sum_{j=1}^N \frac{\exp(q_i k_j^\top)}{\sum_m \exp(q_i k_m^\top)} v_j

由于 Softmax 使用指数函数进行计算,qiq_ikjk_j 被绑定在指数函数内部,因此必须先算出整个 N×NN \times N 的注意力矩阵,才能去乘 VV

Linear Attention 的核心思想是:用一个非负的核函数 ϕ()\phi(\cdot) 来近似或替代指数函数,消除 qqkk 的非线性耦合,变成一个可以提取出来的乘积形式。现在使用 Linear Attention 改写输出公式,将分子替换为 Sim(qi,kj)=ϕ(qi)ϕ(kj)\text{Sim}(q_i, k_j) = \phi(q_i) \phi(k_j)^\top,那么没有规范化的 logits 的计算公式变成:

ot=j=1t(ϕ(qt)ϕ(kj))vjo_t = \sum_{j=1}^t \left( \phi(q_t) \phi(k_j)^\top \right) v_j^\top

根据结合律改变括号的位置:

ot=ϕ(qt)(j=1tϕ(kj)vj)o_t = \phi(q_t) \left( \sum_{j=1}^t \phi(k_j)^\top v_j^\top \right)

令中间那串累加项为状态矩阵 StS_t

St=j=1tϕ(kj)vj=St1+ϕ(kt)vtS_t = \sum_{j=1}^t \phi(k_j)^\top v_j^\top = S_{t-1} + \phi(k_t)^\top v_t^\top

核函数可以设置为 ReLU/ELUReLU/ELU。进入注意力计算之前,就将核函数执行,即 ϕ(kt)kt\phi(k_t) \to k_t,注意力计算的核心式简化为:

St=St1+ktvtS_t = S_{t-1} + k_t v_t^\top

现在新 token 来时只需要更新固定状态 StS_t,再用 qtq_t 与之相乘进行读取即可。缓存变成了一个固定 d×dd \times d 大小的矩阵 SiS_i,Linear Attention 机制类似 RNN:每来一个 Token,就用当前的 kt,vtk_t, v_t 更新 StS_t,然后用 qtq_t 读取 StS_t。但是,原本随序列长度线性增长的信息被压缩到固定大小矩阵后,会形成有损表示,因此在面对更长文本时可能像 RNN 一样丢失或覆盖历史信息。相较于 KV Cache “用空间换效率”,Linear Attention 更接近“用精度换空间”。

Gated Delta Net

在处理信息丢失和覆盖的的问题时,LSTM 在 RNN 的基础上引入了门控机制,用输入门和遗忘门来选择控制 Memory Cell 更新时的旧记忆保留和新记忆写入:

ct=ftct1+itgtc_t​=f_t​⊙c_{t−1}​+i_t​⊙g_t

通过调整这些门的值,LSTM 能够灵活地保留或更新信息,从而更好地捕捉长期依赖关系,缓解固定维度向量记忆的丢失。

类似的,Gated Data Net 也引入了门控机制。对于上一步留下的状态矩阵 St1S_{t-1},为其乘上一个遗忘门 αt\alpha_t,这个值由输入动态生成,介于 [0,1][0, 1] 之间,用来决定以何种程度保留上一步的记忆。这个值趋于 1 时,记忆被完整保存,反之则抹除。

对于增量项,则引入了 Delta Rule。首先定义预测误差 et=vtSt1Tkte_t=v_t-S_{t-1}^Tk_t。它利用现有的记忆和当前的键 ktk_t 尝试去预测对应的 value,预测误差就是预测结果与实际值之间的误差。

为什么用 St1TktS_{t-1}^Tk_t 来做 value 的预测?先回到 Linear Attention 的状态更新公式 St=j=1tkjvjS_t = \sum_{j=1}^t k_j v_j^\top,此时的 StS_t 是一个 d×dd \times d 的矩阵。我们现在站在第 tt 步,用当前的键向量 ktk_t 乘上 StS_t^\top,得到:

Stkt=(j=1tvjkj)kt=j=1tvj(kjkt)Stkt=vt(ktkt)+j=1t1vj(kjkt)S_t^\top k_t = \left( \sum_{j=1}^t v_j k_j^\top \right) k_t = \sum_{j=1}^t v_j (k_j^\top k_t)S_t^\top k_t = v_t (k_t^\top k_t) + \sum_{j=1}^{t-1} v_j (k_j^\top k_t)

而在高维向量空间中,两个随机生成的向量很大概率是接近正交的,又考虑模型中一般都会做向量归一化,就有 ktkt1k_t^\top k_t \approx 1,且当当前的新键与之前的某个旧键没有关联时, kjkt0(jt)k_j^\top k_t \approx 0(j \ne t)。把这两个式子代进之前的公式,就能得到:

Stktvt1+vj0=vtS_t^\top k_t \approx v_t \cdot 1 + \sum v_j \cdot 0 = v_t

证明了用 St1ktS_{t-1}k_t 来预测 value 的合理性, 再回到预测误差。得到预测值后,我们得到一个用于修正 value 预测误差的量 ktetT=kt(vtSt1kt)k_te_t^T=k_t (v_t - S_{t-1}^\top k_t)^\top。加上这个补偿的 state 可以准确地预测出当前时间步 key 对应的 value。考虑训练时,定义一个目标是让当前的记忆能精准读出当前 value 的损失函数:

Lt(S)=12Sktvt22\mathcal{L}_t(S) = \frac{1}{2} \| S^\top k_t - v_t \|_2^2

对矩阵 SS 求偏导:

SLt(S)=kt(Sktvt)=kt(vtSkt)\nabla_S \mathcal{L}_t(S) = k_t (S^\top k_t - v_t)^\top = - k_t (v_t - S^\top k_t)^\top

给定一个动态更新的学习率 βt\beta_t 按照梯度下降的方式更新 SS,有:

St=St1βtSLt(St1)=St1+βtkt(vtSt1kt)S_t = S_{t-1} - \beta_t \nabla_S \mathcal{L}_t(S_{t-1}) = S_{t-1} + \beta_t k_t (v_t - S_{t-1}^\top k_t)^\top

βtkt(vtSt1kt)\beta_t k_t (v_t - S_{t-1}^\top k_t)^\top 就是 Delta Rule 的增量项,其中 βt\beta_t 就是控制记忆写入的写入门。问题又来了,这个训练时的学习率和控制写入的门之间,是否存在等价关系呢?

βt\beta_t 当成一个介于 [0,1][0, 1] 之间的动态学习率,这个学习率用输入的 token 来预测:

βt=Sigmoid(Wβxt)\beta_t = \text{Sigmoid}(W_\beta \cdot x_t)

训练时梯度下降会推着 WβW_\beta 向着使 state 能更准确预测 value 的方向移动,那么最终训练出的动态“学习率”,在推理时也会倾向于去补偿预测误差——预测误差大说明 state 实际变化梯度大,这个学习率会增大(逼近1)来加大写入力度;预测误差小说明 state 实际变化梯度小,这个学习率会降低(逼近0)来降低写入力度。这个网络在推理时固定下来,用来调整写入力度。

至此 Gated Delta Net 的状态更新函数就确定下来了:

St=αtSt1+βtkt(vtSt1kt)S_t = \alpha_t S_{t-1} + \beta_t k_t (v_t - S_{t-1}^\top k_t)^\top

这个函数能够动态调整每个新 token 运算后“记忆”的遗忘程度和写入程度,缓解了固定大小状态矩阵丢失或覆盖历史信息的问题。