先不记 $Q\odot B$ 和 $K\oslash B$。我们从一次普通读取开始,看遗忘门怎样改变它,最后再把同一个计算整理成矩阵乘法。这两个式子是把已有乘除重新分组得到的,不是另外添加的两条模型规则。
第一步:没有门时,一笔写入怎样参与输出
先把 key、query、value 都简化为一个数字,状态 s 也只有一个数字。从 $s_0=0$ 开始,普通线性注意力每一步做 $s_t=s_{t-1}+v_tk_t$,再用 $o_t=s_tq_t$ 读取。
例如,第 1 个位置写入了 $v_1k_1$。到第 3 个位置用 $q_3$ 读取时,这笔写入对输出的贡献是:
$$
(v_1k_1)q_3=v_1\,(q_3k_1).
$$
括号里的 $q_3k_1$ 是一个匹配分数,再用这个分数乘 $v_1$。下面只追踪这笔写入,暂时不管其他位置写了什么。
第二步:加入门后,这笔写入会经历哪些衰减
GLA 每一步先给旧状态乘一个保留率 $\alpha_t$,再加入新写入:
$$
s_t=\alpha_ts_{t-1}+v_tk_t.
$$
第 1 步新写入 $v_1k_1$,此时它还不乘 $\alpha_1$,因为 $\alpha_1$ 作用于写入前的旧状态。第 2 步它已经属于旧状态,所以乘 $\alpha_2$;第 3 步再乘 $\alpha_3$。这笔内容的变化是:
$$
v_1k_1
\ \longrightarrow\ \alpha_2v_1k_1
\ \longrightarrow\ \alpha_3\alpha_2v_1k_1.
$$
到第 3 步用 query 读取,它对输出的贡献就是:
$$
q_3\alpha_3\alpha_2v_1k_1
=v_1\,\underbrace{(q_3k_1\alpha_2\alpha_3)}_{\text{加入遗忘后的匹配分数}}.
$$
所以,门改变的是匹配分数:原来的 $q_3k_1$,现在还要乘从写入之后到读取时刻之间的保留率 $\alpha_2\alpha_3$。
第三步:用累计乘积,把“中间这一段门”写成一个比值
如果每次读取都重新乘一遍中间所有门,会很不方便。先为每个位置记下“从起点到这里的门全部乘起来”的结果:
$$
b_1=\alpha_1,\qquad
b_2=\alpha_1\alpha_2,\qquad
b_3=\alpha_1\alpha_2\alpha_3.
$$
这就是门的前缀积。“前缀”只是从起点到当前位置的这一段;b 是累计结果,α 才是每一步生成的门。要取第 1 步写入到第 3 步读取之间的门,只需:
$$
\frac{b_3}{b_1}
=\frac{\alpha_1\alpha_2\alpha_3}{\alpha_1}
=\alpha_2\alpha_3.
$$
分母消掉了写入前已经发生的门。比如 $\alpha_1=0.5,\alpha_2=0.2,\alpha_3=0.5$,则 $b_1=0.5,b_3=0.05$,比值是 $0.05/0.5=0.1$,恰好等于后两步的 $0.2\times0.5$。
因此,刚才的匹配分数可以改写为:
$$
q_3k_1\alpha_2\alpha_3
=q_3k_1\frac{b_3}{b_1}.
$$
第四步:这一步重新分组,才出现“query 乘、key 除”
目前式子里的 q、k、b 都只是数字,可以交换相乘的顺序。因此:
$$
\underbrace{q_3k_1\frac{b_3}{b_1}}_{\text{原来的分数乘路径保留率}}
=\underbrace{(q_3b_3)}_{\text{只与读取位置 3 有关}}
\underbrace{\left(\frac{k_1}{b_1}\right)}_{\text{只与写入位置 1 有关}}.
$$
query 乘 $b_3$,是因为路径比值的分子属于读取位置;key 除 $b_1$,是因为分母属于写入位置。 我们只是把同一组乘除分到了两边,结果没有改变。
继续用数字验证,取 $q_3=2,k_1=4$。原来的算法是 $2\times4\times0.2\times0.5=0.8$。重新分组后:
$$
(2\times0.05)\times(4/0.5)=0.1\times8=0.8.
$$
除法使 key 这边的数变大,但它会与 query 那边的小数一起使用,两者合起来仍是正确的遗忘量。不能只把 $k_1/b_1$ 单独看成“把记忆增强了”。
这个分组还有一个计算上的好处:$k_1/b_1$ 对后面所有 query 都相同,可以提前算好;$q_3b_3$ 对它要读取的所有历史 key 都相同,也只需算一次。
第五步:多个特征,就在每个特征上做同样的分组
真正的 q 和 k 各有 $d_k$ 个数字,匹配分数是对应特征相乘后求和。GLA 每个 key 特征也有自己的门,所以刚才的推导要逐特征进行,再把结果相加。
先只看两个特征。用 $q_{r,a}$ 表示第 r 个位置、第 a 个特征的 query 数值;$b_{r,a}$ 是这个特征累计到 r 的门积。第 r 个 query 读取第 t 个 key,分数为:
$$
q_{r,1}k_{t,1}\frac{b_{r,1}}{b_{t,1}}
+q_{r,2}k_{t,2}\frac{b_{r,2}}{b_{t,2}}.
$$
第一项是特征 1 的贡献,第二项是特征 2 的贡献。对每一项重复刚才的重新分组:
$$
(q_{r,1}b_{r,1})\left(\frac{k_{t,1}}{b_{t,1}}\right)
+(q_{r,2}b_{r,2})\left(\frac{k_{t,2}}{b_{t,2}}\right).
$$
现在,这恰好是下面两列数字做内积:
$$
\widehat q_r=
\begin{bmatrix}q_{r,1}b_{r,1}\\q_{r,2}b_{r,2}\end{bmatrix},
\qquad
\widehat k_t=
\begin{bmatrix}k_{t,1}/b_{t,1}\\k_{t,2}/b_{t,2}\end{bmatrix}.
$$
内积就是对应位置相乘,再把乘积相加,得到一个数。用简写记号,逐特征相乘是 ⊙,逐特征相除是 ⊘,所以上面两列数字分别就是:
$$
\widehat q_r=q_r\odot b_r,\qquad
\widehat k_t=k_t\oslash b_t.
$$
这时 $b_r=\alpha_1\odot\cdots\odot\alpha_r$ 是一列累计保留率,不能把不同特征的门混乘成一个数。帽子 $\widehat{\phantom q}$ 只是给重新缩放后的向量起个名字;它们不是额外学习出来的新 query/key。
第六步:把所有位置按行排起来,才得到 Q⊙B 和 K⊘B
小写 $b_t$ 记录一个位置的累计保留率;大写 B 把所有位置的 $b_t$ 横过来,逐行排成一张表。B 的行是 token 位置,列是 key 特征。Q、K 的行列含义与它完全对应。
例如三个 token、两个特征,门及其累计表为:
$$
\text{每一步的门}=\begin{bmatrix}0.5&0.8\\0.2&0.5\\0.5&0.25\end{bmatrix},\qquad
B=\begin{bmatrix}0.5&0.8\\0.1&0.4\\0.05&0.1\end{bmatrix}.
$$
B 第一列依次是 $0.5$、$0.5\times0.2$、$0.5\times0.2\times0.5$;第二列也独立这样计算。这张表仅由门算出,不需要学习额外参数。
刚才已经知道,每个 query 要乘自己位置的累计值。对全部 query 一起做,就是 Q 与 B 对应位置相乘:
$$
\widehat Q=Q\odot B=
\begin{bmatrix}
0.5Q_{11}&0.8Q_{12}\\
0.1Q_{21}&0.4Q_{22}\\
0.05Q_{31}&0.1Q_{32}
\end{bmatrix}.
$$
每个 key 则除以自己位置的累计值。对全部 key 一起做,就是 K 与 B 对应位置相除:
$$
\widehat K=K\oslash B=
\begin{bmatrix}
K_{11}/0.5&K_{12}/0.8\\
K_{21}/0.1&K_{22}/0.4\\
K_{31}/0.05&K_{32}/0.1
\end{bmatrix}.
$$
比如第三个 query 是 $(2,3)$,第一个 key 是 $(4,5)$。重新缩放后,它们分别是 $(0.1,0.3)$ 和 $(8,6.25)$,内积为 $0.1\times8+0.3\times6.25=2.675$。直接按两条特征的遗忘路径计算,也得到 $2\times4\times0.1+3\times5\times0.125=2.675$。
第七步:把所有位置对的分数一次算出来,再汇总 value
现在每个位置的 query 和 key 都准备好了。让 $\widehat Q$ 的每一行与 $\widehat K$ 的每一行做内积,就得到所有位置对的分数。矩阵乘法恰好做这件事:
$$
P=\widehat Q\widehat K^\top=(Q\odot B)(K\oslash B)^\top.
$$
这里转置 $\top$ 只是把 K 的行变成列,让“左边一行乘右边一列”的矩阵乘法规则,变成我们需要的“一个 query 与一个 key 做内积”。$P_{rt}$ 是第 r 个 query 读取第 t 个 value 的权重;P 不是归一化的概率表。
矩阵乘法还会算到未来位置。乘一个下三角 0/1 表 M,把它们清零。三个 token 时:
$$
P\odot M=
\begin{bmatrix}
P_{11}&0&0\\
P_{21}&P_{22}&0\\
P_{31}&P_{32}&P_{33}
\end{bmatrix}.
$$
第一行只读 $v_1$,第二行读 $v_1,v_2$,第三行读 $v_1,v_2,v_3$。例如第三个输出是 $o_3=P_{31}v_1+P_{32}v_2+P_{33}v_3$。把每个位置的输出也按行排好,就是本页 PPT 的完整公式:
$$
\boxed{O=\big((Q\odot B)(K\oslash B)^\top\odot M\big)V.}
$$
这条式子里,乘 B 与除 B 合起来负责遗忘,M 负责不看未来,最后乘 V 负责把内容取回来。对角线处 $b_r/b_r=1$,所以本步刚写入的内容也会参与输出,且不会凭空多乘一个遗忘门。
并行具体发生在哪里
上面的公式对应整段已知输入、零初始状态,B 从序列开头累计,有 L 行。各位置的 Q/K/V 与门生成后,先计算 B,再计算缩放后的 Q/K,最后用矩阵乘法同时处理许多位置对。计算 $o_3$ 不再要求先构造完整 $S_1,S_2,S_3$:中间状态对这笔写入的作用,已经被门积比值明确写在分数里。
B 的前缀积仍有计算成本,可以使用前缀扫描组织;所有位置对也仍有平方级数量。下一页开始才将这套公式用于较短的块:把累计起点改为本块入口,B 只有 C 行,再单独处理入口历史。逐 token 生成时未来输入还不存在,则仍使用递推。
接到分块:如果入口还有旧历史,输出需要加什么
以上为看清主线,先令初始状态为零。若本块入口已有状态 S,入口历史会经历从本块第一步到读取位置 r 的所有门,即 $b_r$。沿用第 83 页的“按列缩放状态”,这部分读取为:
$$
S\operatorname{Diag}(b_r)q_r=S(q_r\odot b_r).
$$
因为对角矩阵乘 q,只是把 q 每个特征乘相应保留率。因此所有 query 对入口历史的读取为 $(Q\odot B)S^\top$,整块输出变成:
$$
O=(Q\odot B)S^\top+
\big((Q\odot B)(K\oslash B)^\top\odot M\big)V.
$$
这与第 87 页的式子一致:入口历史和块内写入分别读取,最后相加。这里 B 是本块累计表,不再是整段累计表。
对齐 PPT 的 log 形式:同一个比值的另一种写法
用 i 表示 query 位置、j 表示 key 位置、a 表示特征,把第七步的矩阵乘法展开:
$$
\begin{aligned}
P_{ij}
&=\sum_{a=1}^{d_k}(Q_{ia}B_{ia})\frac{K_{ja}}{B_{ja}}\\
&=\sum_{a=1}^{d_k}Q_{ia}K_{ja}\frac{B_{ia}}{B_{ja}}\\
&=\sum_{a=1}^{d_k}Q_{ia}K_{ja}\exp(\log B_{ia}-\log B_{ja}).
\end{aligned}
$$
求和号表示把所有特征的贡献相加;最后一步用 $x/y=\exp(\log x-\log y)$。PPT 用 k 作特征下标,含义与这里的 a 一样。直接相除是数学推导;长序列中的累计值可能非常小,实际内核还需对数计算、二级分块等数值处理,详见第 87 页的实现补充。