预备知识:Attention 与 Transformer 在标准的 Transformer 中,对于输入序列 X ∈ R N × d X \in \mathbb{R} ^ {N \times d} X ∈ R N × d ,有计算公式:
Attention ( Q , K , V ) = softmax ( Q K ⊤ d k ) V \text{Attention}(Q, K, V) = \text{softmax}\left(\frac{QK^\top}{\sqrt{d_k}}\right)V
Attention ( Q , K , V ) = softmax ( d k Q K ⊤ ) V
整个计算过程可以这样理解:
输入序列 X X X 经过三组线形投影,用三个可学习的矩阵计算出查询矩阵 Q = X W Q Q=XW^Q Q = X W Q ,键矩阵 K = X W K K=XW^K K = X W K ,值矩阵 V = X W K V=XW^K V = X W K ;
每一次做查询时,用 Q Q Q 与 K K K 做点乘来查自己想“问”的索引,得到 Query 对各个 Key 的分数,用 S o f t m a x Softmax S o f t ma x 做归一之后得到对各个 Key 的概率;
对概率进行采样。有四种常见的采样策略:
Greedy:永远选 logits 最大的 token;
Temperature:温度越低越保守,越高越随机;
Top-k:保留分数最高的 k 个候选;
Top-p:只保留概率达到 p 的集合。
采样后的 Key 与 Value 矩阵做点乘得到查询的输出。
对 Attention 头进行拼装就得到 Tranformer:
左侧部分是 Encoder 。Prompt 经过 N N N 层特征提取层,输出一系列含完整语义的特征向量。这些向量经过变换产生 K K K 和 V V V ,传递给 Decoder 以供后续查询,不需要再重复计算。
右侧部分是 Decoder 。一开始右下角的 “Outputs” 只有一个代表开始的起始标记,这个标记把整个句子向后推一格(Shifted Right)。第一次输入这个起始符,经过嵌入层后作为当前位置的 Query,与之前 Encoder 编码出来的 Key 做交叉注意力,通过查阅之前 Key 的信息得到注意力矩阵,计算出下一个 token 的预测。
假设用 Transformer 来做翻译任务。Inputs (Prompt)是 “一个苹果”,它先用 Encoder 将这句话拆成两个 token [i_t0, i_t1]并映射成 K K K 和 V V V ,分别对应 “一个”,“苹果”。一开始的 “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 = [ o 0 ] O=\begin{bmatrix}o_0\end{bmatrix} O = [ o 0 ] 处理后产生 q 0 , k 0 , v 0 q_0, k_0, v_0 q 0 , k 0 , v 0 。
计算 Q K ⊤ = [ q 0 ] [ k 0 ] = [ s 00 ] QK^\top=\begin{bmatrix}q_0\end{bmatrix}\begin{bmatrix}k_0\end{bmatrix}=\begin{bmatrix}s_{00}\end{bmatrix} Q K ⊤ = [ q 0 ] [ k 0 ] = [ s 00 ] ,做 Softmax 得到 [ p 00 ] \begin{bmatrix}p_{00}\end{bmatrix} [ p 00 ] ,采样之后乘 [ v 0 ] \begin{bmatrix}v_0\end{bmatrix} [ v 0 ] ,得到输出 o 1 o_1 o 1 ,追加到 O O O 后面;
Step2:O = [ o 0 o 1 ] O=\begin{bmatrix}o_0\\o_1\end{bmatrix} O = [ o 0 o 1 ] 重新处理产生完整的 [ q 0 q 1 ] , [ k 0 k 1 ] , [ v 0 v 1 ] \begin{bmatrix}q_0\\q_1\end{bmatrix}, \begin{bmatrix}k_0\\k_1\end{bmatrix}, \begin{bmatrix}v_0\\v_1\end{bmatrix} [ q 0 q 1 ] , [ k 0 k 1 ] , [ v 0 v 1 ]
计算 Q K ⊤ = [ q 0 q 1 ] [ k 0 k 1 ] = [ s 00 s 01 s 10 s 11 ] 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} Q K ⊤ = [ q 0 q 1 ] [ k 0 k 1 ] = [ s 00 s 10 s 01 s 11 ] ,做 Softmax 得到 [ p 00 p 01 p 10 p 11 ] \begin{bmatrix}p_{00}&p_{01}\\p_{10}&p_{11}\end{bmatrix} [ p 00 p 10 p 01 p 11 ] ,采样之后乘 [ v 0 v 1 ] \begin{bmatrix}v_0\\v_1\end{bmatrix} [ v 0 v 1 ] ,得到输出结果的最末项 o 2 o_2 o 2 ,追加到 O O O 后面;
Step3:O = [ o 0 o 1 o 2 ] O=\begin{bmatrix}o_0\\o_1\\o_2\end{bmatrix} O = o 0 o 1 o 2 重新处理产生完整的 [ q 0 q 1 q 2 ] , [ k 0 k 1 k 2 ] , [ v 0 v 1 v 2 ] \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} q 0 q 1 q 2 , k 0 k 1 k 2 , v 0 v 1 v 2
计算 Q K ⊤ = [ q 0 q 1 q 2 ] [ k 0 k 1 k 2 ] = [ s 00 s 01 s 02 s 10 s 11 s 12 s 20 s 21 s 22 ] 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} Q K ⊤ = q 0 q 1 q 2 [ k 0 k 1 k 2 ] = s 00 s 10 s 20 s 01 s 11 s 21 s 02 s 12 s 22 ,做 Softmax 得到 [ p 00 p 01 p 02 p 10 p 11 p 12 p 20 p 21 p 22 ] \begin{bmatrix}p_{00}&p_{01}&p_{02}\\p_{10}&p_{11}&p_{12}\\p_{20}&p_{21}&p_{22}\end{bmatrix} p 00 p 10 p 20 p 01 p 11 p 21 p 02 p 12 p 22 ,采样之后乘 [ v 0 v 1 v 2 ] \begin{bmatrix}v_0\\v_1\\v_2\end{bmatrix} v 0 v 1 v 2 ,得到输出结果的最末项 o 3 o_3 o 3 ,追加到 O O O 后面。
在第 t t t 步计算中,前面已经计算出的 { v 0 , … , v t − 1 } \{v_0, \dots, v_{t-1}\} { v 0 , … , v t − 1 } 、{ k 0 , … , k t − 1 } \{k_0, \dots, k_{t-1}\} { k 0 , … , k t − 1 } ,以及 Q K T QK^T Q K T 分数矩阵左上角的 ( t − 1 ) × ( t − 1 ) (t-1) \times (t-1) ( t − 1 ) × ( t − 1 ) 方阵,都不需要重复计算。不使用 Cache 时,每一次预测都把提示词与前缀输出序列全部送入 Decoder 做 forward,对于长度为 N N N 的序列,总体计算开销会显著增加。
引入 KV Cache 后:
Step1: 仅输入 O = [ o 0 ] O=\begin{bmatrix}o_0\end{bmatrix} O = [ o 0 ] 处理产生 q 0 , k 0 , v 0 q_0, k_0, v_0 q 0 , k 0 , v 0 。将 k 0 , v 0 k_0, v_0 k 0 , v 0 写入缓存,此时K ⊤ = [ k 0 ] , V = [ v 0 ] K^\top=\begin{bmatrix}k_0\end{bmatrix}, V=\begin{bmatrix}v_0\end{bmatrix} K ⊤ = [ k 0 ] , V = [ v 0 ] 。
计算 Q K ⊤ = [ q 0 ] [ k 0 ] = [ s 00 ] QK^\top=\begin{bmatrix}q_0\end{bmatrix}\begin{bmatrix}k_0\end{bmatrix}=\begin{bmatrix}s_{00}\end{bmatrix} Q K ⊤ = [ q 0 ] [ k 0 ] = [ s 00 ] ,做 Softmax 得到 [ p 00 ] \begin{bmatrix}p_{00}\end{bmatrix} [ p 00 ] ,采样之后乘 [ v 0 ] \begin{bmatrix}v_0\end{bmatrix} [ v 0 ] ,得到输出 o 1 o_1 o 1 ,追加到 O O O 后面;
Step2: 仅输入最新的 o 1 o_1 o 1 处理产生 q 1 , k 1 , v 1 q_1, k_1, v_1 q 1 , k 1 , v 1 。追加更新缓存,此时K ⊤ = [ k 0 k 1 ] , V = [ v 0 v 1 ] K^\top=\begin{bmatrix}k_0&k_1\end{bmatrix}, V=\begin{bmatrix}v_0\\v_1\end{bmatrix} K ⊤ = [ k 0 k 1 ] , V = [ v 0 v 1 ] 。
计算 Q K ⊤ = [ q 1 ] [ k 0 k 1 ] = [ s 10 s 11 ] QK^\top=\begin{bmatrix}q_1\end{bmatrix}\begin{bmatrix}k_0&k_1\end{bmatrix}=\begin{bmatrix}s_{10}&s_{11}\end{bmatrix} Q K ⊤ = [ q 1 ] [ k 0 k 1 ] = [ s 10 s 11 ] (注:这刚好是分数矩阵第二行的有效下三角部分,天然避开了原本需要 Mask 的 s 01 s_{01} s 01 ),做 Softmax 得到 [ p 10 p 11 ] \begin{bmatrix}p_{10}&p_{11}\end{bmatrix} [ p 10 p 11 ] ,采样之后乘 [ v 0 v 1 ] \begin{bmatrix}v_0\\v_1\end{bmatrix} [ v 0 v 1 ] ,直接得到输出 o 2 o_2 o 2 ,追加到 O O O 后面;
Step3: 仅输入最新的 o 2 o_2 o 2 处理产生 q 2 , k 2 , v 2 q_2, k_2, v_2 q 2 , k 2 , v 2 。追加更新缓存,此时K ⊤ = [ k 0 k 1 k 2 ] , V = [ v 0 v 1 v 2 ] K^\top=\begin{bmatrix}k_0&k_1&k_2\end{bmatrix}, V=\begin{bmatrix}v_0\\v_1\\v_2\end{bmatrix} K ⊤ = [ k 0 k 1 k 2 ] , V = v 0 v 1 v 2 。
计算 Q K ⊤ = [ q 2 ] [ k 0 k 1 k 2 ] = [ s 20 s 21 s 22 ] 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} Q K ⊤ = [ q 2 ] [ k 0 k 1 k 2 ] = [ s 20 s 21 s 22 ] (注:这刚好是分数矩阵第三行的有效下三角部分 ),做 Softmax 得到 [ p 20 p 21 p 22 ] \begin{bmatrix}p_{20}&p_{21}&p_{22}\end{bmatrix} [ p 20 p 21 p 22 ] ,采样之后乘 [ v 0 v 1 v 2 ] \begin{bmatrix}v_0\\v_1\\v_2\end{bmatrix} v 0 v 1 v 2 ,直接得到输出 o 3 o_3 o 3 ,追加到 O O O 后面。
这使得 Decode 阶段对于长度为 N N N 的序列,单步计算复杂度从的 O ( N 2 ) O(N^2) O ( N 2 ) (计算 N × N N \times N N × N 矩阵)降至线性的 O ( N ) O(N) O ( N ) (计算大小为 N N N 的一行),是典型的用空间换时间。
Prompt Prefill
Transformer 在实际工作时还会有提示词输入,这种情况下怎么将提示词的信息呈现给模型呢?回到之前的 Transformer 架构,左侧的 Decoder 会有一个 Input 入口,输入的内容会被传递给 Decoder 处理生成对应的 Key 和 Value。
那么在有提示词的情况下如何让 Decoder 访问提示词的信息呢?直接将提示词塞给 Eecoder 处理后传给 Decoder 就可以了。生成的 Key 和 Value 预填充进 KV Cache 中,后面直接参与 Attention 计算,就可以让每一次输入的 o i o_i o i 都可以看到提示词的信息了。
仍然走一遍计算流程,在之前的基础上,我们加上长度同样为 3 的提示词序列:
Step 1:预填充阶段,完整输入 Prompt 序列 X = [ x 0 x 1 x 2 ] X=\begin{bmatrix}x_0\\x_1\\x_2\end{bmatrix} X = x 0 x 1 x 2 ,一次性并行处理产生完整的查询、键、值矩阵:
Q = [ q 0 q 1 q 2 ] , K ⊤ = [ k 0 k 1 k 2 ] , V = [ v 0 v 1 v 2 ] 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} Q = q 0 q 1 q 2 , K ⊤ = [ k 0 k 1 k 2 ] , V = v 0 v 1 v 2
建立初始缓存,将全量的 K K K 和 V V V 写入缓存。
K c a c h e ⊤ = [ k 0 k 1 k 2 ] , V c a c h e = [ v 0 v 1 v 2 ] 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} K c a c h e ⊤ = [ k 0 k 1 k 2 ] , V c a c h e = v 0 v 1 v 2
计算 Q K c a c h e ⊤ = [ q 0 q 1 q 2 ] [ k 0 k 1 k 2 ] = [ s 00 − ∞ − ∞ s 10 s 11 − ∞ s 20 s 21 s 22 ] 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} Q K c a c h e ⊤ = q 0 q 1 q 2 [ k 0 k 1 k 2 ] = s 00 s 10 s 20 − ∞ s 11 s 21 − ∞ − ∞ s 22
做 Softmax 得到 [ p 00 0 0 p 10 p 11 0 p 20 p 21 p 22 ] \begin{bmatrix}p_{00}&0&0\\p_{10}&p_{11}&0\\p_{20}&p_{21}&p_{22}\end{bmatrix} p 00 p 10 p 20 0 p 11 p 21 0 0 p 22 ,乘 V c a c h e [ v 0 v 1 v 2 ] V_{cache}\begin{bmatrix}v_0\\v_1\\v_2\end{bmatrix} V c a c h e v 0 v 1 v 2 ,得到隐藏状态 H = [ h 0 h 1 h 2 ] H=\begin{bmatrix}h_0\\h_1\\h_2\end{bmatrix} H = h 0 h 1 h 2 。
提取最后一行包含全量上下文的 h 2 h_2 h 2 ,送入输出层,得到生成的第 1 个回答 Token o 3 o_3 o 3 ,追加到序列末尾。
Step 2:完成 Prefill 后,切换到 Decoder 递推模式。此时不需要再重新计算 x 0 , x 1 , x 2 x_0, x_1, x_2 x 0 , x 1 , x 2 对应的历史 Key/Value,只需要提取缓存;由于未来的词还未生成,当前步对历史 token 的注意力不再需要显式的完整 Mask 矩阵。
仅输入最新的 o 3 o_3 o 3 处理产生 q 3 , k 3 , v 3 q_3, k_3, v_3 q 3 , k 3 , v 3 。追加缓存,此时:
K c a c h e ⊤ = [ k 0 k 1 k 2 k 3 ] , V c a c h e = [ v 0 v 1 v 2 v 3 ] 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} K c a c h e ⊤ = [ k 0 k 1 k 2 k 3 ] , V c a c h e = v 0 v 1 v 2 v 3
计算 Q K c a c h e ⊤ = [ q 3 ] [ k 0 k 1 k 2 k 3 ] = [ s 30 s 31 s 32 s 33 ] 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} Q K c a c h e ⊤ = [ q 3 ] [ k 0 k 1 k 2 k 3 ] = [ s 30 s 31 s 32 s 33 ]
做 Softmax 得到 [ p 30 p 31 p 32 p 33 ] \begin{bmatrix}p_{30}&p_{31}&p_{32}&p_{33}\end{bmatrix} [ p 30 p 31 p 32 p 33 ] ,乘 V c a c h e [ v 0 v 1 v 2 v 3 ] V_{cache}\begin{bmatrix}v_0\\v_1\\v_2\\v_3\end{bmatrix} V c a c h e v 0 v 1 v 2 v 3 ,直接得到隐藏状态 h 3 h_3 h 3 ,送入输出层,得到生成的第 2 个回答 Token:o 4 o_4 o 4 ,追加到序列末尾。
Step 3:仅输入最新的 o 4 o_4 o 4 处理产生 q 4 , k 4 , v 4 q_4, k_4, v_4 q 4 , k 4 , v 4 。追加缓存,此时:
K c a c h e ⊤ = [ k 0 k 1 k 2 k 3 k 4 ] , V c a c h e = [ v 0 v 1 v 2 v 3 v 4 ] 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} K c a c h e ⊤ = [ k 0 k 1 k 2 k 3 k 4 ] , V c a c h e = v 0 v 1 v 2 v 3 v 4
计算 Q K c a c h e ⊤ = [ q 4 ] [ k 0 k 1 k 2 k 3 k 4 ] = [ s 40 s 41 s 42 s 43 s 44 ] 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} Q K c a c h e ⊤ = [ q 4 ] [ k 0 k 1 k 2 k 3 k 4 ] = [ s 40 s 41 s 42 s 43 s 44 ]
做 Softmax 得到 [ p 40 p 41 p 42 p 43 p 44 ] \begin{bmatrix}p_{40}&p_{41}&p_{42}&p_{43}&p_{44}\end{bmatrix} [ p 40 p 41 p 42 p 43 p 44 ] ,乘 V c a c h e [ v 0 v 1 v 2 v 3 v 4 ] V_{cache}\begin{bmatrix}v_0\\v_1\\v_2\\v_3\\v_4\end{bmatrix} V c a c h e v 0 v 1 v 2 v 3 v 4 ,直接得到隐藏状态 h 4 h_4 h 4 ,送入输出层,得到生成的第 3 个回答 Token:o 5 o_5 o 5 ,追加到序列末尾。
Prefill 过程偏计算密集 ,其一次性计算大量输入 token 的矩阵乘,尤其在使用 Full Attention 的时候有 O ( N 2 ) O(N^2) O ( N 2 ) 复杂度;Decode 过程则偏访存密集 ,每一步都需要读取模型权重和更新 KV Cache。
Linear Attention
当使用 KV Cache 时,缓存会随着序列输出的不断追加线形增长。为了避免重复计算,所有的历史 K V KV K V 都会被完整保存在显存中。在有限的显存下,不可能让 Cache 无限增长,因此要找到一种方法对历史进行压缩。
标准的 Full Attention 中,完整的输出公式是这样的:
o i = ∑ j = 1 N exp ( q i k j ⊤ ) ∑ m exp ( q i k m ⊤ ) v j o_i = \sum_{j=1}^N \frac{\exp(q_i k_j^\top)}{\sum_m \exp(q_i k_m^\top)} v_j
o i = j = 1 ∑ N ∑ m exp ( q i k m ⊤ ) exp ( q i k j ⊤ ) v j
由于 Softmax 使用指数函数进行计算,q i q_i q i 和 k j k_j k j 被绑定在指数函数内部,因此必须先算出整个 N × N N \times N N × N 的注意力矩阵,才能去乘 V V V 。
Linear Attention 的核心思想是:用一个非负的核函数 ϕ ( ⋅ ) \phi(\cdot) ϕ ( ⋅ ) 来近似或替代指数函数,消除 q q q 和 k k k 的非线性耦合,变成一个可以提取出来的乘积形式。现在使用 Linear Attention 改写输出公式,将分子替换为 Sim ( q i , k j ) = ϕ ( q i ) ϕ ( k j ) ⊤ \text{Sim}(q_i, k_j) = \phi(q_i) \phi(k_j)^\top Sim ( q i , k j ) = ϕ ( q i ) ϕ ( k j ) ⊤ ,那么没有规范化的 logits 的计算公式变成:
o t = ∑ j = 1 t ( ϕ ( q t ) ϕ ( k j ) ⊤ ) v j ⊤ o_t = \sum_{j=1}^t \left( \phi(q_t) \phi(k_j)^\top \right) v_j^\top
o t = j = 1 ∑ t ( ϕ ( q t ) ϕ ( k j ) ⊤ ) v j ⊤
根据结合律改变括号的位置:
o t = ϕ ( q t ) ( ∑ j = 1 t ϕ ( k j ) ⊤ v j ⊤ ) o_t = \phi(q_t) \left( \sum_{j=1}^t \phi(k_j)^\top v_j^\top \right)
o t = ϕ ( q t ) ( j = 1 ∑ t ϕ ( k j ) ⊤ v j ⊤ )
令中间那串累加项为状态矩阵 S t S_t S t :
S t = ∑ j = 1 t ϕ ( k j ) ⊤ v j ⊤ = S t − 1 + ϕ ( k t ) ⊤ v t ⊤ S_t = \sum_{j=1}^t \phi(k_j)^\top v_j^\top = S_{t-1} + \phi(k_t)^\top v_t^\top
S t = j = 1 ∑ t ϕ ( k j ) ⊤ v j ⊤ = S t − 1 + ϕ ( k t ) ⊤ v t ⊤
核函数可以设置为 R e L U / E L U ReLU/ELU R e LU / E LU 。进入注意力计算之前,就将核函数执行,即 ϕ ( k t ) → k t \phi(k_t) \to k_t ϕ ( k t ) → k t ,注意力计算的核心式简化为:
S t = S t − 1 + k t v t ⊤ S_t = S_{t-1} + k_t v_t^\top
S t = S t − 1 + k t v t ⊤
现在新 token 来时只需要更新固定状态 S t S_t S t ,再用 q t q_t q t 与之相乘进行读取即可。缓存变成了一个固定 d × d d \times d d × d 大小的矩阵 S i S_i S i ,Linear Attention 机制类似 RNN:每来一个 Token,就用当前的 k t , v t k_t, v_t k t , v t 更新 S t S_t S t ,然后用 q t q_t q t 读取 S t S_t S t 。但是,原本随序列长度线性增长的信息被压缩到固定大小矩阵后,会形成有损表示 ,因此在面对更长文本时可能像 RNN 一样丢失或覆盖历史信息。相较于 KV Cache “用空间换效率”,Linear Attention 更接近“用精度换空间”。
Gated Delta Net
在处理信息丢失和覆盖的的问题时,LSTM 在 RNN 的基础上引入了门控机制,用输入门和遗忘门来选择控制 Memory Cell 更新时的旧记忆保留和新记忆写入:
c t = f t ⊙ c t − 1 + i t ⊙ g t c_t=f_t⊙c_{t−1}+i_t⊙g_t
c t = f t ⊙ c t − 1 + i t ⊙ g t
通过调整这些门的值,LSTM 能够灵活地保留或更新信息,从而更好地捕捉长期依赖关系,缓解固定维度向量记忆的丢失。
类似的,Gated Data Net 也引入了门控机制。对于上一步留下的状态矩阵 S t − 1 S_{t-1} S t − 1 ,为其乘上一个遗忘门 α t \alpha_t α t ,这个值由输入动态生成,介于 [ 0 , 1 ] [0, 1] [ 0 , 1 ] 之间,用来决定以何种程度保留上一步的记忆。这个值趋于 1 时,记忆被完整保存,反之则抹除。
对于增量项,则引入了 Delta Rule 。首先定义预测误差 e t = v t − S t − 1 T k t e_t=v_t-S_{t-1}^Tk_t e t = v t − S t − 1 T k t 。它利用现有的记忆和当前的键 k t k_t k t 尝试去预测对应的 value,预测误差就是预测结果与实际值之间的误差。
为什么用 S t − 1 T k t S_{t-1}^Tk_t S t − 1 T k t 来做 value 的预测?先回到 Linear Attention 的状态更新公式 S t = ∑ j = 1 t k j v j ⊤ S_t = \sum_{j=1}^t k_j v_j^\top S t = ∑ j = 1 t k j v j ⊤ ,此时的 S t S_t S t 是一个 d × d d \times d d × d 的矩阵。我们现在站在第 t t t 步,用当前的键向量 k t k_t k t 乘上 S t ⊤ S_t^\top S t ⊤ ,得到:
S t ⊤ k t = ( ∑ j = 1 t v j k j ⊤ ) k t = ∑ j = 1 t v j ( k j ⊤ k t ) S t ⊤ k t = v t ( k t ⊤ k t ) + ∑ j = 1 t − 1 v j ( k j ⊤ k t ) 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)
S t ⊤ k t = ( j = 1 ∑ t v j k j ⊤ ) k t = j = 1 ∑ t v j ( k j ⊤ k t ) S t ⊤ k t = v t ( k t ⊤ k t ) + j = 1 ∑ t − 1 v j ( k j ⊤ k t )
而在高维向量空间中,两个随机生成的向量很大概率是接近正交的,又考虑模型中一般都会做向量归一化,就有 k t ⊤ k t ≈ 1 k_t^\top k_t \approx 1 k t ⊤ k t ≈ 1 ,且当当前的新键与之前的某个旧键没有关联时, k j ⊤ k t ≈ 0 ( j ≠ t ) k_j^\top k_t \approx 0(j \ne t) k j ⊤ k t ≈ 0 ( j = t ) 。把这两个式子代进之前的公式,就能得到:
S t ⊤ k t ≈ v t ⋅ 1 + ∑ v j ⋅ 0 = v t S_t^\top k_t \approx v_t \cdot 1 + \sum v_j \cdot 0 = v_t
S t ⊤ k t ≈ v t ⋅ 1 + ∑ v j ⋅ 0 = v t
证明了用 S t − 1 k t S_{t-1}k_t S t − 1 k t 来预测 value 的合理性, 再回到预测误差。得到预测值后,我们得到一个用于修正 value 预测误差的量 k t e t T = k t ( v t − S t − 1 ⊤ k t ) ⊤ k_te_t^T=k_t (v_t - S_{t-1}^\top k_t)^\top k t e t T = k t ( v t − S t − 1 ⊤ k t ) ⊤ 。加上这个补偿的 state 可以准确地预测出当前时间步 key 对应的 value。考虑训练时,定义一个目标是让当前的记忆能精准读出当前 value 的损失函数:
L t ( S ) = 1 2 ∥ S ⊤ k t − v t ∥ 2 2 \mathcal{L}_t(S) = \frac{1}{2} \| S^\top k_t - v_t \|_2^2
L t ( S ) = 2 1 ∥ S ⊤ k t − v t ∥ 2 2
对矩阵 S S S 求偏导:
∇ S L t ( S ) = k t ( S ⊤ k t − v t ) ⊤ = − k t ( v t − S ⊤ k t ) ⊤ \nabla_S \mathcal{L}_t(S) = k_t (S^\top k_t - v_t)^\top = - k_t (v_t - S^\top k_t)^\top
∇ S L t ( S ) = k t ( S ⊤ k t − v t ) ⊤ = − k t ( v t − S ⊤ k t ) ⊤
给定一个动态更新的学习率 β t \beta_t β t 按照梯度下降的方式更新 S S S ,有:
S t = S t − 1 − β t ∇ S L t ( S t − 1 ) = S t − 1 + β t k t ( v t − S t − 1 ⊤ k t ) ⊤ 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
S t = S t − 1 − β t ∇ S L t ( S t − 1 ) = S t − 1 + β t k t ( v t − S t − 1 ⊤ k t ) ⊤
β t k t ( v t − S t − 1 ⊤ k t ) ⊤ \beta_t k_t (v_t - S_{t-1}^\top k_t)^\top β t k t ( v t − S t − 1 ⊤ k t ) ⊤ 就是 Delta Rule 的增量项,其中 β t \beta_t β t 就是控制记忆写入的写入门 。问题又来了,这个训练时的学习率和控制写入的门之间,是否存在等价关系呢?
将 β t \beta_t β t 当成一个介于 [ 0 , 1 ] [0, 1] [ 0 , 1 ] 之间的动态学习率,这个学习率用输入的 token 来预测:
β t = Sigmoid ( W β ⋅ x t ) \beta_t = \text{Sigmoid}(W_\beta \cdot x_t)
β t = Sigmoid ( W β ⋅ x t )
训练时梯度下降会推着 W β W_\beta W β 向着使 state 能更准确预测 value 的方向移动,那么最终训练出的动态“学习率”,在推理时也会倾向于去补偿预测误差——预测误差大说明 state 实际变化梯度大,这个学习率会增大(逼近1)来加大写入力度;预测误差小说明 state 实际变化梯度小,这个学习率会降低(逼近0)来降低写入力度。这个网络在推理时固定下来,用来调整写入力度。
至此 Gated Delta Net 的状态更新函数就确定下来了:
S t = α t S t − 1 + β t k t ( v t − S t − 1 ⊤ k t ) ⊤ S_t = \alpha_t S_{t-1} + \beta_t k_t (v_t - S_{t-1}^\top k_t)^\top
S t = α t S t − 1 + β t k t ( v t − S t − 1 ⊤ k t ) ⊤
这个函数能够动态调整每个新 token 运算后“记忆”的遗忘程度和写入程度,缓解了固定大小状态矩阵丢失或覆盖历史信息的问题。