RNN的复兴04:线性注意力并行计算DPLR

作者: 引线小白-本文永久链接:httpss://www.limoncc.com/post/9e070b6858f0e490/
知识共享许可协议: 本博客采用署名-非商业-禁止演绎4.0国际许可证

摘要: 本文意在理清注意力机制优化之DPLR并行计算。若有错误,请大家指正。
关键词: DPLR,线性注意力,KDA,记忆,RWKV

一、导言

1.1、引言

无论是KDA论文,还是RWKV论文在论述并行算法时都是极度的不友好,符号过于丑陋。描述也是一略而过,看的是云里雾里。本文使用清晰的符号和下角标,尽可能优雅描述并行计算的每个部分,带你快速看懂像现代RNN的并行算法。

1.2、回顾DeltaNet的IPLR

经典的DeltaNet的状态更新矩阵涉及的是单位矩阵+秩一矩阵(IPLR),有形如下面的结构 $\displaystyle \bm{H}_{\tau}=\bm{E}_{\tau}-\bm{x}_{\tau}\bm{y}_{\tau}^\T$,这样的矩阵有 WY表示方法。

$$\begin{align}
\bm{P}_t=\prod_{\tau=1}^t\bm{H}_{\tau}=\bm{E}-\bm{W}^\T\bm{Y}
\end{align}$$

可以将连乘运算转换为连加运算,其中 $\bm{W}^\T=[\bm{w}_1,\cdots,\bm{w}_t]$, $\bm{Y}^\T=[\bm{y}_1,\cdots,\bm{y}_t]$。 其中 $\bm{w}_\tau \mathbf{y}_\tau^\T$ 是节点 $\tau$ 对全局状态的影响,$\sum$ 表示所有节点影响的叠加,负号则是历史路径的衰减效应。

证明:
$$\begin{align}
\bm{P}_t
&= \bigg[\bm{E}-\sum_{\tau}^{t-1} \bm{w}_\tau \bm{y}_\tau ^\T\bigg] \bigg[\bm{E}-\bm{x}_{t}\bm{y}_{t}^\T \bigg]\\
&=\bm{E}-\sum_{\tau}^{t-1} \bm{w}_\tau \bm{y}_\tau ^\T-\bm{x}_{t}\bm{y}_{t}^\T +\sum_{\tau}^{t-1} \bm{w}_\tau \bm{y}_\tau ^\T\bm{x}_{t}\bm{y}_{t}^\T\\
&=\bm{E}-\sum_{\tau}^{t-1} \bm{w}_\tau \bm{y}_\tau ^\T - \underbrace{\bigg[\bm{x}_{t}-\sum_{\tau}^{t-1} \bm{w}_\tau \bm{y}_\tau ^\T\bm{x}_{t}\bigg]}_{\bm{w}_t}\bm{y}_{t}^\T\\
&=\bm{E}-\sum_{\tau}^t \bm{w}_\tau \bm{y}_\tau ^\T\\
&=\bm{E}-\bm{W}^\T\bm{Y}
\end{align}$$

证明过程中还实现了 $\bm{w}_t$ 的构造方法:

$$\begin{align}
\bm{w}_t = \bm{x}_{t}-\sum_{\tau}^{t-1} \bm{w}_\tau \bm{y}_\tau ^\T\bm{x}_{t}=\bm{x}_{t}-\sum_{\tau}^{t-1} \bm{y}_\tau ^\T\bm{x}_{t}\bm{w}_\tau
\end{align}$$

具体分析可以参考笔者的文章 RNN的复兴03:现代RNN的并行计算IPLR

二、现代RNN的DPLR

对于 RWKV7[^3] 或 KDA 这样的现代RNN,它们记忆更新结构更多的形如 $\displaystyle \bm{H}_{\tau}=\bm{\varLambda}_{\tau}-\bm{x}_{\tau}\bm{y}_{\tau}^\T$, 对角矩阵取代了单位矩阵,再要变换为上述 WY 表示就比较麻烦了。先定义:$\displaystyle \bm{D}_t=\prod_{\tau=1}^t\bm{\varLambda}_{\tau}$,其实有

$$\begin{align}
\bm{D}_{t} = \bm{D}_{1:t}=\prod_{\tau=1}^t\bm{\varLambda}_{\tau}
=\prod_{\tau=1}^t\mathrm{diag}(\bm{a}_\tau)
=\mathrm{diag}\bigg[\prod_{\tau=1}^t\bm{a}_\tau\bigg]
\end{align}$$

这样实际上就有

$$\begin{align}
\bm{P}_t = \prod_{\tau=1}^t\bm{H}_{\tau}
=\bm{D}_t - \bm{W}_t^\T \bm{Y}_t
\end{align}$$

现在问题是 WY到底是什么?定义 $\bm{W}_t^\T=[\bm{w}_1,\cdots,\bm{w}_t]$, $\bm{Y}_t^\T=[\bm{\nu}_1,\cdots,\bm{\nu}_\tau^{(t)},\cdots,\bm{\nu}_t]$,注意 $\bm{\nu}_\tau^{(t)}$ 同时依赖时刻 $\tau$ 和时刻 $t$ 。下面展开分析:

$$\begin{align}
\bm{P}_t
&=\bm{P}_{t-1}\bigg[\bm{\varLambda}_t-\bm{x}_{t}\bm{y}_{t}^\T \bigg]
=\bigg[\bm{D}_{t-1}-\bm{W}_{t-1}^\T \bm{Y}_{t-1}\bigg] \bigg[\bm{\varLambda}_t-\bm{x}_{t}\bm{y}_{t}^\T \bigg]\\
&= \bigg[\bm{D}_{t-1}-\sum_{\tau}^{t-1} \bm{w}_\tau \bm{\nu}_\tau ^{(t-1)\T}\bigg] \bigg[\bm{\varLambda}_t-\bm{x}_{t}\bm{y}_{t}^\T \bigg]\\
&=\bm{D}_{t}
-\sum_{\tau}^{t-1} \bm{w}_\tau \bm{\nu}_\tau ^{(t-1)\T}\bm{\varLambda}_t
-\bm{D}_{t-1}\bm{x}_{t}\bm{y}_{t}^\T
+\sum_{\tau}^{t-1} \bm{w}_\tau \bm{\nu}_\tau ^{(t-1)\T}\bm{x}_{t}\bm{y}_{t}^\T\\
&=\bm{D}_{t}
-\sum_{\tau}^{t-1} \bm{w}_\tau \big(\bm{\varLambda}_t\bm{\nu}_\tau ^{(t-1)} \big)^\T
-\underbrace{\bigg[\bm{D}_{t-1}\bm{x}_{t}
-\sum_{\tau}^{t-1} \bm{w}_\tau \bm{\nu}_\tau ^{(t-1)\T}\bm{x}_{t}\bigg]}_{\bm{w}_t}\bm{y}_{t}^\T\\
&=\bm{D}_{t}
-\sum_{\tau}^{t-1} \bm{w}_\tau \big(\bm{\nu}_\tau ^{(t)} \big)^\T-\bm{w}_{t}\bm{y}_t^\T\\
&=\bm{D}_{t}
-\sum_{\tau}^{t-1} \bm{w}_\tau \big(\bm{\nu}_\tau ^{(t)} \big)^\T-\bm{w}_{t}\big(\bm{\nu}_t^{(t)}\big)^{\T}\\
&=\bm{D}_{t}
-\sum_{\tau}^{t} \bm{w}_\tau \big(\bm{\nu}_\tau ^{(t)} \big)^\T\\
&=\bm{D}_t - \bm{W}_t^\T \bm{Y}_t
\end{align}$$

其中定义 $\displaystyle \bm{w}_t \triangleq \bm{D}_{t-1}\bm{x}_{t} - \sum_{\tau}^{t-1}\big[\big(\bm{\nu}_\tau^{(t-1)}\big)^\T\bm{x}_{t}\big] \bm{w}_\tau$,同时比较易得

$$\begin{align}
\bm{\nu}_{\tau}^{(t)}
=\bm{\Lambda}_{t}\bm{\Lambda}_{t-1}\cdots\bm{\Lambda}_{\tau+1}\cdot\bm{y}_{\tau}
\Longleftarrow
\begin{cases}
\bm{\nu}_{\tau}^{(t)} &= \bm{\Lambda}_t \bm{\nu}_{\tau}^{(t-1)}\\
\bm{\nu}_{t}^{(t)} &=\bm{y}_{t}
\end{cases}
\end{align}$$

由于 $\bm{\varLambda}_{\tau}$ 是对角矩阵,那么

$$\begin{align}
\bm{\nu}_\tau^{(t)} =\bm{D}_{\tau+1:t}\cdot\bm{y}_\tau, \quad \tau < t
\end{align}$$

这样就得到了一个较为通用的 DPLR 矩阵的 WY表示方法。$\bm{y}_\tau$ 是第 $\tau$ 时刻写入的“原始基底”。$\bm{\varLambda}_{\tau+1} \cdots \bm{\varLambda}_t$ 是从写入时刻 $\tau$ 之后,一直到当前时刻 $t$ 的累积衰减因子。也就是说相对于 DeltaNet 的 IPLR 的 $\bm{Y}^\T=[\bm{y}_1,\cdots,\bm{y}_t]$,现在的 $\bm{Y}_t^\T=[\bm{D}_{2:t}\cdot\bm{y}_1,\cdots,\bm{D}_{\tau+1:t}\cdot\bm{y}_\tau,\cdots,\bm{y}_t]$, 也就是说 $\bm{y}$ 节点的信息需要随时间衰减。不过 $\bm{\nu}_1^{(t)}=\bm{D}_{2:t}\cdot\bm{y}_1=\bm{\Lambda}_{t}\bm{\Lambda}_{t-1}\cdots\bm{\Lambda}_{2}\cdot\bm{y}_{1}$, 那么 $\bm{\Lambda}_{1}$ 的作用在哪?$\bm{\Lambda}_{1}$ 更多作用在遗忘算子 $\bm{w}_t$ 上。

三、DPLR 并行计算实现

3.1、遗忘与写入共享键向量情形

有了 WY 表示,那么我们如何实现RNN 状态更新公式的并行计算?在遗忘键向量 $\bm{y}_t$ 和写入键向量
$\bm{k}_t$ 相同的情况下,对于 $\displaystyle \bm{S}_t = \bm{S}_{t-1}\big[\bm{\varLambda}_{t}-\bm{x}_{t}\bm{y}_{t}^\T\big] + \bm{v}_t\bm{k}_t^\T$ 有如下较为通用的并行更新计算公式:

$$\begin{align}
\bm{S}_t = \bm{S}_{0}\big[\bm{D}_t - \bm{W}_t^\T \bm{Y}_t\big] + \bm{U}_t^\T\bm{K}_t
\end{align}$$

其中 $\bm{W}_t^\T=[\bm{w}_1,\cdots,\bm{w}_t]$, $\bm{Y}_t^\T=[\bm{\nu}_1,\cdots,\bm{\nu}_\tau^{(t)},\cdots,\bm{\nu}_t]$, $\bm{U}_t^\T=[\bm{u}_1,\cdots,\bm{u}_t]$, $\bm{K}_t^\T=[\bm{\kappa}_1,\cdots,\bm{\kappa}_{\tau}^{(t)},\cdots,\bm{\kappa}_t]$

证明

$$\begin{align}
\bm{S}_t
&=\Big[\bm{S}_{t-2}\big[\bm{\varLambda}_{t-1}-\bm{x}_{t-1}\bm{y}_{t-1}^\T\big] + \bm{v}_{t-1}\bm{k}_{t-1}^\T\Big]
\big[\bm{\varLambda}_{t}-\bm{x}_{t}\bm{y}_{t}^\T\big]+\bm{v}_t\bm{k}_t^\T\\
&=\bm{S}_{t-2}\big[\bm{\varLambda}_{t-1}-\bm{x}_{t-1}\bm{y}_{t-1}^\T\big]\big[\bm{\varLambda}_{t}-\bm{x}_{t}\bm{y}_{t}^\T\big]
+\bm{v}_{t-1}\bm{k}_{t-1}^\T\big[\bm{\varLambda}_{t}-\bm{x}_{t}\bm{y}_{t}^\T\big]
+\bm{v}_t\bm{k}_t^\T\\
&=\bm{S}_0\underbrace{\prod_{\tau=1}^t \Big[\bm{\varLambda}_{\tau}-\bm{x}_{\tau}\bm{y}_{\tau}^\T\Big]}_{\text{累积转移矩阵}}
+\underbrace{\sum_{\tau=1}^t \Big[\bm{v}_\tau \bm{k}_\tau^\T \prod_{i=\tau+1}^t \big[\bm{\varLambda}_{i}-\bm{x}_{i}\bm{y}_{i}^\T\big]\Big]}_{\text{累积更新项}}
\end{align}$$

其中定义 $\displaystyle \bm{\varLambda}_{t+1}-\bm{x}_{t+1} \bm{y}_{t+1}^\T=\bm{E}-\bm{0}=\bm{E}$,对于左侧的累积转移矩阵,根据前面的推导易得结论。

$$\begin{align}
\bm{S}_0\prod_{\tau=1}^t \Big[\bm{\varLambda}_{\tau}-\bm{x}_{\tau}\bm{y}_{\tau}^\T\Big]
=\bm{S}_0\Big[\bm{D}_t-\sum_{\tau}^t \bm{w}_\tau \bm{y}_\tau ^\T\Big]
=\bm{S}_0\big[\bm{D}_t-\bm{W}_t^\T\bm{Y}_t\big]
\end{align}$$

现在关注右侧累积更新项 $\bm{L}(t)$, 对于衰减与更新共享键向量情形 $\bm{y}_t=\bm{k}_t$:

$$\begin{align}
\bm{L}(t)
&=\sum_{\tau=1}^t \Big[\bm{v}_\tau \bm{k}_\tau^\T \prod_{i=\tau+1}^t \big[\bm{\varLambda}_{i}-\bm{x}_{i}\bm{y}_{i}^\T\big]\Big]\\
&=\bm{L}(t-1)\big[\bm{\varLambda}_{t} - \bm{x}_t \bm{y}_t^\T\big]
+\bm{v}_t\bm{k}_t^\T\\
&=\sum_{\tau=1}^{t-1} \bm{u}_\tau \bm{\kappa}_\tau^{(t-1)\T}\big[\bm{\varLambda}_{t} - \bm{x}_t \bm{y}_t^\T\big]
+\bm{v}_t\bm{k}_t^\T\\
&=\sum_{\tau=1}^{t-1} \bm{u}_\tau \bm{\kappa}_\tau^{(t-1)\T}\bm{\varLambda}_{t}
-\sum_{\tau=1}^{t-1} \bm{u}_\tau \bm{\kappa}_\tau^{(t-1)\T}\bm{x}_t \bm{y}_t^\T
+\bm{v}_t\bm{k}_t^\T\\
&=\sum_{\tau=1}^{t-1} \bm{u}_\tau \big(\bm{\varLambda}_{t}\bm{\kappa}_\tau^{(t-1)}\big)^\T
-\sum_{\tau=1}^{t-1} \bm{u}_\tau \bm{\kappa}_\tau^{(t-1)\T}\bm{x}_t \bm{k}_t^\T
+\bm{v}_t\bm{k}_t^\T\Longleftarrow \bm{y}_t=\bm{k}_t\\
&=\sum_{\tau=1}^{t-1} \bm{u}_\tau \big(\bm{\kappa}_\tau^{(t)}\big)^\T
+\Big[\underbrace{\bm{v}_t-\sum_{\tau=1}^{t-1}\bm{u}_\tau\big[\bm{\kappa}_{\tau}^{(t-1)\T}\bm{x}_t\big]}_{\bm{u}_t}\Big]\bm{k}_t^\T\\
&=\sum_{\tau=1}^{t} \bm{u}_\tau \big(\bm{\kappa}_\tau^{(t)}\big)^\T\\
&=\bm{U}^\T\bm{K}
\end{align}$$

其中 $\displaystyle \bm{u}_t \triangleq \bm{v}_t-\sum_{\tau=1}^{t-1}\bm{u}_\tau\big[\bm{\kappa}_{\tau}^{(t-1)\T}\bm{x}_t\big]=\bm{v}_t-\sum_{\tau=1}^{t-1}\big[\bm{\kappa}_{\tau}^{(t-1)\T}\bm{x}_t\big]\bm{u}_\tau$ 同时易得

$$\begin{align}
\bm{\kappa}_{\tau}^{(t)}
=\bm{\Lambda}_{t}\bm{\Lambda}_{t-1}\cdots\bm{\Lambda}_{\tau+1}\cdot\bm{k}_{\tau}
\Longleftarrow
\begin{cases}
\bm{\kappa}_{\tau}^{(t)} &= \bm{\Lambda}_t \bm{\kappa}_{\tau}^{(t-1)}\\
\bm{\kappa}_{t}^{(t)} &=\bm{k}_{t}
\end{cases}
\end{align}$$

那么有

$$
\begin{align}
\bm{\kappa}_\tau^{(t)} =\bm{D}_{\tau+1:t}\cdot\bm{k}_\tau, \quad \tau < t
\end{align}$$

3.2、遗忘与写入的键向量独立情形

如果形如 RWKV 在遗忘和写入中使用独立的键向量,那么情况将会复杂点。为方便叙述, 这里引入一个新的量 反馈记忆:当前输入键 $\bm{x}_t$ 对记忆状态 $\bm{S}_{t-1}$ 的读取:

$$\begin{align}
\bm{\zeta}_t\triangleq \bm{S}_{t-1}\bm{x}_t
\end{align}$$

将 $\bm{\zeta}_t$ 代入原状态更新公式 $\bm{S}_t = \bm{S}_{t-1}\big[\bm{\varLambda}_t - \bm{x}_t\bm{y}_t^\T\big] + \bm{v}_t\bm{k}_t^\T$ ,并从 $\tau=1$ 到 $t$ 展开得到:

$$\begin{align}
\bm{S}_t
&=\bm{S}_{t-1}\bm{\varLambda}_t - \bm{\zeta}_t\bm{y}_t^\T + \bm{v}_t\bm{k}_t^\T\\
&=\big[\bm{S}_{t-2}\bm{\varLambda}_{t-1} - \bm{\zeta}_{t-1}\bm{y}_{t-1}^\T +\bm{v}_{t-1}\bm{k}_{t-1}^\T\big]\bm{\varLambda}_t - \bm{\zeta}_t\bm{y}_t^\T + \bm{v}_t\bm{k}_t^\T\\
&=\bm{S}_{t-2}\bm{\varLambda}_{t-1}\bm{\varLambda}_t
+\bm{v}_{t-1}\bm{k}_{t-1}^\T\bm{\varLambda}_t + \bm{v}_t\bm{k}_t^\T
-\bm{\zeta}_{t-1}\bm{y}_{t-1}^\T\bm{\varLambda}_t - \bm{\zeta}_t\bm{y}_t^\T\\
&=\bm{S}_{t-3}\bm{\varLambda}_{t-2}\bm{\varLambda}_{t-1}\bm{\varLambda}_t\\
&+ \bm{v}_{t-2}\bm{k}_{t-2}^\T\bm{\varLambda}_{t-1}\bm{\varLambda}_t + \bm{v}_{t-1}\bm{k}_{t-1}^\T\bm{\varLambda}_t + \bm{v}_t\bm{k}_t^\T\\
&- \bm{\zeta}_{t-2}\bm{y}_{t-2}^\T\bm{\varLambda}_{t-1}\bm{\varLambda}_t-\bm{\zeta}_{t-1}\bm{y}_{t-1}^\T\bm{\varLambda}_t - \bm{\zeta}_t\bm{y}_t^\T\\
&=\bm{S}_0\prod_{\tau=1}^t\bm{\varLambda}_\tau + \sum_{\tau=1}^t \bm{v}_\tau \bm{\kappa}_{\tau}^{(t)\T}-\sum_{\tau=1}^t\bm{\zeta}_\tau\bm{\nu}_\tau^{(t)\T}\\
&=\bm{S}_0\bm{D}_t + \sum_{\tau=1}^t \bm{v}_\tau \bm{\kappa}_{\tau}^{(t)\T}-\sum_{\tau=1}^t\bm{\zeta}_\tau\bm{\nu}_\tau^{(t)\T}
\end{align}$$

其中衰减基底 $\bm{\nu}_\tau^{(t)}= \bm{D}_{\tau+1:t}\bm{y}_\tau$ 和 $\bm{\kappa}_\tau^{(t)} = \bm{D}_{\tau+1:t}\bm{k}_\tau$。通常情况下 $\bm{S}_0= \bm{0}$, 同时令 $\bm{Z}_t^\T=[\bm{\zeta}_1,\cdots,\bm{\zeta}_t]$ 。这样就有

$$\begin{align}
\bm{S}_t
=\bm{S}_0\bm{D}_t + \sum_{\tau=1}^t \bm{v}_\tau \bm{\kappa}_{\tau}^{(t)\T}-\sum_{\tau=1}^t\bm{\zeta}_\tau\bm{\nu}_\tau^{(t)\T}
=\underbrace{\bm{S}_0\bm{D}_t}_{\small\text{衰减的初始信息}} + \underbrace{\underbrace{\bm{V}_t^\T\bm{K}_t}_{\small\text{写入的新信息}}- \underbrace{\bm{Z}_t^\T \bm{Y}_t}_{\small\text{状态中反馈信息}}}_{\small\text{写入的净增信息}}
\end{align}$$

其中: $\bm{V}_t^\T = [\bm{v}_1,\cdots,\bm{v}_t]$ , $\bm{K}_t^\T=[\bm{\kappa}_1,\cdots,\bm{\kappa}_{\tau}^{(t)},\cdots,\bm{\kappa}_t]$ , $\bm{Y}_t^\T=[\bm{\nu}_1,\cdots,\bm{\nu}_\tau^{(t)},\cdots,\bm{\nu}_t]$ , $\bm{Z}_t^\T=[\bm{\zeta}_1,\cdots,\bm{\zeta}_t]$ , $\bm{\kappa}_{\tau}^{(t)}=\bm{D}_{\tau+1:t}\cdot\bm{k}_{\tau}$ , $\bm{\nu}_{\tau}^{(t)}=\bm{D}_{\tau+1:t}\cdot\bm{y}_{\tau}$ 。

对于更新公式,也可以从之前的任意点开始计算,这是分块并行的基础。例如从中间点 $\iota$ 处开始

$$\begin{align}
\bm{S}_t
&=\bm{S}_0\bm{D}_t + \sum_{\tau=1}^t \bm{v}_\tau \bm{\kappa}_{\tau}^{(t)\T}-\sum_{\tau=1}^t\bm{\zeta}_\tau\bm{\nu}_\tau^{(t)\T}\\
&=\bm{S}_0\bm{D}_t + \sum_{\tau=1}^\iota \bm{v}_\tau \bm{\kappa}_{\tau}^{(t)\T}-\sum_{\tau=1}^\iota\bm{\zeta}_\tau\bm{\nu}_\tau^{(t)\T}
+\sum_{\tau=\iota+1}^t \bm{v}_\tau \bm{\kappa}_{\tau}^{(t)\T}-\sum_{\tau=\iota+1}^t\bm{\zeta}_\tau\bm{\nu}_\tau^{(t)\T}\\
&=\bm{S}_0\bm{D}_{\iota} \bm{D}_{\iota+1:t} + \sum_{\tau=1}^\iota \bm{v}_\tau \bm{D}_{\tau+1:t}\bm{k}_\tau^\T-\sum_{\tau=1}^\iota\bm{\zeta}_\tau\bm{D}_{\tau+1:t}\bm{y}_\tau^\T
+\sum_{\tau=\iota+1}^t \bm{v}_\tau \bm{\kappa}_{\tau}^{(t)\T}-\sum_{\tau=\iota+1}^t\bm{\zeta}_\tau\bm{\nu}_\tau^{(t)\T}\\
&=\bigg[\bm{S}_0\bm{D}_{\iota}+ \sum_{\tau=1}^\iota \bm{v}_\tau \bm{D}_{\tau+1:\iota}\bm{k}_\tau^\T-\sum_{\tau=1}^\iota\bm{\zeta}_\tau\bm{D}_{\tau+1:\iota}\bm{y}_\tau^\T\bigg]\bm{D}_{\iota+1:t}
+\sum_{\tau=\iota+1}^t \bm{v}_\tau \bm{\kappa}_{\tau}^{(t)\T}-\sum_{\tau=\iota+1}^t\bm{\zeta}_\tau\bm{\nu}_\tau^{(t)\T}\\
&=\bm{S}_\iota\bm{D}_{\iota+1:t} + \sum_{\tau=\iota+1}^t \bm{v}_\tau \bm{\kappa}_{\tau}^{(t)\T}-\sum_{\tau=\iota+1}^t\bm{\zeta}_\tau\bm{\nu}_\tau^{(t)\T}\\
\end{align}$$

3.3、等价分析
3.3.1、$\bm{\zeta}_t$ 分解

现在分析一下3.1和3.2在遗忘与写入共享键向量的等价关系。这样有利于深刻理解现代RNN设计的精妙之处。这里的关键之处在于理解 $\bm{\zeta}_t$。由定义 $\bm{\zeta}_t\triangleq\bm{S}_{t-1}\bm{x}_t$。在 $\bm{k}_t=\bm{y}_t$(即遗忘键 = 写入键)的前提下,3.1 节已给出 $\bm{S}_{t-1}$ 的 WY 形式:

$$\begin{align}
\bm{S}_{t-1}=\bm{S}_0\big[\bm{D}_{t-1}-\bm{W}_{t-1}^\T\bm{Y}_{t-1}\big]+\bm{U}_{t-1}^\T\bm{K}_{t-1}
\end{align}$$

右乘 $\bm{x}_t$:

$$\begin{align}
\bm{\zeta}_t=\bm{S}_{t-1}\bm{x}_t=\bm{S}_0\underbrace{\big[\bm{D}_{t-1}\bm{x}_t-\bm{W}_{t-1}^\T\bm{Y}_{t-1}\bm{x}_t\big]}_{\bm{w}_t}+\bm{U}_{t-1}^\T\bm{K}_{t-1}\bm{x}_t
\end{align}$$

其中第一项正是 3.1 节构造的 $\bm{w}_t$($\displaystyle \bm{w}_t\triangleq\bm{D}_{t-1}\bm{x}_t-\sum_{\tau}^{t-1}(\bm{\nu}_\tau^{(t-1)\T}\bm{x}_t)\bm{w}_\tau$,其中 $\displaystyle \sum_{\tau}^{t-1}(\bm{\nu}_\tau^{(t-1)\T}\bm{x}_t)\bm{w}_\tau=\bm{W}_{t-1}^\T\bm{Y}_{t-1}\bm{x}_t$ 。第二项展开为:

$$\begin{align}
\bm{U}_{t-1}^\T\bm{K}_{t-1}\bm{x}_t=\sum_{\tau=1}^{t-1}\bm{u}_\tau\big[\bm{\kappa}_\tau^{(t-1)\T}\bm{x}_t\big]
\end{align}$$

在 $\bm{k}_t=\bm{y}_t$ 时有 3.1 节 $\bm{u}_t$ 的构造 $\bm{u}_t\triangleq\bm{v}_t-\sum_{\tau=1}^{t-1}\big[\bm{\kappa}_\tau^{(t-1)\T}\bm{x}_t\big]\bm{u}_\tau$,移项即得:

$$\begin{align}
\sum_{\tau=1}^{t-1}\bm{u}_\tau\big[\bm{\kappa}_\tau^{(t-1)\T}\bm{x}_t\big]=\bm{v}_t-\bm{u}_t
\end{align}$$

因此得到核心分解式

$$\begin{align}
\bm{\zeta}_t=\bm{S}_0\bm{w}_t+\big(\bm{v}_t-\bm{u}_t\big)
\end{align}$$
写成矩阵形式就是

$$\begin{align}
\bm{Z}_t = \bm{W}_t\bm{S}_0^\T+\big[\bm{V}_t-\bm{U}_t\big]
\end{align}$$

3.3.2、等价分析

在 $\bm{k}_t=\bm{y}_t$ 时,$\bm{\kappa}_\tau^{(t)}=\bm{\nu}_\tau^{(t)}$(衰减基底统一),3.2 节结论为:

$$\begin{align}
\bm{S}_t=\bm{S}_0\bm{D}_t+\sum_{\tau=1}^t\bm{v}_\tau\bm{\nu}_\tau^{(t)\T}-\sum_{\tau=1}^t\bm{\zeta}_\tau\bm{\nu}_\tau^{(t)\T}
\end{align}$$

将 $\bm{\zeta}_\tau=\bm{S}_0\bm{w}_\tau+(\bm{v}_\tau-\bm{u}_\tau)$ 代入第三项:

$$\begin{align}
\sum_{\tau=1}^t\bm{\zeta}_\tau\bm{\nu}_\tau^{(t)\T}=\sum_{\tau=1}^t\big[\bm{S}_0\bm{w}_\tau+(\bm{v}_\tau-\bm{u}_\tau)\big]\bm{\nu}_\tau^{(t)\T}=\bm{S}_0\bm{W}_t^\T\bm{Y}_t+\sum_{\tau=1}^t(\bm{v}_\tau-\bm{u}_\tau)\bm{\nu}_\tau^{(t)\T}
\end{align}$$

于是:

$$\begin{aligned}
\bm{S}_t&=\bm{S}_0\bm{D}_t+\sum_{\tau=1}^t\bm{v}_\tau\bm{\nu}_\tau^{(t)\T}-\bm{S}_0\bm{W}_t^\T\bm{Y}_t-\sum_{\tau=1}^t(\bm{v}_\tau-\bm{u}_\tau)\bm{\nu}_\tau^{(t)\T}\\
&=\bm{S}_0\big[\bm{D}_t-\bm{W}_t^\T\bm{Y}_t\big]+\sum_{\tau=1}^t\bm{u}_\tau\bm{\nu}_\tau^{(t)\T}\\
&=\bm{S}_0\big[\bm{D}_t-\bm{W}_t^\T\bm{Y}_t\big]+\bm{U}_t^\T\bm{K}_t
\end{aligned}$$

此即 3.1 节在 $\bm{k}_t=\bm{y}_t$ 下的结果,等价性成立。当 $\bm{S}_0=\bm{0}$(常见初始化)时,$\bm{\zeta}_t=\bm{v}_t-\bm{u}_t$,反馈记忆退化为纯残差,此时 3.2 节的 $\bm{S}_t=\bm{V}_t^\T\bm{K}_t-\bm{Z}_t^\T\bm{Y}_t$ 在 $\bm{k}_t=\bm{y}_t$ 下直接给出 $\bm{S}_t=\sum_\tau\bm{u}_\tau\bm{\nu}_\tau^\T=\bm{U}_t^\T\bm{K}_t$,与 3.1 节完全一致。

3.3.3、关键解释

3.2 节引入 $\bm{\zeta}_t=\bm{S}_{t-1}\bm{x}_t$ 表面是为处理”遗忘键 $\bm{y}_t$ 与写入键 $\bm{k}_t$ 独立”而做的技术性记号,但在 $\bm{y}_t=\bm{k}_t$ 时,它有着相当明确物理意义:当前查询 $\bm{x}_t$ 在旧状态中读取到的反馈信息。或者说记忆状态对当前事件的记忆。

  • $\bm{S}_0\bm{w}_t$ 部分:对应初始状态 $\bm{S}_0$ 经历了”先衰减后被秩一修正”的累积效应,恰是 WY 表示中 $\bm{w}_t$ 所编码的内容: 初始状态对当前查询能反馈出的已被遗忘记忆

  • $\bm{v}_t-\bm{u}_t$ 部分:对应要写入值 $\bm{v}_t$ 扣除要写入到记忆净增量,正是记忆状对当前事件的记忆。

换言之,新增记忆 = 要写记忆 - 反馈记忆 + 被遗忘的初始状态记忆

$$\begin{align}
\bm{u}_t = \bm{v}_t - \bm{\zeta}_t+ \bm{S}_0\bm{w}_t
\end{align}$$

1、这个公式蕴含了一个重要应用, $\bm{u}_t$ 实际上具备恢复初始记忆中被遗忘内容的能力。这样就可以根据需要在 $\bm{S}_0$ 中动态注入记忆,从而帮助模型在必要时检索信息。这个算法层面的实现比工程上的RAG要通用的多。

2、当然恢复初始记忆中被遗忘内容的能力依赖于 $\bm{w}_t$, 而 $\bm{w}_t$ 是否是充分且完备的这需要打一个问号?这很可能是现代RNN下轮的改进方向。

当然如何把知识转换为 $\bm{S}_0$ 去挂载是一个重要的课题。目前改论文《Doc-to-LoRA:Learning to Instantly Internalize Contexts》中使用一个轻量网络把知识转换为 $\bm{S}_0$ 方法可以参考。如果对RWKV的状态更新公式还不熟悉,可以看笔者的文章 RNN的复兴02_什么是记忆

四、信息传播算子

现代RNN其实蕴含了 Songlin Yang等[^1]论文中说的图上的信息传播算子[^4], 下面来详细分析,这将有助于深入理解现代RNN到底在做什么。

4.1、记忆遗忘算子 $\bm{W}$

来考察 $\bm{W}^\T=[\bm{w}_1,\cdots,\bm{w}_t]$的计算,基于 $\displaystyle \bm{w}_t \triangleq \bm{D}_{t-1}\bm{x}_{t} - \sum_{\tau}^{t-1}(\bm{\nu}_\tau ^{(t-1)\T}\bm{x}_{t}) \bm{w}_\tau$, 展开分析,很快就能看到 $\bm{W}$ 编码了节点信息的传播。注意到 $\bm{D}_0=\bm{E}$, 先展开几项,以便获得直觉:

$$\begin{align}
\bm{w}_1&=\bm{x}_1 \\
\bm{w}_2&=\bm{D}_1\bm{x}_2 - \bm{\nu}_1^{(1)\T}\bm{x}_2\cdot \bm{w}_1 \\
\bm{w}_3&=\bm{D}_2\bm{x}_3 - \bm{\nu}_1^{(2)\T}\bm{x}_3\cdot \bm{w}_1-\bm{\nu}_2^{(2)\T}\bm{x}_3\cdot \bm{w}_2 \\
&\vdots\\
\bm{w}_t&=\bm{D}_{t-1}\bm{x}_t - \bm{\nu}_1^{(t-1)\T}\bm{x}_t\cdot \bm{w}_1-\bm{\nu}_2^{(t-1)\T}\bm{x}_t\cdot \bm{w}_2-\cdots-\bm{\nu}_{t-1}^{(t-1)\T}\bm{x}_t\cdot \bm{w}_{t-1}\\
\end{align}$$

写成矩阵形式

$$\begin{align}
\begin{bmatrix}
\bm{w}_1 ^\T\\
\bm{w}_2^\T\\
\bm{w}_3^\T\\
\vdots\\
\bm{w}_t^\T\\
\end{bmatrix}
=\begin{bmatrix}
\bm{x}_1^\T\\
\bm{D}_1\bm{x}_2^\T\\
\bm{D}_2\bm{x}_3^\T\\
\vdots\\
\bm{D}_{t-1}\bm{x}_t^\T\\
\end{bmatrix}
+\begin{bmatrix}
0&0&0&\cdots&0\\
-\bm{\nu}_1^{(1)\T}\bm{x}_2^\T&0&0&\cdots&0\\
-\bm{\nu}_1^{(2)\T}\bm{x}_3&-\bm{\nu}_2^{(2)\T}\bm{x}_3&0&\cdots&0\\
\vdots\\
-\bm{\nu}_1^{(t-1)\T}\bm{x}_t&-\bm{\nu}_2^{(t-1)\T}\bm{x}_t&\cdots&-\bm{\nu}_{t-1}^{(t-1)\T}\bm{x}_t&0
\end{bmatrix}
\begin{bmatrix}
\bm{w}_1^\T\\
\bm{w}_2^\T\\
\bm{w}_3^\T\\
\vdots\\
\bm{w}_t^\T\\
\end{bmatrix}
\end{align}$$

$\def\Afg{\displaystyle{\mathop{A}^{\scriptsize\curvearrowleft}}}\\
\def\Awr{\displaystyle{\mathop{A}^{\small\leadsto}}}$

如果令 $\bm{b}_t=\bm{D}_{t-1}\bm{x}_t$, $\bm{B}^\T=[\bm{x}_1,\cdots,\bm{D}_{t-1}\bm{x}_t]$,于是有 $\displaystyle\bm{W} = \bm{B}+\bm{\Afg}\bm{W}$,易得:

$$\begin{align}
\bm{W} &=\big[\bm{E}-\bm{\Afg}\big]^{-1}\bm{B}\\
\end{align}$$

令 $\bm{T}=\big[\bm{E}-\bm{\Afg}\big]^{-1}$,其中下三角矩阵 $\bm{\Afg}$ 中的元素 ${\Afg}_{ij} = -\bm{\nu}_{j}^{(i-1)\T} \bm{x}_i=-\big[\bm{D}_{j+1:i-1}\cdot\bm{y}_j\big]^\T \bm{x}_i$ , 其中 $i>j$ 。这样就得到了 $\bm{W}$ 的计算方法。这里用 $\curvearrowleft$ 表示遗忘。下三角矩阵求逆有前向替代法Neumann级数展开方法。具体分析可以参考笔者的文章 RNN的复兴03:线性注意力并行计算IPLR

为什么说 $\bm{W}$ 是记忆遗忘算子,使用 Neumann级数[^5] 展开有

$$\begin{align}
\bm{W}
=\big[\bm{E}-\bm{\Afg}\big]^{-1}\bm{B}
=\big[\bm{E}+\bm{\Afg}+\bm{\Afg}^2+\bm{\Afg}^3+\cdots+\bm{\Afg}^\infty\big]\cdot\bm{B}
\end{align}$$

观测其中一行 $\bm{w}_{\tau}$ 实际有

其中 $\tau>i>s>j$。来几个例子,做做数学符号体操, 对了解细节大有裨益,以Neumann级数展开视角,当 $t=4$ 时考察 $\bm{W}$ 其实有

$$\begin{align}
\bm{W}
&=\underbrace{\begin{bmatrix}
\bm{x}_1^\T\\
\bm{D}_1 \bm{x}_2^\T\\
\bm{D}_2 \bm{x}_3^\T\\
\bm{D}_3 \bm{x}_4^\T\\
\end{bmatrix}
}_{自身贡献}
+\underbrace{\begin{bmatrix}
0 & 0 & 0 & 0\
\Afg_{21} & 0 & 0 & 0\\
\Afg_{31} & \Afg_{32} & 0 & 0\\
\Afg_{41} & \Afg_{42} & \Afg_{43} & 0\\
\end{bmatrix}
}_{\bm{\Afg}:相邻贡献}
\begin{bmatrix}
\bm{x}_1^\T\\
\bm{D}_1 \bm{x}_2^\T\\
\bm{D}_2 \bm{x}_3^\T\\
\bm{D}_3 \bm{x}_4^\T\\
\end{bmatrix}\\
&+\underbrace{\begin{bmatrix}
0 & 0 & 0 & 0\
0 & 0 & 0 & 0\\
\Afg_{32}\Afg_{21} & 0 & 0 & 0\\
\Afg_{42}\Afg_{21}+\Afg_{43}\Afg_{31} & \Afg_{43}\Afg_{32} & 0 & 0\\
\end{bmatrix}
}_{\bm{\Afg}^2:跨步影响}
\begin{bmatrix}
\bm{x}_1^\T\\
\bm{D}_1 \bm{x}_2^\T\\
\bm{D}_2 \bm{x}_3^\T\\
\bm{D}_3 \bm{x}_4^\T\\
\end{bmatrix}
+\underbrace{\begin{bmatrix}
0 & 0 & 0 & 0\
0 & 0 & 0 & 0\\
0 & 0 & 0 & 0\\
\Afg_{43}\Afg_{32}\Afg_{21} & 0 & 0 & 0\\
\end{bmatrix}
}_{\bm{\Afg}^3:长程依赖}
\begin{bmatrix}
\bm{x}_1^\T\\
\bm{D}_1 \bm{x}_2^\T\\
\bm{D}_2 \bm{x}_3^\T\\
\bm{D}_3 \bm{x}_4^\T\\
\end{bmatrix}\\
&+
\underbrace{\begin{bmatrix}
0 & 0 & 0 & 0\
0 & 0 & 0 & 0\\
0 & 0 & 0 & 0\\
0 & 0 & 0 & 0\\
\end{bmatrix}
}_{\bm{A}^4}
\begin{bmatrix}
\bm{x}_1^\T\\
\bm{D}_1 \bm{x}_2^\T\\
\bm{D}_2 \bm{x}_3^\T\\
\bm{D}_3 \bm{x}_4^\T\\
\end{bmatrix}+\bm{O}\bm{B}\\
\end{align}$$

展开就有

$$\begin{align}
\bm{w}_1 &= \bm{b}_1\\
\bm{w}_2 &= \bm{b}_2 + \Afg_{21}(\bm{b}_1)\\
\bm{w}_3 &= \bm{b}_3 + \Afg_{31}(\bm{b}_1) + \Afg_{32}(\bm{b}_2) + \Afg_{32}\Afg_{21}(\bm{b}_1)\\
\bm{w}_4 &= \bm{b}_4 + \Afg_{41}(\bm{b}_1) + \Afg_{42}(\bm{b}_2) + \Afg_{43}(\bm{b}_3) \\
&+ \Afg_{42}\Afg_{21}(\bm{b}_1)+\Afg_{43}\Afg_{31}(\bm{b}_1)+\Afg_{43}\Afg_{32}(\bm{b}_2)\\
&+\Afg_{43}\Afg_{32}\Afg_{21}(\bm{b}_1)\\
\end{align}$$

如果以 $\bm{b}_\tau=\bm{D}_{i-1}\bm{x}_i$ 为节点,以下三角矩阵 $\bm{\Afg}$ 中的元素 $\Afg_{ij} = -\bm{\nu}_j^{(i-1)\T}\bm{x}_i= - \big[\bm{D}_{j+1:i-1}\cdot\bm{y}_{j}\big]^\T\bm{x}_i$ 乘积为边, 其中 $i>j$ 。那么 $\bm{w}_\tau$ 的计算可视为一个图:

  • 1、$\bm{w}_\tau$ 表示节点 $\tau$ 的累积路径信息, 编码历史记忆衰减信息。
  • 2、$\bm{x}_\tau$ 是节点的特征向量,影响信息传播的方向
  • 3、$\bm{D}_{\tau-1}$ 是节点更新强度,控制信息权重
  • 4、$\Afg_{ij}$ 是边权值,是历史记忆的衰减系数

这样的路径求和将顺序依赖转化为独立路径的叠加,$\bm{W}$ 的本质是图上的信息传播算子,是记忆遗忘的累积效应。

  • 1、当 $\bm{D}_{\tau-1} \approx 0$ 时:节点孤立(无信息传播)
  • 2、当 $\mathbf{x}_i \perp \mathbf{y}_j$ 时:边权为零(无信息传递)
  • 3、当 $|\mathbf{x}_\tau| \to 0$ 时:节点影响消失

DPLRDeltaNetIPLR本质的不同就是在节点和边上引入了遗忘机制。$\bm{b}_i=\bm{D}_{i-1}\bm{x}_i$ 是自身影响的衰减、$\Afg_{ij} = -\bm{\nu}_j^\T\bm{x}_i= - \big[\bm{D}_{j+1:t}\cdot\bm{y}_{j}\big]^\T\bm{x}_i$ 是相互影响的衰减,进而 记忆遗忘算子 $\bm{W}$ 实际编码了记忆 $\bm{S}_t$ 遗忘转移规律。

4.2、记忆写入算子 $\bm{U}$

现在来考察另外一个重要中间变量:$\displaystyle \bm{u}_t \triangleq \bm{v}_t-\sum_{\tau=1}^{t-1}\bm{u}_\tau\big[\bm{\kappa}_{\tau}^{(t-1)\T}\bm{x}_t\big]=\bm{v}_t-\sum_{\tau=1}^{t-1}\big[\bm{\kappa}_{\tau}^{(t-1)\T}\bm{x}_t\big]\bm{u}_\tau$

$$\begin{align}
\bm{u}_1&= \bm{v}_1\\
\bm{u}_2&= \bm{v}_2 -\bm{\kappa}_1 ^{(1)\T}\bm{x}_{2}\cdot\bm{u}_1\\
\bm{u}_3&= \bm{v}_3 -\bm{\kappa}_1 ^{(2)\T}\bm{x}_{3}\cdot\bm{u}_1-\bm{\kappa}_2 ^{(2)\T}\bm{x}_{3}\cdot\bm{u}_2\\
&\vdots\\
\bm{u}_t&=\bm{v}_t-\bm{\kappa}_1 ^{(t-1)\T}\bm{x}_{t}\cdot\bm{u}_1-\bm{\kappa}_2 ^{(t-1)\T}\bm{x}_{t}\cdot\bm{u}_2-
\cdots-\bm{\kappa}_{t-1} ^{(t-1)\T}\bm{x}_{t}\cdot\bm{u}_{t-1}
\end{align}$$

写成矩阵形式

$$\begin{align}
\begin{bmatrix}
\bm{u}_1 ^\T\\
\bm{u}_2^\T\\
\bm{u}_3^\T\\
\vdots\\
\bm{u}_t^\T\\
\end{bmatrix}
=\begin{bmatrix}
\bm{v}_1^\T\\
\bm{v}_2^\T\\
\bm{v}_3^\T\\
\vdots\\
\bm{v}_t^\T\\
\end{bmatrix}
+\begin{bmatrix}
0&0&0&\cdots&0\\
-\bm{\kappa}_1 ^{(1)\T}\bm{x}_{2}&0&0&\cdots&0\\
-\bm{\kappa}_1 ^{(2)\T}\bm{x}_{3}&-\bm{\kappa}_2 ^{(2)\T}\bm{x}_{3}&0&\cdots&0\\
\vdots\\
-\bm{\kappa}_1 ^{(t-1)\T}\bm{x}_{t}&-\bm{\kappa}_2 ^{(t-1)\T}\bm{x}_{3}&\cdots&-\bm{\kappa}_{t-1} ^{(t-1)\T}\bm{x}_{t}&0
\end{bmatrix}
\begin{bmatrix}
\bm{u}_1^\T\\
\bm{u}_2^\T\\
\bm{u}_3^\T\\
\vdots\\
\bm{u}_t^\T\\
\end{bmatrix}
\end{align}$$

有和 $\bm{W}$一样,可以定义同样的下三角矩阵 $\bm{\Awr}$ 中的元素 $\Awr_{ij} = -\bm{\kappa}_{j} ^{(i-1)\T}\bm{x}_{i}=-\big[\bm{D}_{j+1:i-1}\cdot\bm{k}_j\big]^\T \bm{x}_i$,或者叫邻接矩阵,这里用 $\leadsto$ 表示写入,其中 $i>j$。定义 $\bm{U}^\T=[\bm{u}_1,\cdots,\bm{u}_t]$, $\bm{V}^\T=[\bm{v}_1,\cdots,\bm{v}_t]$

$$\begin{align}
\bm{U} = \bm{V}+\bm{\Awr}\bm{U}
\end{align}$$

$$\begin{align}
\bm{U}
= \big[\bm{E}-\bm{\Awr}\big]^{-1}\bm{V}
=\bm{T}\bm{V}
\end{align}$$

以Neumann级数展开视角

$$\begin{align}
\bm{U}
=\big[\bm{E}-\bm{\Awr}\big]^{-1}\bm{V}
=\big[\bm{E}+\bm{\Awr}+\bm{\Awr}^2+\bm{\Awr}^3+\cdots+\bm{\Awr}^\infty\big]\cdot\bm{V}
\end{align}$$

观察其中一行

$$\begin{align}
\bm{u}_\tau =
\underbrace{\bm{v}_\tau \vphantom{\sum\sum_{j}}}_{\text{自身贡献}}
+\underbrace{\sum_{i}\Awr_{\tau i}\bm{v}_\tau \vphantom{\sum\sum_{j}}}_{\text{相邻贡献}}
+\underbrace{\sum_{i}\sum_{j}\Awr_{\tau i}\Awr_{ij}\bm{v}_j \vphantom{\sum\sum_{j}}}_{\text{跨步影响}}
+\underbrace{\sum_{i}\sum_{s}\sum_{j}\Awr_{\tau i}\Awr_{is}\Awr_{sj}\bm{v}_j \vphantom{\sum\sum_{j}}}_{\text{长程依赖}}
+\cdots
\end{align}$$

当 $t=4$ 时考察 $\bm{U}$ 其实有

$$\begin{align}
\bm{U}
&=
\underbrace{\begin{bmatrix}
\bm{v}_1^\T\\
\bm{v}_2^\T\\
\bm{v}_3^\T\\
\bm{v}_4^\T\\
\end{bmatrix}
}_{自身贡献}
+\underbrace{\begin{bmatrix}
0 & 0 & 0 & 0\
\Awr_{21} & 0 & 0 & 0\\
\Awr_{31} & \Awr_{32} & 0 & 0\\
\Awr_{41} & \Awr_{42} & \Awr_{43} & 0\\
\end{bmatrix}
}_{\bm{\Awr}:相邻贡献}
\begin{bmatrix}
\bm{v}_1^\T\\
\bm{v}_2^\T\\
\bm{v}_3^\T\\
\bm{v}_4^\T\\
\end{bmatrix}\\
&+\underbrace{\begin{bmatrix}
0 & 0 & 0 & 0\
0 & 0 & 0 & 0\\
\Awr_{32}\Awr_{21} & 0 & 0 & 0\\
\Awr_{42}\Awr_{21}+\Awr_{43}\Awr_{31} & \Awr_{43}\Awr_{32} & 0 & 0\\
\end{bmatrix}
}_{\bm{\Awr}^2:跨步影响}
\begin{bmatrix}
\bm{v}_1^\T\\
\bm{v}_2^\T\\
\bm{v}_3^\T\\
\bm{v}_4^\T\\
\end{bmatrix}
+\underbrace{\begin{bmatrix}
0 & 0 & 0 & 0\
0 & 0 & 0 & 0\\
0 & 0 & 0 & 0\\
\Awr_{43}\Awr_{32}\Awr_{21} & 0 & 0 & 0\\
\end{bmatrix}
}_{\bm{\Awr}^3:长程依赖}
\begin{bmatrix}
\bm{v}_1^\T\\
\bm{v}_2^\T\\
\bm{v}_3^\T\\
\bm{v}_4^\T\\
\end{bmatrix}\\
&+
\underbrace{\begin{bmatrix}
0 & 0 & 0 & 0\
0 & 0 & 0 & 0\\
0 & 0 & 0 & 0\\
0 & 0 & 0 & 0\\
\end{bmatrix}
}_{\bm{\Awr}^4}
\begin{bmatrix}
\bm{v}_1^\T\\
\bm{v}_2^\T\\
\bm{v}_3^\T\\
\bm{v}_4^\T\\
\end{bmatrix}+\bm{O}\bm{C}\\
\end{align}$$

展开有
$$\begin{align}
\bm{u}_1 &= \bm{v}_1\\
\bm{u}_2 &= \bm{v}_2 + \Awr_{21}(\bm{v}_1)\\
\bm{u}_3 &= \bm{v}_3 + \Awr_{31}(\bm{v}_1) + \Awr_{32}(\bm{v}_2) + \Awr_{32}\Awr_{21}(\bm{v}_1)\\
\bm{u}_4 &= \bm{v}_4 + \Awr_{41}(\bm{v}_1) + \Awr_{42}(\bm{v}_2) + \Awr_{43}(\bm{v}_3) \\
&+ \Awr_{42}\Awr_{21}(\bm{v}_1)+\Awr_{43}\Awr_{31}(\bm{v}_1)+\Awr_{43}\Awr_{32}(\bm{v}_2)\\
&+\Awr_{43}\Awr_{32}\Awr_{21}(\bm{v}_1)\\
\end{align}$$

如果以 $\bm{v}_\tau$ 为节点,以下三角矩阵 $\bm{\Awr}$ 中的元素 $\Awr_{ij} = -\bm{\kappa}_{j} ^\T\bm{x}_{i}=-\big[\bm{D}_{j+1:i-1}\cdot\bm{k}_j\big]^\T \bm{x}_i$ 乘积为边, 其中 $i>j$。那么 $\bm{u}_\tau$ 的计算可视为一个图:

  • 1、$\bm{u}_\tau$ 表示节点 $\tau$ 的更新路径信息, 编码了记忆注入路径
  • 2、$\bm{v}_\tau$ 是节点的特征向量,影响信息传播的方向
  • 3、这里节点更新强度暂时是1,控制新信息权重,实际上一般都会对 $\bm{v}_\tau$ 乘以标量以控制强度
  • 4、$\Awr_{ij}$ 是边权值,是记忆衰减系数

这样的路径求和将顺序依赖转化为独立路径的叠加,$\bm{U}$ 的本质是图上的新信息传播算子,是新信息对记忆更新的累积效应,它决定了新信息如何写入到记忆 $\bm{S}_t$ 中。

  • 1、当 更新强度为零时:节点孤立(无信息传播)
  • 2、当 $\mathbf{x}_i \perp \mathbf{k}_j$ 时:边权为零(无信息传递)
  • 3、当 $|\mathbf{v}_\tau| \to 0$ 时:节点影响消失
4.3、记忆反馈算子 $\bm{Z}$

对于 RWKV 有一个专门的记忆反馈算子 $\bm{Z}$, 已知 $\bm{\zeta}_t \triangleq \bm{S}_{t-1}\bm{x}_t$。把 3.2 节中 $\bm{S}_{t-1}$ 的展开式代入:

$$\begin{align}
\bm{S}_{t-1} = \bm{S}_0\bm{D}_{t-1} + \sum_{\tau=1}^{t-1}\bm{v}_\tau\bm{\kappa}_\tau^{(t-1)\T} - \sum_{\tau=1}^{t-1}\bm{\zeta}_\tau\bm{\nu}_\tau^{(t-1)\T}
\end{align}$$

右乘 $\bm{x}_t$:

$$\begin{align}
\bm{\zeta}_t
=\underbrace{\bm{S}_0\bm{D}_{t-1}\bm{x}_t\vphantom{\sum_{\tau=1}^{t-1}}}_{\small\text{初始状态贡献}}
+\underbrace{\sum_{\tau=1}^{t-1}\bm{v}_\tau\big(\bm{\kappa}_\tau^{(t-1)\T}\bm{x}_t\big)}_{\small\text{历史写入贡献}}
-\underbrace{\sum_{\tau=1}^{t-1}\bm{\zeta}_\tau\big(\bm{\nu}_\tau^{(t-1)\T}\bm{x}_t\big)}_{\small\text{历史反馈贡献}}
\end{align}
$$

移项整理有:

$$\begin{align}
\bm{\zeta}_t
=\bm{S}_0\bm{D}_{t-1}\bm{x}_t
+\sum_{\tau=1}^{t-1}\big(\bm{\kappa}_\tau^{(t-1)\T}\bm{x}_t\big)\cdot \bm{v}_\tau
-\sum_{\tau=1}^{t-1}\big(\bm{\nu}_\tau^{(t-1)\T}\bm{x}_t\big)\cdot\bm{\zeta}_\tau
\end{align}
$$

与 $\bm{W}$、$\bm{U}$ 的推导完全同构,先把 $\bm{\zeta}_t$ 的递推展开几项,以便获得直觉。注意到 $\bm{D}_0=\bm{E}$,且 $t=1$ 时两个求和为空:

$$\begin{align}
\bm{\zeta}_1&=\bm{S}_0\bm{x}_1\\
\bm{\zeta}_2&=\bm{S}_0\bm{D}_1\bm{x}_2 + \bm{\kappa}_1^{(1)\T}\bm{x}_2\cdot\bm{v}_1-\bm{\nu}_1^{(1)\T}\bm{x}_2\cdot\bm{\zeta}_1\\
\bm{\zeta}_3&=\bm{S}_0\bm{D}_2\bm{x}_3 + \bm{\kappa}_1^{(2)\T}\bm{x}_3\cdot\bm{v}_1+\bm{\kappa}_2^{(2)\T}\bm{x}_3\cdot\bm{v}_2-\bm{\nu}_1^{(2)\T}\bm{x}_3\cdot\bm{\zeta}_1-\bm{\nu}_2^{(2)\T}\bm{x}_3\cdot\bm{\zeta}_2\\
&\vdots\\
\bm{\zeta}_t&=\bm{S}_0\bm{D}_{t-1}\bm{x}_t + \sum_{\tau=1}^{t-1}\bm{\kappa}_\tau^{(t-1)\T}\bm{x}_t\cdot\bm{v}_\tau-\sum_{\tau=1}^{t-1}\bm{\nu}_\tau^{(t-1)\T}\bm{x}_t\cdot\bm{\zeta}_\tau
\end{align}$$

写成矩阵形式
$$\begin{align}
\begin{bmatrix}
\bm{\zeta}_1^\T\\
\bm{\zeta}_2^\T\\
\bm{\zeta}_3^\T\\
\vdots\\
\bm{\zeta}_t^\T
\end{bmatrix}
=\begin{bmatrix}
\bm{x}_1^\T\bm{S}_0^\T\\
\bm{D}_1\bm{x}_2^\T\bm{S}_0^\T\\
\bm{D}_2\bm{x}_3^\T\bm{S}_0^\T\\
\vdots\\
\bm{D}_{t-1}\bm{x}_t^\T\bm{S}_0^\T\\
\end{bmatrix}
-\begin{bmatrix}
0&0&0&\cdots&0\\
-\bm{\kappa}_1 ^{(1)\T}\bm{x}_{2}&0&0&\cdots&0\\
-\bm{\kappa}_1 ^{(2)\T}\bm{x}_{3}&-\bm{\kappa}_2 ^{(2)\T}\bm{x}_{3}&0&\cdots&0\\
\vdots\\
-\bm{\kappa}_1 ^{(t-1)\T}\bm{x}_{t}&-\bm{\kappa}_2 ^{(t-1)\T}\bm{x}_{3}&\cdots&-\bm{\kappa}_{t-1} ^{(t-1)\T}\bm{x}_{t}&0
\end{bmatrix}
\begin{bmatrix}
\bm{v}_1^\T\\
\bm{v}_2^\T\\
\bm{v}_3^\T\\
\vdots\\
\bm{v}_t^\T\\
\end{bmatrix}\\
+\begin{bmatrix}
0&0&0&\cdots&0\\
-\bm{\nu}_1^{(1)\T}\bm{x}_2&0&0&\cdots&0\\
-\bm{\nu}_1^{(2)\T}\bm{x}_3&-\bm{\nu}_2^{(2)\T}\bm{x}_3&0&\cdots&0\\
\vdots\\
-\bm{\nu}_1^{(t-1)\T}\bm{x}_t&-\bm{\nu}_2^{(t-1)\T}\bm{x}_t&\cdots&-\bm{\nu}_{t-1}^{(t-1)\T}\bm{x}_t&0
\end{bmatrix}
\begin{bmatrix}
\bm{\zeta}_1^\T\\
\bm{\zeta}_2^\T\\
\bm{\zeta}_3^\T\\
\vdots\\
\bm{\zeta}_t^\T
\end{bmatrix}
\end{align}$$

观察括号里的标量内积,它们是历史信息 $\bm{k}_\tau$ 或 $\bm{y}_\tau$ 与当前查询 $\bm{x}_t$ 的衰减内积。我们据此定义两个严格下三角矩阵 $\bm{\Awr}$ 和 $\bm{\Afg}$(约定 $i > j$ 时有值,对角线为 0):

$$\begin{align}
\Awr_{ij} &\triangleq-\bm{\kappa}_{j} ^{(i-1)\T}\bm{x}_{i}=-\big[\bm{D}_{j+1:i-1}\cdot\bm{k}_j\big]^\T \bm{x}_i\\
{\Afg}_{ij} &\triangleq -\bm{\nu}_{j}^{(i-1)\T} \bm{x}_i=-\big[\bm{D}_{j+1:i-1}\cdot\bm{y}_j\big]^\T \bm{x}_i
\end{align}$$

定义 $\bm{b}_t=\bm{D}_{t-1}\bm{x}_t$, $\bm{B}_t^\T=[\bm{x}_1,\cdots,\bm{D}_{t-1}\bm{x}_t]$, 写成矩阵形式
$$\begin{align}
\bm{Z}=\bm{B}\bm{S}_0^\T-\bm{\Awr}\bm{V}+\bm{\Afg}\bm{Z}
\end{align}$$

求解有

$$\begin{align}
\bm{Z} = \big[\bm{E}-\bm{\Afg}\big]^{-1}\big[\bm{B}\bm{S}_0^\T-\bm{\Awr}\bm{V}\big]
\end{align}$$

这里记 $\displaystyle \bm{c}_t \triangleq \bm{S}_0\bm{D}_{t-1}\bm{x}_t + \sum_{\tau=1}^{t-1}\big(\bm{\kappa}_\tau^{(t-1)\T}\bm{x}_t\big)\bm{v}_\tau$, 定义 $\bm{Z}^\T=[\bm{\zeta}_1,\cdots,\bm{\zeta}_t]$,$\bm{C}^\T=[\bm{c}_1,\cdots,\bm{c}_t]$。

于是有
$$\begin{align}
\bm{Z} &= \bm{C}+\bm{\Afg}\bm{Z}\\
\bm{Z} &= \big[\bm{E}-\bm{\Afg}\big]^{-1}\bm{C}=\bm{T}\bm{C}
\end{align}$$

值得注意的事实:$\bm{Z}$ 与 $\bm{W}$ 共享同一个传播矩阵 $\bm{T}=\big[\bm{E}-\bm{\Afg}\big]^{-1}$。$\bm{W}=\bm{T}\bm{B}$、$\bm{Z}=\bm{T}\bm{C}$,而 $\bm{U}=\bm{T}_\kappa\bm{V}$ 使用以 $\bm{\kappa}$ 为边的另一传播矩阵 $\bm{T}_\kappa=\big[\bm{E}-\bm{\Awr}\big]^{-1}$。换言之,由 $\bm{\nu}$(衰减后的 $\bm{y}$)驱动的遗忘结构同时支配着擦除算子反馈算子,而由 $\bm{\kappa}$(衰减后的 $\bm{k}$)驱动的结构支配着写入算子。有一个漂亮的分解:

$$\begin{align}
\bm{Z}=\bm{T}\bm{C}=\bm{T}\bm{B}\bm{S}_0^\T-\bm{T}\bm{\Awr}\bm{V}
=\underbrace{\bm{W}\bm{S}_0^\T\vphantom{\sum}}_{\text{初始状态的遗忘传播}}
+\underbrace{-\bm{T}\bm{\Awr}\bm{V}\vphantom{\sum}}_{\text{经遗忘修正的衰减读出}}
\end{align}$$

记忆反馈 = 被遗忘算子传播的初始状态 + 经同一遗忘结构修正的衰减关联读出:$\bm{\kappa}$ 一侧负责写入关联,$\bm{\nu}$ 一侧负责在读取路径上施加遗忘。为什么说 $\bm{Z}$ 是记忆反馈算子,使用 Neumann级数 展开有
$$\begin{align}
\bm{Z}
=\big[\bm{E}-\bm{\Afg}\big]^{-1}\bm{C}
=\big[\bm{E}+\bm{A}+\bm{A}^2+\bm{A}^3+\cdots+\bm{A}^\infty\big]\cdot\bm{C}
\end{align}$$

观测其中一行 $\bm{\zeta}_\tau$ 实际有

$$\begin{align}
\bm{\zeta}_\tau =
\underbrace{\bm{c}_\tau \vphantom{\sum_{j}}}_{\text{源项}}
+\underbrace{\sum_{i}\Afg_{\tau i}\bm{c}_i \vphantom{\sum_{j}}}_{\text{相邻贡献}}
+\underbrace{\sum_{i}\sum_{j}\Afg_{\tau i}\Afg_{ij}\bm{c}_j \vphantom{\sum_{j}}}_{\text{跨步影响}}
+\underbrace{\sum_{i}\sum_{s}\sum_{j}\Afg_{\tau i}\Afg_{is}\Afg_{sj}\bm{c}_j \vphantom{\sum_{j}}}_{\text{长程依赖}}
+\cdots
\end{align}$$

当 $t=4$ 时考察 $\bm{Z}$ 其实有

$$\begin{align}
\bm{Z}
&=
\underbrace{\begin{bmatrix}
\bm{c}_1^\T\\
\bm{c}_2^\T\\
\bm{c}_3^\T\\
\bm{c}_4^\T\\
\end{bmatrix}
}_{\text{源项}}
+\underbrace{\begin{bmatrix}
0 & 0 & 0 & 0\
\Afg_{21} & 0 & 0 & 0\\
\Afg_{31} & \Afg_{32} & 0 & 0\\
\Afg_{41} & \Afg_{42} & \Afg_{43} & 0\\
\end{bmatrix}
}_{\bm{\Afg}:\text{相邻贡献}}
\begin{bmatrix}
\bm{c}_1^\T\\
\bm{c}_2^\T\\
\bm{c}_3^\T\\
\bm{c}_4^\T\\
\end{bmatrix}\\
&+\underbrace{\begin{bmatrix}
0 & 0 & 0 & 0\
0 & 0 & 0 & 0\\
\Afg_{32}\Afg_{21} & 0 & 0 & 0\\
\Afg_{42}\Afg_{21}+\Afg_{43}\Afg_{31} & \Afg_{43}\Afg_{32} & 0 & 0\\
\end{bmatrix}
}_{\bm{\Afg}^2:\text{跨步影响}}
\begin{bmatrix}
\bm{c}_1^\T\\
\bm{c}_2^\T\\
\bm{c}_3^\T\\
\bm{c}_4^\T\\
\end{bmatrix}
+\underbrace{\begin{bmatrix}
0 & 0 & 0 & 0\
0 & 0 & 0 & 0\\
0 & 0 & 0 & 0\\
\Afg_{43}\Afg_{32}\Afg_{21} & 0 & 0 & 0\\
\end{bmatrix}
}_{\bm{\Afg}^3:\text{长程依赖}}
\begin{bmatrix}
\bm{c}_1^\T\\
\bm{c}_2^\T\\
\bm{c}_3^\T\\
\bm{c}_4^\T\\
\end{bmatrix}
+\cdots
\end{align}$$

展开就有
$$\begin{align}
\bm{\zeta}_1 &= \bm{c}_1\\
\bm{\zeta}_2 &= \bm{c}_2 + \Afg_{21}(\bm{c}_1)\\
\bm{\zeta}_3 &= \bm{c}_3 + \Afg_{31}(\bm{c}_1) + \Afg_{32}(\bm{c}_2) + \Afg_{32}\Afg_{21}(\bm{c}_1)\\
\bm{\zeta}_4 &= \bm{c}_4 + \Afg_{41}(\bm{c}_1) + \Afg_{42}(\bm{c}_2) + \Afg_{43}(\bm{c}_3) \\
&+ \Afg_{42}\Afg_{21}(\bm{c}_1)+\Afg_{43}\Afg_{31}(\bm{c}_1)+\Afg_{43}\Afg_{32}(\bm{c}_2)
+\Afg_{43}\Afg_{32}\Afg_{21}(\bm{c}_1)
\end{align}$$

从图论视角看,$\bm{Z}$ 与 $\bm{W}$ 在同一张图上传播——节点相同、边权相同,仅源项从 $\bm{B}=\bm{D}_{t-1}\bm{x}_t$ 替换为包含历史写入贡献的 $\bm{C}$。遗忘结构支配着擦除算子与反馈算子。区别在于:

  • $\bm{W}$ 传播的是纯衰减后的查询特征 $\bm{D}_{t-1}\bm{x}_t$,编码”当前查询在遗忘路径上的累积效应”
  • $\bm{Z}$ 传播的是”初始状态贡献 + 历史写入贡献” $\bm{c}_t$,编码”当前查询在记忆中的实际读取结果”

换言之,$\bm{W}$ 是遗忘路径的骨架,$\bm{Z}$ 则是该骨架上承载了具体记忆内容的实例化。这正是3.3节核心分解式 $\bm{\zeta}_t=\bm{S}_0\bm{w}_t+(\bm{v}_t-\bm{u}_t)$ 的图论体现。

这样的路径求和将顺序依赖转化为独立路径的叠加,$\bm{Z}$ 的本质是图上的记忆读出算子,是记忆反馈的累积效应,它决定了记忆 $\bm{S}_t$ 中的存量信息如何被检索、并以何种强度反馈到当前时刻。特别地,$\bm{\zeta}_t=\bm{S}_{t-1}\bm{x}_t$ 正是 Delta Rule 中”写前先读”的检索信号:在写入新信息之前,先读出当前记忆在该查询方向上的响应,随后的增量更新正是对这份读出结果的修正。

  • 1、当 $\bm{c}_\tau \approx \bm{0}$ 时($\bm{S}_0\approx\bm{0}$ 且历史写入与当前查询正交):节点孤立(无反馈信号)
  • 2、当 $\mathbf{x}_i \perp \mathbf{y}_j$ 时:边权为零(无信息传递)
  • 3、当 $|\mathbf{x}_\tau| \to 0$ 时:节点影响消失
  • 4、当 $\bm{D}_{\tau-1} \approx \bm{0}$ 时:初始状态贡献消失(衰减清空了初始记忆)

还有一个值得玩味的退化视角:若 $\bm{S}_0=\bm{0}$、$\bm{D}=\bm{E}$(无衰减)、且 $\bm{y}=\bm{0}$(无反馈项,$\bm{\Awr}=\bm{O}$),则 $\bm{\zeta}_t=\sum_{\tau<t}(\bm{k}_\tau^\T\bm{x}_t)\bm{v}_\tau$ 恰好退化为因果线性注意力的读出。换言之,$\bm{Z}$ 是带衰减与反馈修正的广义线性注意力读出算子

4.4、三大算子(擦、写、读)的统一图论框架

综合前述分析,现代RNN的三个关键中间变量可统一为图上的信息传播算子:

算子 递推结构 传播矩阵 源项 物理含义
$\bm{W}$ $\bm{W}=\bm{B}+\bm{\Afg}\bm{W}$ $\bm{T}=(\bm{E}-\bm{\Afg})^{-1}$ $\bm{B}=\bm{D}_{t-1}\bm{x}_t$ 遗忘路径累积
$\bm{U}$ $\bm{U}=\bm{V}+\bm{\Awr}\bm{U}$ $\bm{T}_\kappa=(\bm{E}-\bm{\Awr})^{-1}$ $\bm{V}$ 写入路径累积
$\bm{Z}$ $\bm{Z}=\bm{C}+\bm{\Afg}\bm{Z}$ $\bm{T}=(\bm{E}-\bm{\Afg})^{-1}$ $\bm{C}=\bm{B}\bm{S}_0^\T-\bm{\Awr}\bm{V}$ 反馈路径累积

其中邻接矩阵元素:
$$\begin{align}
\Afg_{ij} = -\big[\bm{D}_{j+1:i-1}\cdot\bm{y}_j\big]^\T \bm{x}_i,\qquad
\Awr_{ij} = -\big[\bm{D}_{j+1:i-1}\cdot\bm{k}_j\big]^\T \bm{x}_i,\qquad i>j
\end{align}$$
关键洞察:在共享键向量情形($\bm{y}_t=\bm{k}_t$)下,$\bm{\Afg}=\bm{\Awr}$,三个算子共享同一传播矩阵 $\bm{T}$,此时 $\bm{W}$、$\bm{U}$、$\bm{Z}$ 的差异完全由源项决定。而当遗忘与写入键独立时(如RWKV),则需两个不同的传播矩阵。

擦除与读取共享同一遗忘结构 $\bm{T}$,写入使用 $\bm{T}_\kappa$;当 $\bm{y}=\bm{k}$(即 $\bm{\nu}$ 与 $\bm{\kappa}$ 同源)时,三者退化为经典 DeltaNet 的 IPLR 形式;DPLR/KDA 的贡献正是把节点强度 $\bm{D}$ 与边权 $\bm{\nu}$、$\bm{\kappa}$ 中的衰减机制精细化(对角/通道级门控),这与 Gated DeltaNet 系列把”擦除控制”与”写入控制”解耦的思路一脉相承。而三者最终都归结为下三角矩阵求逆 $\big[\bm{E}-\bm{A}\big]^{-1}$,前向替代与 Neumann 级数(分块并行)两条计算路径同时可用,这正是现代 RNN(DeltaNet/DPLR 家族)能在序列维度上高效并行的根源。

$\def\Afg{\displaystyle{\mathop{A}^{\scriptsize\curvearrowleft}}}\\\def\Awr{\displaystyle{\mathop{A}^{\small\leadsto}}}$

五、状态矩阵并行计算

5.1、DPLR核心矛盾

经过上述分析就能并行计算 DPLR 了?并不能!因为$\bm{\kappa}_{\tau}^{(t)}=\bm{D}_{\tau+1:t}\cdot\bm{k}_{\tau}$ 和 $\bm{\nu}_{\tau}^{(t)}=\bm{D}_{\tau+1:t}\cdot\bm{y}_{\tau}$,一个向量依赖两个变量,这导致 $\bm{\Afg}$ 和 $\bm{\Awr}$ 实际上有三个维度,除了序列维度 $\tau$ 和 特征维度 $d$, 还依赖总的 token 数 $t$ 。这导致无法把 RNN 递推转化为可并行的矩阵乘法。具体来说对于类似Lightning Attention和DetlaNet的这样的线性注意力,通常存在

$$\begin{align}
\bm{o}_t^\T = \sum_{i=1}^{t} (\bm{q}_t^\T \bm{k}_i) \bm{v}_i^\T
\end{align}$$

可以轻松写成矩阵形式 $\bm{O} = (\bm{Q}\bm{K}^\T \odot \bm{M})\bm{V}$。其中 $\bm{M}$ 是掩码矩阵。在 DPLR 中,状态更新是

$$\begin{align}
\bm{S}_t
=\bm{S}_0\bm{D}_t + \sum_{\tau=1}^t \bm{v}_\tau \bm{\kappa}_{\tau}^{(t)\T}-\sum_{\tau=1}^t\bm{\zeta}_\tau\bm{\nu}_\tau^{(t)\T}
\end{align}$$
展开 $\bm{S}_{t}$ 代入输出 $\bm{o}_t^\T = \bm{q}_t^\T\bm{S}_{t}^\T$,得到:

$$\begin{align}
\bm{o}_t^\T
&= \bm{q}_t^\T\bm{D}_t\bm{S}_0^\T + \sum_{\tau=1}^t\bm{q}_t^\T \bm{\kappa}_{\tau}^{(t)}\bm{v}_{\tau}^\T -\sum_{\tau=1}^t\bm{q}_t^\T\bm{\nu}_\tau^{(t)}\bm{\zeta}_{\tau}^\T\\
&=\bm{q}_t^\T\bm{D}_t\bm{S}_0^\T + \sum_{\tau=1}^t\bm{q}_t^\T \bm{D}_{\tau+1:t}\cdot\bm{k}_{\tau}\bm{v}_{\tau}^\T -\sum_{\tau=1}^t\bm{q}_t^\T\bm{D}_{\tau+1:t}\cdot\bm{y}_{\tau}\bm{\zeta}_{\tau}^\T
\end{align}$$

下面形式其实是错误的,除非只计算一个token。

$$\begin{align}
\bm{O}_t = \bm{Q}_t\bm{D}_t\bm{S}_t^\T + (\bm{Q}_t\bm{K}_t^\T\odot \bm{M})\bm{V}-(\bm{Q}_t\bm{Y}_t^\T\odot \bm{M})\bm{Z}
\end{align}$$

因为 $\bm{K}_t$ 和 $\bm{Y}_t$ 的矩阵形式是不存在的。

5.2、衰减商技巧

为了得到下面的形式的计算

$$\begin{align}
\bm{S}_t
=\bm{S}_0\bm{D}_t + \sum_{\tau=1}^t \bm{v}_\tau \bm{\kappa}_{\tau}^\T-\sum_{\tau=1}^t\bm{\zeta}_\tau\bm{\nu}_\tau^\T
=\bm{S}_0\bm{D}_t + \bm{V}_t^\T\bm{K}_t- \bm{Z}_t^\T \bm{Y}_t
\end{align}$$

需要重新定义$\bm{\kappa}_{\tau}^{(t)}$ 和 $\bm{\nu}_{\tau}^{(t)}$,以便摘掉 $(t)$ 的帽子,实际有

$$\begin{align}
\bm{\kappa}_{\tau}^{(t)}
&=\bm{D}_{\tau+1:t}\cdot\bm{k}_{\tau}
=\prod_{i=\tau+1}^t\bm{\varLambda}_{i}\cdot\bm{k}_{\tau}\\
&=\bm{\varLambda}_{\tau+1}\cdots\bm{\varLambda}_{t}\cdot\bm{k}_{\tau}
=\mathrm{diag}[\bm{a}_{\tau+1}\cdots\odot\cdots \bm{a}_t]\cdot\bm{k}_{\tau}\\
&=\mathrm{diag}\bigg[\bigodot_{i=\tau+1}^t\bm{a}_{i}\bigg]\cdot\bm{k}_{\tau}\\
&=\mathrm{diag}\bigg[\big[\prod_{i=\tau+1}^t a_{i1},\cdots,\prod_{i=\tau+1}^t a_{ic},\cdots\big]\bigg]\cdot\bm{k}_{\tau}\\
&=\mathrm{diag}[(\bm{a}_{1}\cdots\bm{a}_t)/(\bm{a}_{1}\cdots \bm{a}_{\tau})]\cdot\bm{k}_{\tau}\\
&=\mathrm{diag}\bigg[\bigodot_{i=1}^t\bm{a}_{i}/\bigodot_{i=1}^{\tau}\bm{a}_{i}\bigg]\cdot\bm{k}_{\tau}\\
&=\bigg(\prod_{i=1}^t\bm{\varLambda}_{i}\oslash \prod_{i=1}^\tau\bm{\varLambda}_{i}\bigg)\cdot\bm{k}_{\tau}\\
&=\big(\bm{D}_{t}\oslash \bm{D}_{\tau}\big)\cdot\bm{k}_{\tau}\\
&=\bm{D}_{t}\bm{D}_{\tau}^{-1}\cdot\bm{k}_{\tau}
\end{align}$$

为了解决上溢和下溢问题,同时精简符号,不妨定义 $g_{tc}=\sum_{i=1}^{t}\ln a_{ic}\;\in\;(-\infty,\,0]$ , $
t\mapsto g_{tc}$ 。 单调不增,负性与单调性直接来自 $a_{ic}\leqslant 1$。这样有

$$\begin{align}
\bm g_t\;\triangleq\;\sum_{i=1}^{t}\ln\bm a_i,\quad \bm g_0\triangleq\bm 0
\end{align}$$

那么有

$$\begin{align}
\bm{\kappa}_{\tau}^{(t)}
&=\bm{D}_{\tau+1:t}\cdot\bm{k}_{\tau}
=\prod_{i=\tau+1}^t\bm{\varLambda}_{i}\cdot\bm{k}_{\tau}\\
&=\mathrm{diag}\bigg[\big[\prod_{i=\tau+1}^t a_{i1},\cdots,\prod_{i=\tau+1}^t a_{ic},\cdots\big]\bigg]\cdot\bm{k}_{\tau}\\
&=\mathrm{diag}\bigg[\exp\big[g_{t1},\cdots,g_{tc},\cdots\big]\div\exp\big[g_{\tau 1},\cdots,g_{\tau c},\cdots\big]\bigg]\cdot\bm{k}_{\tau}\Leftarrow g_{tc}=\sum_{i=1}^{t}\ln a_{ic}\\
&=\mathrm{diag}\bigg[\exp\big[\bm{g}_{t}-\bm{g}_{\tau}\big]\bigg]\cdot\bm{k}_{\tau}\Leftarrow\bm g_t\;\triangleq\;\sum_{i=1}^{t}\ln\bm a_i\\
&=\mathrm{diag}\big[\exp\big[\bm{g}_{t}\big]\big]\oslash\mathrm{diag}\big[\exp\big[\bm{g}_{\tau}\big]\big]\cdot\bm{k}_{\tau}\\
&=\big(\exp\big[\bm{g}_{t}\big]\oslash\exp\big[\bm{g}_{\tau}\big]\big)\odot\bm{k}_{\tau}\\
&=\big(\bm{D}_{t}\oslash \bm{D}_{\tau}\big)\cdot\bm{k}_{\tau}\\
&=\bm{D}_{t}\bm{D}_{\tau}^{-1}\cdot\bm{k}_{\tau}
\end{align}$$

其实有

$$\begin{align}
\bm{D}_{t}\bm{D}_{\tau}^{-1}
=\bigodot_{i=\tau+1}^{t}\bm a_i
=\underbrace{\bigodot_{i=1}^{t}\bm a_i}_{\exp\bm g_t}\;\oslash\;\underbrace{\bigodot_{i=1}^{\tau}\bm a_i}_{\exp\bm g_\tau}
=\exp\bm g_t\oslash\exp\bm g_\tau
=\exp(\bm g_t-\bm g_\tau)
\end{align}$$

这里用了 $\exp$ 与 $\oslash$ 的互化(逐元素 $e^{x}/e^{y}=e^{x-y}$)。区间积由此坍缩成两次查表、一次减法

5.3、输出并行形式
5.3.1、朴素的并行

把 $\bm{\kappa}_{\tau}^{(t)}=\bm{D}_{t}\bm{D}_{\tau}^{-1}\bm{k}_{\tau}$、$\bm{\nu}_{\tau}^{(t)}=\bm{D}_{t}\bm{D}_{\tau}^{-1}\bm{y}_{\tau}$ 代回输出表达式 $\bm{o}_t^\T=\bm{q}_t^\T\bm{D}_t\bm{S}_0^\T+\sum_{\tau=1}^t\bm{q}_t^\T\bm{\kappa}_\tau^{(t)}\bm{v}_\tau^\T-\sum_{\tau=1}^t\bm{q}_t^\T\bm{\nu}_\tau^{(t)}\bm{\zeta}_\tau^\T$。

$$\begin{align}
\bm{o}_t^\T
&=\bm{q}_t^\T\bm{D}_t\bm{S}_0^\T+\sum_{\tau=1}^t\bm{q}_t^\T\bm{D}_{t}\bm{D}_{\tau}^{-1}\bm{k}_{\tau}\bm{v}_\tau^\T-\sum_{\tau=1}^t\bm{q}_t^\T\bm{D}_{t}\bm{D}_{\tau}^{-1}\bm{y}_{\tau}\bm{\zeta}_\tau^\T
\end{align}$$

一个关键观察随之浮现:$\bm{q}_t^\T\bm{\kappa}_\tau^{(t)}=\big(\bm{q}_t^\T\bm{D}_t\big)\big(\bm{D}_{\tau}^{-1}\bm{k}_{\tau}\big)$,尾部公因子 $\bm{D}_t$ 只依赖查询时刻,头部因子 $\bm{D}_{\tau}^{-1}$ 只依赖写入时刻。也就是说,原来区间积 $\bm{D}_{\tau+1:t}$ 被拆成了两端点各自的绝对量之比,节点 $\tau$ 与 $t$ 不再耦合。据此重新定义三组向量:

$$\tilde{\bm{q}}_t \triangleq \bm{D}_t^\T\bm{q}_t=\exp(\bm{g}_t)\odot\bm{q}_t,\qquad
\tilde{\bm{k}}_\tau \triangleq \bm{D}_\tau^{-1}\bm{k}_\tau=\exp(-\bm{g}_\tau)\odot\bm{k}_\tau,\qquad
\tilde{\bm{y}}_\tau \triangleq \bm{D}_\tau^{-1}\bm{y}_\tau=\exp(-\bm{g}_\tau)\odot\bm{y}_\tau$$

立即得到矩阵并行输出公式:

$$\begin{align}
\bm{O}
&=\underbrace{\vphantom{\Big(}\tilde{\bm{Q}}\,\bm{S}_0^\T}_{\text{初始状态读出}}
+\underbrace{\Big(\tilde{\bm{Q}}\tilde{\bm{K}}^\T\odot\bm{M}\Big)\bm{V}}_{\text{写入路径读出}}
-\underbrace{\Big(\tilde{\bm{Q}}\tilde{\bm{Y}}^\T\odot\bm{M}\Big)\bm{Z}}_{\text{反馈路径扣减}}
\end{align}$$

其中 $\bm{M}$ 为因果掩码,$\bm{Z}^\T=[\bm{\zeta}_1,\cdots,\bm{\zeta}_t]$ 已由4.3节给出闭式 $\bm{Z}=\big[\bm{E}-\bm{\Afg}\big]^{-1}\bm{C}$。无法写成矩阵乘法的核心矛盾被彻底化解:序列维度的耦合被完全剥离到了三个独立可预计算的量上——$\tilde{\bm{Q}},\tilde{\bm{K}},\tilde{\bm{Y}}$ 都只是原始向量与一个仅依赖自身时刻的对角缩放,剩下的全是现代GPU最擅长的稠密矩阵乘法。

5.3.2、分块并行

上述公式在代数上是精确的,却隐藏着一个陷阱。注意到 $\tilde{\bm{k}}_\tau$、$\tilde{\bm{y}}_\tau$ 里出现的是 $\exp(-\bm{g}_\tau)$,而 $\bm{g}_\tau=\sum_{i\le\tau}\ln\bm{a}_i\in(-\infty,0]^c$ 单调不增——序列越长,$|g_{\tau c}|$ 越大,$\exp(-g_{\tau c})$ 指数级膨胀,即便 $\tilde{\bm{q}}_t=\exp(\bm{g}_t)\odot\bm{q}_t$ 一侧同步收缩使其乘积 $\exp(\bm{g}_t-\bm{g}_\tau)$ 良态,两端点单独来看早就分别在半精度下发生了上溢或下溢。bf16 动态范围约 $10^{\pm38}$,几万个token下来任何一侧都撑不住。所以还是要回到 块内并行,块间依赖上来。这里定义一些分块符号方便叙述: 块长 $l$, 块数 $B=t/l$, 块号$b\in\{1\cdots B\}$。进而有块首(start) $s=(b-1)l$,块尾(end) $e=bl$。

$$\begin{align}
\underbrace{1 \cdots s}_{\text{历史(已折叠进 }\bm S_{b-1}\text{)}}\;\Big|\;\underbrace{s+1 \cdots \tau \cdots \iota \cdots e}_{\text{本块 }b\text{(}l\text{ 个 token)}}\;\Big|\;\underbrace{e+1 \cdots t}_{\text{未来(不许看)}}
\end{align}$$

下标 身份 在公式里的角色
$s$ 块首(start) 所有帽子的公共参照点:近端从 $s$ 走到 $i$,远端从 $\tau$ 退回 $s$
$e$ 块尾(end) 状态推进的落点;末块 $e=t$
$\iota$ 查询行($s< \iota \leqslant e$) “现在的我”,读出的一侧
$\tau$ 键列($s<\tau\leqslant \iota$) “过去被我看的那个”,被写入的一侧
$t$ 总长/末时刻 只在末块以 $e=t$ 身份出现,全程不参与块内运算

块内的衰减商技巧

$$\begin{align}
\bm{D}_{\iota}\bm{D}_{\tau}^{-1}
&=\bigodot_{i=\tau+1}^{\iota}\bm a_i
=\underbrace{\bigodot_{i=1}^{\iota}\bm a_i}_{\exp\bm g_\iota}\;\oslash\;\underbrace{\bigodot_{i=1}^{\tau}\bm a_i}_{\exp\bm g_\tau}
=\exp\bm g_\iota\oslash\exp\bm g_\tau
=\exp(\bm g_\iota-\bm g_\tau)\\
&=\underbrace{\exp\big[\bm g_\iota-\bm g_s\big]}_{\scriptsize\text{近
端}}\odot\underbrace{\exp\big[\bm g_s-\bm g_\tau\big]}_{\scriptsize\text{远端}}
\end{align}$$
也就是有

$$\begin{align}
\bm{D}_{\iota}\bm{D}_{\tau}^{-1}
=\bm{D}_{\iota}\bm{D}_{s}^{-1}\bm{D}_{s}\bm{D}_{\tau}^{-1}
=\underbrace{\exp\big[\bm g_\iota-\bm g_s\big]}_{\scriptsize\text{近
端}}\odot\underbrace{\exp\big[\bm g_s-\bm g_\tau\big]}_{\scriptsize\text{远端}}
\end{align}$$

为什么用 $\exp\big[\bm g_\iota-\bm g_s\big]\odot\exp\big[\bm g_s-\bm g_\tau\big]$ 替代 $\exp(\bm g_\iota-\bm g_\tau)$,一张图说明问题:

$$\begin{align}
1 \cdots s\;\Big|\;\underbrace{\underbrace{s+1 \cdots \tau}_{\exp\big[\bm g_s-\bm g_\tau\big]} \cdots \iota}_{\exp\big[\bm g_\iota-\bm g_s\big]} \cdots e\;\Big|\;e+1 \cdots t
\\
\underbrace{\underbrace{1 \cdots s\;\Big|\;s+1 \cdots \tau}_{\normalsize\exp\big[\bm{g}_\tau\big]} \cdots \iota}_{\normalsize\exp\big[\bm g_\iota\big]} \cdots e\;\Big|\;e+1 \cdots t
\end{align}$$

$\exp\big[\bm g_s-\bm g_\tau\big]$ 显然要比 $\exp\big[\bm{g}_\tau\big]$ 的数值更加稳定,因为连乘数量少。那么:

$$\begin{align}
\bm{o}_{\iota}^\T
&=\bm{q}_\iota^\T\bm{D}_\iota\bm{S}_0^\T+\sum_{\tau=1}^\iota\bm{q}_\iota^\T\bm{D}_{\iota}\bm{D}_{\tau}^{-1}\bm{k}_{\tau}\bm{v}_\tau^\T-\sum_{\tau=1}^\iota\bm{q}_\iota^\T\bm{D}_{\iota}\bm{D}_{\tau}^{-1}\bm{y}_{\tau}\bm{\zeta}_\tau^\T\\
&=\bm{q}_\iota^\T\bm{D}_\iota\bm{D}_{s}^{-1}\bm{D}_{s}\bm{S}_0^\T+\sum_{\tau=1}^\iota\bm{q}_\iota^\T\bm{D}_{\iota}\bm{D}_{\tau}^{-1}\bm{k}_{\tau}\bm{v}_\tau^\T-\sum_{\tau=1}^\iota\bm{q}_\iota^\T\bm{D}_{\iota}\bm{D}_{\tau}^{-1}\bm{y}_{\tau}\bm{\zeta}_\tau^\T\\
&=\bm{q}_\iota^\T\bm{D}_\iota\bm{D}_{s}^{-1}\bm{D}_{s}\bm{S}_0^\T
+\sum_{\tau=1}^s\bm{q}_\iota^\T\bm{D}_{\iota}\bm{D}_{s}^{-1}\bm{D}_{s}\bm{D}_{\tau}^{-1}\bm{k}_{\tau}\bm{v}_\tau^\T-\sum_{\tau=1}^s\bm{q}_\iota^\T\bm{D}_{\iota}\bm{D}_{s}^{-1}\bm{D}_{s}\bm{D}_{\tau}^{-1}\bm{y}_{\tau}\bm{\zeta}_\tau^\T\\
&+\sum_{\tau=s+1}^\iota\bm{q}_\iota^\T\bm{D}_{\iota}\bm{D}_{s}^{-1}\bm{D}_{s}\bm{D}_{\tau}^{-1}\bm{k}_{\tau}\bm{v}_\tau^\T-\sum_{\tau=s+1}^\iota\bm{q}_\iota^\T\bm{D}_{\iota}\bm{D}_{s}^{-1}\bm{D}_{s}\bm{D}_{\tau}^{-1}\bm{y}_{\tau}\bm{\zeta}_\tau^\T\\
&=\bm{q}_\iota^\T\bm{D}_\iota\bm{D}_{s}^{-1}\bigg(\bm{D}_{s}\bm{S}_0^\T
+\sum_{\tau=1}^s\bm{D}_{s}\bm{D}_{\tau}^{-1}\bm{k}_{\tau}\bm{v}_\tau^\T-\sum_{\tau=1}^s\bm{D}_{s}\bm{D}_{\tau}^{-1}\bm{y}_{\tau}\bm{\zeta}_\tau^\T\bigg)\\
&+\sum_{\tau=s+1}^\iota\bm{q}_\iota^\T\bm{D}_{\iota}\bm{D}_{s}^{-1}\bm{D}_{s}\bm{D}_{\tau}^{-1}\bm{k}_{\tau}\bm{v}_\tau^\T-\sum_{\tau=s+1}^\iota\bm{q}_\iota^\T\bm{D}_{\iota}\bm{D}_{s}^{-1}\bm{D}_{s}\bm{D}_{\tau}^{-1}\bm{y}_{\tau}\bm{\zeta}_\tau^\T\\
&=\bm{q}_\iota^\T\bm{D}_\iota\bm{D}_{s}^{-1}\bm{S}_s^\T
+\sum_{\tau=s+1}^\iota\bm{q}_\iota^\T\bm{D}_{\iota}\bm{D}_{s}^{-1}\bm{D}_{s}\bm{D}_{\tau}^{-1}\bm{k}_{\tau}\bm{v}_\tau^\T-\sum_{\tau=s+1}^\iota\bm{q}_\iota^\T\bm{D}_{\iota}\bm{D}_{s}^{-1}\bm{D}_{s}\bm{D}_{\tau}^{-1}\bm{y}_{\tau}\bm{\zeta}_\tau^\T\\
&=\tilde{\bm{q}}_\iota^\T\bm{S}_s^\T+\sum_{\tau=s+1}^\iota\tilde{\bm{q}}_\iota^\T\tilde{\bm{k}}_\tau\bm{v}_\tau^\T-\sum_{\tau=s+1}^\iota\tilde{\bm{q}}_\iota^\T\tilde{\bm{y}}_\tau\bm{\zeta}_\tau^\T\\
&=\tilde{\bm{q}}_\iota^\T\bm{S}_{b-1}^\T+\sum_{\tau=s+1}^\iota\tilde{\bm{q}}_\iota^\T\tilde{\bm{k}}_\tau\bm{v}_\tau^\T-\sum_{\tau=s+1}^\iota\tilde{\bm{q}}_\iota^\T\tilde{\bm{y}}_\tau\bm{\zeta}_\tau^\T\\
&=\tilde{\bm{q}}_\iota^\T\bigg[\bm{S}_s+\sum_{\tau=s+1}^\iota\bm{v}_\tau\tilde{\bm{k}}_\tau^\T-\sum_{\tau=s+1}^\iota\bm{\zeta}_\tau\tilde{\bm{y}}_\tau^\T\bigg]^\T\\
\end{align}$$

也就是说重新定义块间三组向量:

$$\tilde{\bm{q}}_\iota \triangleq\exp\big[\bm{g}_\iota-\bm{g}_s\big]\odot\bm{q}_\iota,\,
\tilde{\bm{k}}_\tau \triangleq \exp\big[\bm{g}_s-\bm{g}_\tau\big]\odot\bm{k}_\tau,\,
\tilde{\bm{y}}_\tau \triangleq \exp\big[\bm{g}_s-\bm{g}_\tau\big]\odot\bm{y}_\tau$$

写成矩阵形式:

$$\begin{align}
\tilde{\bm{Q}}_{b}\triangleq\bm{Q}_{b}\odot\exp\big[\bm{G}_\iota-\bm{G}_s\big],\,
\tilde{\bm{K}}_{b}\triangleq\bm{K}_{b}\odot\exp\big[\bm{G}_s-\bm{G}_\tau\big],\,
\tilde{\bm{Y}}_{b}\triangleq\bm{Y}_{b}\odot\exp\big[\bm{G}_s-\bm{G}_\tau\big]
\end{align}$$

这个矩阵更新公式就是

$$\begin{align}
\bm{O}_{b}
= \tilde{\bm{Q}}_{b}\bm{S}_{b-1}^\T
+\Big(\tilde{\bm{Q}}_{b}\tilde{\bm{K}}_{b}^\T\odot\bm{M}\Big)\bm{V}_{b}
-\Big(\tilde{\bm{Q}}_{b}\tilde{\bm{Y}}_{b}^\T\odot\bm{M}\Big)\bm{Z}_{b}
\end{align}$$

其中
$$\begin{align}
\big(\tilde{\bm{Q}}_{b}\tilde{\bm{K}}_{b}^\T\odot \bm{M}\big)[\iota]
&=\sum_{\tau=s+1}^e\bm{q}_\iota^\T \bm{D}_{\iota}\bm{D}_\tau^{-1}\bm{k}_\tau \bm{v}_\tau\\
&=\sum_{\tau=s+1}^e\bm{q}_\iota^\T\big(\exp\big[\bm{g}_\iota-\bm{g}_s\big]\odot\exp\big[\bm{g}_s-\bm{g}_\tau\big]\odot\bm{k}_\tau\big) \bm{v}_\tau\\
&=\sum_{\tau=s+1}^e\big[\exp\big[\bm{g}_\iota-\bm{g}_s\big]\odot\bm{q}_\iota\big]^\T\cdot\big[\exp\big[\bm{g}_s-\bm{g}_\tau\big]\odot\bm{k}_\tau\big]\cdot\bm{v}_\tau\\
&=\Bigg[\bigg[\bm{Q}_{b}\odot\exp\big[\bm{G}_\iota-\bm{G}_s\big]\bigg]\cdot\bigg[\bm{K}_{b}\odot\exp\big[\bm{G}_s-\bm{G}_\tau\big]\bigg]^\T\odot \bm{M}\Bigg][\iota]
\end{align}$$

最后再来考察一下记忆更新

$$\begin{align}
\bm{S}_{b}&=\bm{S}_{s}\bm{D}_{s+1:e}+ \sum_{\tau=s+1}^e \bm{v}_\tau \bigg[\exp\big[\bm{g}_e-\bm{g}_\tau\big]\odot\bm{k}_\tau\bigg]^\T-\sum_{\tau=s+1}^e\bm{\zeta}_\tau\bigg[\exp\big[\bm{g}_e-\bm{g}_\tau\big]\odot\bm{y}_\tau\bigg]^\T\\
&=\bm{S}_{s}\odot\exp\big[\bm{g}_e-\bm{g}_s\big]
+\sum_{\tau=s+1}^e \bm{v}_\tau \bigg[\exp\big[\bm{g}_e-\bm{g}_\tau\big]\odot\bm{k}_\tau\bigg]^\T
-\sum_{\tau=s+1}^e\bm{\zeta}_\tau\bigg[\exp\big[\bm{g}_e-\bm{g}_\tau\big]\odot\bm{y}_\tau\bigg]^\T\\
&=\bm{S}_s\mathrm{Diag}[\exp[\bm g_e-\bm g_s]]
+\sum_{\tau=s+1}^e \bm{v}_\tau \tilde{\bm{k}}_\tau^\T
-\sum_{\tau=s+1}^e\bm{\zeta}_\tau\tilde{\bm{y}}_\tau^\T\\
&=\bm{S}_s \bm{D}_{s+1:e} + \bm{\bar{V}}_b^\T\bm{\bar{K}}_b-\bm{\bar{Z}}_b^\T\bm{\bar{Y}}_b
\end{align}$$

也就说记忆更新还是要单独计算。或者取 $\bm{G}$系列矩阵的最后一行

5.3.3、块间状态递推

先从一个简单例子开始,分成三块

$$\begin{align}
\begin{bmatrix}
\bm{O}_1\\
\bm{O}_2\\
\bm{O}_3\\
\end{bmatrix}&=
\left(\begin{bmatrix}
\tilde{\bm{Q}}_1\\
\tilde{\bm{Q}}_2\\
\tilde{\bm{Q}}_3\\
\end{bmatrix}
\begin{bmatrix}
\bm{S}_0^\T & \bm{S}_1^\T & \bm{S}_2^\T
\end{bmatrix}\quad\odot
\begin{bmatrix}
\bm{I}_1 & \bm{0} &\bm{0} \\
\bm{0} & \bm{I}_2 &\bm{0} \\
\bm{0} & \bm{0} &\bm{I}_3 \\
\end{bmatrix}\right)\begin{bmatrix}
\bm{E}_1\\
\bm{E}_2\\
\bm{E}_3\\
\end{bmatrix}\\&+
\left(
\begin{bmatrix}
\tilde{\bm{Q}}_1\\
\tilde{\bm{Q}}_2\\
\tilde{\bm{Q}}_3\\
\end{bmatrix}
\begin{bmatrix}
\tilde{\bm{K}}_1^\T & \tilde{\bm{K}}_2^\T & \tilde{\bm{K}}_3^\T
\end{bmatrix}\,\odot
\begin{bmatrix}
\bm{M}_1 & \bm{0} &\bm{0} \\
\bm{0} & \bm{M}_2 &\bm{0} \\
\bm{0} & \bm{0} &\bm{M}_3 \\
\end{bmatrix}\right)
\begin{bmatrix}
\bm{V}_1\\
\bm{V}_2\\
\bm{V}_3\\
\end{bmatrix}
\\&-
\left(
\begin{bmatrix}
\tilde{\bm{Q}}_1\\
\tilde{\bm{Q}}_2\\
\tilde{\bm{Q}}_3\\
\end{bmatrix}
\begin{bmatrix}
\tilde{\bm{Y}}_1^\T & \tilde{\bm{Y}}_2^\T & \tilde{\bm{Y}}_3^\T
\end{bmatrix}\,\,\,\odot
\begin{bmatrix}
\bm{M}_1 & \bm{0} &\bm{0} \\
\bm{0} & \bm{M}_2 &\bm{0} \\
\bm{0} & \bm{0} &\bm{M}_3 \\
\end{bmatrix}\right)
\begin{bmatrix}
\bm{Z}_1\\
\bm{Z}_2\\
\bm{Z}_3\\
\end{bmatrix}
\end{align}$$

观察局部有

$$\begin{align}
\bm{O}_{b}
= \tilde{\bm{Q}}_{b}\bm{S}_{b}^\T
+\Big(\tilde{\bm{Q}}_{b}\tilde{\bm{K}}_{b}^\T\odot\bm{M}\Big)\bm{V}_{b}
-\Big(\tilde{\bm{Q}}_{b}\tilde{\bm{Y}}_{b}^\T\odot\bm{M}\Big)\bm{Z}_{b}
\end{align}$$

对写入项与扣减项同时做”近端—远端”二次拆分,$\exp[\bm g_e-\bm g_\tau]=\exp[\bm g_e-\bm g_s]\odot\exp[\bm g_s-\bm g_\tau]$,其中 $\exp[\bm g_e-\bm g_s]$ 与 $\tau$ 无关,可整体提出为块端对角因子。定义块端对角矩阵与块端键矩阵:

$$\begin{align}
\bm{\varPhi}_{b}\triangleq\mathrm{Diag}\big[\exp\big[\bm{g}_e-\bm{g}_s\big]\big],\qquad
\bar{\bm{K}}_{b}\triangleq\bm{K}_{b}\odot\exp\big[\bm{G}_e-\bm{G}_\tau\big],\qquad
\bar{\bm{Y}}_{b}\triangleq\bm{Y}_{b}\odot\exp\big[\bm{G}_e-\bm{G}_\tau\big]
\end{align}$$

则块间状态递推收敛为如下紧凑形式:

$$\begin{align}
\bm{S}_{b}=\bm{S}_{s}\,\bm{\varPhi}_{b}+ \bm{\bar{V}}_b^\T\bm{\bar{K}}_b-\bm{\bar{Z}}_b^\T\bm{\bar{Y}}_b
\end{align}$$

这正是块粒度下的 $\bm{S}_t=\bm{S}_0\bm{D}_t+\bm{V}^\T\bm{K}-\bm{Z}^\T\bm{Y}$:三项分别为衰减的块初始状态、块内净写入、块内净擦除,块与块之间仅通过 $\bm{S}_s$ 传递依赖,块内全部计算可稠密矩阵化。

5.3.4、衰减矩阵计算

先把根上的对象写出来。对块内 token $s+1,\cdots,e$(设 $l=4$)。$\bm{Q}$ 侧近端因子

$$\begin{align}
\bm{P}\triangleq\exp\big[\bm{G}_\iota-\bm{G}_s\big]=
\begin{bmatrix}
\exp[\bm g_{s+1}-\bm g_s]\\
\exp[\bm g_{s+2}-\bm g_s]\\
\exp[\bm g_{s+3}-\bm g_s]\\
\exp[\bm g_{e\quad}-\bm g_s]\\
\end{bmatrix}
\in(0,1]^{l\times c}
\end{align}$$

$\bm{K/Y }$ 侧远端因子

$$\begin{align}
\bm{F}\triangleq\exp\big[\bm{G}_s-\bm{G}_\tau\big]=
\begin{bmatrix}
\exp[\bm g_s-\bm g_{s+1}]\\
\exp[\bm g_s-\bm g_{s+2}]\\
\exp[\bm g_s-\bm g_{s+3}]\\
\exp[\bm g_s-\bm g_{e\quad}]\\
\end{bmatrix}
\in[1,\mathrm{e}^{|\bm g_s-\bm g_e|}]^{l\times c}
\end{align}$$

状态更新因子

$$\begin{align}
\bm{R}\triangleq\exp\big[\bm{G}_e-\bm{G}_\tau\big]=
\begin{bmatrix}
\exp[\bm g_e-\bm g_{s+1}]\\
\exp[\bm g_e-\bm g_{s+2}]\\
\exp[\bm g_e-\bm g_{s+3}]\\
\bm 1\\
\end{bmatrix}\in(0,1]^{l\times c}
\end{align}$$

三张切片不是独立的,可被一个公共向量缝合。记切片$\bm{Q}$ 侧近端因子的最后一行

$$\begin{align}
\bm{\phi}_b\;\triangleq\;\exp\big[\bm g_e-\bm g_s\big]\quad(\leqslant 1)
\end{align}$$

则逐行验证(对每个 $\tau$):
$$\exp[\bm g_e-\bm g_\tau]=\underbrace{\exp[\bm g_e-\bm g_s]}_{\bm\phi_b}\odot\underbrace{\exp[\bm g_s-\bm g_\tau]}_{\bm F\text{ 的对应行}}
\quad\Longrightarrow\quad
\boxed{\;\bm R=\bm\phi_b\odot\bm F\;}$$

即:状态更新因子 = 远端因子矩阵 ⊙ 近端因子矩阵的最后一行(广播)。同一个 $\bm\phi_b\leqslant1$,一手把 $\bm F$ 的膨胀因子压回 $(0,1]$,一手把 $\bm S_s$ 收缩——压状态与驯因子是同一个动作

状态更新因子不必从零算,计算只是一个广播缩放:

$$\begin{align}
\bar{\bm K}_b
=\bm{\tilde{K}}_b\odot\bm{R}
=\underbrace{\big(\bm K_b\odot\bm F\big)}_{\tilde{\bm K}_b\,\text{输出路径已算}}\odot\bm\phi_b,\qquad
\bar{\bm Y}_b=\tilde{\bm Y}_b\odot\bm\phi_b,\qquad
\bm\varPhi_b=\mathrm{Diag}(\bm\phi_b)
\end{align}$$

记忆更新的全部新原料只有一个向量 $\bm\phi_b$。增量成本 = 一次广播乘(把 $\tilde{\bm K}_b,\tilde{\bm Y}_b$ 换规范到 $e$)+ 一次对角缩放($\bm S_s\bm{\varPhi}_b$)+ 两次矩阵乘($\bm V_b^\T\bar{\bm K}_b$、$\bm Z_b^\T\bar{\bm Y}_b$)。

5.3.5、KDA/GDN 标准形式

当遗忘与写入共享键向量($\bm{y}_t=\bm{k}_t$)时,借助 3.3 节的核心分解 $\bm{\zeta}_t=\bm{S}_0\bm{w}_t+(\bm{v}_t-\bm{u}_t)$,块内以 $\bm{S}_s$ 为参照重写为 $\bm{\zeta}_\tau=\bm{S}_s\bm{w}_\tau^{(b)}+(\bm{v}_\tau-\bm{u}_\tau)$,即 $\bm{Z}_b=\bm{W}_b\bm{S}_s^\T+\bm{V}_b-\bm{U}_b$。代入递推式:

$$\begin{align}
\bm{S}_{b}
&=\bm{S}_s\bm{\varPhi}_b+\bm{V}_b^\T\bar{\bm{K}}_b-\big[\bm{W}_b\bm{S}_s^\T+\bm{V}_b-\bm{U}_b\big]^\T\bar{\bm{K}}_b\\
&=\bm{S}_s\bm{\varPhi}_b+\big[\bm{U}_b-\bm{W}_b\bm{S}_s^\T\big]^\T\bar{\bm{K}}_b
\end{align}$$

输出侧同理,利用 $\bm{\tilde{K}}_b^\T(\bm{V}_b-\bm{Z}_b)=\tilde{\bm{K}}_b^\T[\bm{U}_b-\bm{W}_b\bm{S}_s^\T]$,得

$$\begin{align}
\bm{O}_b=\tilde{\bm{Q}}_b\bm{S}_s^\T+\Big(\tilde{\bm{Q}}_b\tilde{\bm{K}}_b^\T\odot\bm{M}\Big)\Big[\bm{U}_b-\bm{W}_b\bm{S}_s^\T\Big]
\end{align}$$

这与 Kimi Linear[^2] 论文中的 chunkwise 形式完全一致。$\bm{U}_b-\bm{W}_b\bm{S}_s^\T$ 把擦除与写入打包为一个等效写入,$\bm{Z}$ 无需跨块传递,块间只需维护 $\bm{S}_s$。而在 RWKV7 这类独立键情形下,$\bm{\kappa}$ 与 $\bm{\nu}$ 分属两条传播链($\bm{T}_\kappa$ 与 $\bm{T}$),$\bm{Z}$ 需按 4.3 节闭式在块内单独求解,这正是 DPLR 比 KDA 多出的那份通用性代价。

5.3.5、分块并行算法全貌
步骤 计算内容 并行性 关键操作
1、块内预计算 $\tilde{\bm{Q}}_b,\tilde{\bm{K}}_b,\tilde{\bm{Y}}_b$(对数衰减差查表) 全块并行 $\exp[\bm g_\iota-\bm g_s]$、$\exp[\bm g_s-\bm g_\tau]$
2、块内三算子 $\bm{W}_b=\bm{T}_b\tilde{\bm{B}}_b$、$\bm{U}_b=\bm{T}_{\kappa,b}\bm{V}_b$、$\bm{Z}_b=\bm{T}_b\bm{C}_b$ 块间并行、块内前向替代 下三角矩阵求逆 $[\bm{E}-\bm{A}]^{-1}$
3、块间状态推进 $\bm{S}_{b+1}=\bm{S}_s\bm{\varPhi}_b+\bar{\bm{K}}_b^\T\bm{V}_b-\bar{\bm{Y}}_b^\T\bm{Z}_b$ 串行(块粒度扫描) 一次块端衰减 + 两次矩阵乘
4、块内输出读出 $\bm{O}_b=\tilde{\bm{Q}}_b\bm{S}_s^\T+\Big(\tilde{\bm{Q}}_b\tilde{\bm{K}}_b^\T\odot\bm{M}\Big)\Big[\bm{U}_b-\bm{W}_b\bm{S}_s^\T\Big]$ 全块并行 因果掩码下的稠密矩阵乘

DPLR 的并行化路径可以完整地串起来块内并行,块间依赖

第一步,用 $\bm{D}_t\oslash\bm{D}_\tau$ 的衰减商技巧把区间积坍缩为两次查表一次减法,将 $\bm{\kappa}_\tau^{(t)}$、$\bm{\nu}_\tau^{(t)}$ 中对 $(\tau,t)$ 的双重依赖剥离为端点各自的绝对量,化解了”状态矩阵无法写成矩阵乘法”的核心矛盾;
第二步,用块内参照点 $\bm s$ 把大区间拆为”近端收缩 × 远端有界”的乘积,解决半精度上下溢;
第三步,块内沿用 4.4 节统一图论框架,擦除 $\bm W$、写入 $\bm U$、反馈 $\bm Z$ 三个算子各自归结为一次下三角求逆,前向替代与 Neumann 分块两条路径通用;
第四步,块间以 $\bm{S}_{b}=\bm{S}_{s}\,\bm{\varPhi}_{b}+ \bm{\bar{V}}_b^\T\bm{\bar{K}}_b-\bm{\bar{Z}}_b^\T\bm{\bar{Y}}_b$ 串行扫描,共享键情形下 $\bm{U}-\bm{W}\bm{S}_s^\T$ 进一步消去 $\bm{Z}$、减少矩阵乘次数。

六、评述

最后回到 3.3.3 节留下的两个开放问题。

1、其一,$\bm{u}_t$ 恢复初始记忆的能力依赖 $\bm{w}_t$ 的充分性,从图论视角看 $\bm{W}$ 只传播了衰减骨架上的查询特征,若 $\bm{S}_0$ 中的记忆方向与所有历史查询正交,则 $\bm{w}_t$ 无法将其读出——这意味着”知识挂载到 $\bm{S}_0$”的设计必须与查询分布对齐,Doc-to-LoRA 一类方法本质是在学习这个对齐。

2、其二,$\bm{\kappa}$ 与 $\bm{\nu}$ 的解耦程度决定了 $\bm{T}$ 与 $\bm{T}_\kappa$ 的分化程度:KDA 通过令二者同源换来近 2× 的核加速,RWKV7 则保留独立键换取更细的记忆编辑自由度,Gated DeltaNet-2 进一步在通道级解耦擦除与写入——三者的取舍构成了当前线性注意力架构设计的核心。DPLR 并行计算的意义不止于工程实现,它把”记忆如何遗忘、写入、反馈”这三件事翻译成了同一张图上可并行的路径求和,让算法层的状态操纵与硬件层的张量核心咬合在一起。

最后线性注意力表达能力的终点在哪里,且看下回分解。

参考文献:

[^1]: Yang, S., Wang, B., Zhang, Y., Shen, Y., & Kim, Y. (2025, January 15). Parallelizing linear transformers with the delta rule over sequence length. arXiv. https://doi.org/10.48550/arXiv.2406.06484
[^2]: Kimi Team. “Kimi Linear: An Expressive, Efficient Attention Architecture.” arXiv:2510.26692.
[^3]: Peng, B., et al. “RWKV-7 ‘Goose’ with Expressive Dynamic State Evolution.” COLM 2025.
[^4]: Yang, S. “DeltaNet Explained (Part II).” Blog post.
[^5]: Huang, Y., et al. “Multiplication-Only Matrix Inversion Approximation for Chunk-wise Linear Attention.” arXiv:2606.06034.


版权声明
引线小白创作并维护的柠檬CC博客采用署名-非商业-禁止演绎4.0国际许可证。
本文首发于柠檬CC [ https://www.limoncc.com ] , 版权所有、侵权必究。
本文永久链接httpss://www.limoncc.com/post/9e070b6858f0e490/
如果您需要引用本文,请参考:
引线小白. (Aug. 18, 2026). 《RNN的复兴04:线性注意力并行计算DPLR》[Blog post]. Retrieved from https://www.limoncc.com/post/9e070b6858f0e490
@online{limoncc-9e070b6858f0e490,
title={RNN的复兴04:线性注意力并行计算DPLR},
author={引线小白},
year={2026},
month={Aug},
date={18},
url={\url{https://www.limoncc.com/post/9e070b6858f0e490}},
}

'