causal attentionのKVキャッシュを固定サイズの行列に圧縮する最近のLinear Attentionの系譜の各種手法の原理を系統的にまとめる。本記事は本件について筆者自身の数理的な技術背景の理解の過程を残したもので、記号や導出が一般的な解説と異なる可能性があり、必ずしも読みやすいものではないかもしれない。
Masked self-attention
inputの表現の配列が$x_i \in \mathbb{R}^d, i=1,...,L$であるとする。任意長の配列の要素間の相互作用をもとに配列を更新する学習可能な機構として、masked self-attentionは非常に強力でよく使われている。
masked self-attention $x_i \rightarrow h_i$は、一般的な定義で
\begin{eqnarray}
h_i &=& \sum_{j=1}^{L} a_{i, j} v_j \\
a_{i, j} &=& \frac{M_{i, j} e^{q_i^T k_j}}{\sum_{s=1}^{L} M_{i, s} e^{q_i^T k_s}} \\
q_i &=& W_Q x_i \\
k_i &=& W_K x_i \\
v_i &=& W_V x_i
\end{eqnarray}
の演算である。$W_q, W_k, W_v$は$i$に依存しない学習可能な行列であり、$M_{i, j}$がattention mask。Q, K, Vはquery, key, valueの記号であり、配列インデクス$i$のクエリ$q_i$について、配列全体のkey $k_1,...,k_L$との内積$q_i^T k_j$のexponantialを重み(attention map)として、配列全体のvalue $v_1, ..., v_k$の重みつき平均をとったものが$h_i$である。self-attentionは本質的にはただの配列全体のvalueの重み付き平均である。その重みがqueryとkeyの内積で決まる。attention maskは$\{0, 1\}$の行列で、$0$のquery, keyの組み合わせのattentionを$0$にする。$M_{i, j} = 0$であれば、$h_i$の計算において、$x_j$は計算に使用しないということになる。
$X =[x_1,...,x_L] \in \mathbb{R}^{d \times L}, H =[h_1,...,h_L] \in \mathbb{R}^{d \times L}$として行列で表現すると、
\begin{eqnarray}
H &=& V A^T \\
A &=& \mathrm{softmax}_{\mathrm{col}}(M \odot (Q^T K)) \\
Q &=& W_Q X \\
K &=& W_K X \\
V &=& W_V X
\end{eqnarray}
となる。ただし、$\mathrm{softmax}_{\mathrm{col}}$は行列の要素ごとのexponentialを計算して、列方向に和が$1$になるように規格化する演算で。$\odot$はアダマール積。
Causal attention
GPTなどで使用されているdecoder-only型のTransformerではcausal maskというattention maskのself-attentionを使用する。このmaskでは、$h_i$の計算は$x_1,...,x_{i}$だけで行い、$x_{i+1}$以降は使用しないという後(未来)の知識を参照しえないという意味で"Causal"ということである。causal maskのself-attentionをcalsal attentionと言ったり、calsal attentionを使うTransformerをCausal Transformerと言ったりする。次のtokenを予測する形式のTransfomerは基本的にこのマスクを使っている。
$M_{i, j}$は以下。
M_{i,j} = \begin{cases}1 & (j \le i) \\ 0 & (j > i)\end{cases}
これをmasked self-attentionの定義に入れて式を整理すると、
h_i = \sum_{j=1}^{i} \frac{e^{q_i^T k_j}}{\sum_{s=1}^{i} e^{q_i^T k_s}} v_j
となる。$h_1,...,h_{i-1}$を計算する過程で$x_1,...,x_{i-1}$から計算された$k_1,...,k_{i-1}$および$v_1,...,v_{i-1}$が計算済みとすると、新たに$x_{i}$から$q_{i}, k_{i}, v_{i}$だけを計算することで$q_{i}, k_{i}, v_{i}$と計算済みの$k_1,...,k_{i-1}, v_1,...,v_{i-1}$から$h_i$を計算できる。attention mapの保持は必要がなく、$q_1,...,q_{j-1}$も不要である。逐次的に計算する際に計算済みのkeyとvalueだけを保持するKVキャッシュを利用することで、次の要素を効率的に計算できる。見方を変えると、Causal TransformerではKVキャッシュが次の要素を予測する ためのcontextの全てである。
この逐次計算について行列で書くと、$K_i=[k_1,...,k_{i-1}, k_{i}], V_{i} = [v_1, ..., v_{i-1}, v_i]$とすれば、
h_i = V_i \mathrm{softmax}_{\mathrm{col}} (q_i^T K_i)^T
である。
Linear attention
Linear attentionはattention mapを計算する際に$0 \le a_{i, j} \le 1$かつkey方向の和が常に1になるように規格化する非線形演算のsoftmaxを排除して単純化したcausal attentionである。
h_i = \sum_{j=1}^{i} (q_i^T k_j) v_j
行列で書くと、
h_i = V_i K_i^T q_i
となる。$V_i K_i^T \in \mathbb{R}^{d \times d}$による$q_i$の線形変換のような形で表現される。$S_i = V_i K_i^T$とすると、
S_i = [v_1,...,v_{i-1}, v_i] \begin{bmatrix} k_1^T \\ \vdots \\ k_{i-1}^T \\ k_i^T\end{bmatrix} = \sum_{s=1}^{i} v_s k_s^T
であり、
S_i = S_{i-1} + v_i k_i^T
である。
つまり、
h_i = S_i q_i = (S_{i-1} + v_i k_i^T) q_i
となるから、$h_i$は$S_{i-1}$と$q_i, k_i, v_i$だけで計算できる。Linear attentionではKVキャッシュは不要で、$S_i$だけ逐次更新して保持すればよい。言い換えると、$h_i$を計算するためのcontextであるKVキャッシュは$S_i \in \mathbb{R}^{d \times d}$の行列へ圧縮される。KVキャッシュは$K_i, V_i \in \mathbb{R}^{d \times i}$であり前の配列長に依存して大きくなるが、$S_i$は常に固定の大きさである。
causal attentionからsoftmaxを排除しただけなので、通常のmasked self-attention同様に、逐次計算でなくcausal mask $M$を使って配列全体の$H = [S_1 q_1, S_2 q_2,...,S_L q_L]$を行列演算で一度に計算することもできる。
\begin{eqnarray}
H &=& V (M \odot (Q^T K))^T \\
&=& V (M^T \odot (K^T Q))
\end{eqnarray}
Mamba2
NVIDIAのopen-source modelのNemotronのnano 2以降で使われている。
Mamba2は数式としてはLinear Attentionの逐次更新のシンプルな拡張である。
\begin{eqnarray}
h_i &=& S_i q_i \\
S_i &=& \alpha_i S_{i-1} + v_i k_i^T \\
\alpha_i &=& \sigma (f_{\theta} (x_i))
\end{eqnarray}
Linear Attentionとの違いは、$S_i$の更新式で$S_{i-1}$に係数$\alpha_i$が付いただけである。
$\alpha_i$は$x_i$から計算される、配列の別の要素には依存しない係数であり、具体的には学習可能なMLP等の関数とsigmoidで計算される。$0 \le \alpha_i \le 1$であり、過去のcontextをどの程度記憶し続けるかを表すgateのような役割である。
Mamba2は逐次更新の式で定義されるが、行列演算で一気に計算することもできる。
\begin{eqnarray}
S_i &=&\alpha_i S_{i-1} + v_i k_i^T \\
&=& \alpha_i (\alpha_{i-1} S_{i-2} + v_{i-1} k_{i-1}^T) + v_i k_i^T \\
&\vdots& \\
&=& \left(\prod_{t=1}^{i} \alpha_{t} \right) S_0 + \sum_{s=1}^{i} \left( \prod_{t=s+1}^{i} \alpha_{t} \right) v_s k_s^T
\end{eqnarray}
なので、
\gamma_i = \prod_{t=1}^i \alpha_t
とすれば、
S_i = \gamma_i S_0 + \sum_{s=1}^{i} \frac{\gamma_i}{\gamma_s} v_s k_s^T
であり、$S_0 = O$とすれば、
\begin{eqnarray}
h_i &=& \sum_{s=1}^{i} \frac{\gamma_i}{\gamma_s} v_s k_s^T q_i\\
&=& V_i \begin{bmatrix}\frac{\gamma_i}{\gamma_1} & 0 & \cdots & 0 & 0 \\ 0 & \frac{\gamma_i}{\gamma_2} & \cdots & 0 & 0 \\
\vdots & \vdots & \ddots & \vdots & \vdots \\ 0 & 0 & \cdots & \frac{\gamma_i}{\gamma_{i-1}} & 0 \\ 0 & 0 & \cdots & 0 & \frac{\gamma_i}{\gamma_i} \end{bmatrix} K_i^T q_i
\end{eqnarray}
である。これは、
h_i = V_i \left(\begin{bmatrix} \frac{\gamma_i}{\gamma_1} \\ \frac{\gamma_i}{\gamma_2} \\ \vdots \\ \frac{\gamma_i}{\gamma_{i-1}} \\ \frac{\gamma_i}{\gamma_i} \end{bmatrix} \odot (K_i^T q_i)\right)
と書くこともできる。同様に$i-1$については、
h_{i-1} = V_{i-1} \left(\begin{bmatrix} \frac{\gamma_{i-1}}{\gamma_1} \\ \frac{\gamma_{i-1}}{\gamma_2} \\ \vdots \\ \frac{\gamma_{i-1}}{\gamma_{i-2}} \\ \frac{\gamma_{i-1}}{\gamma_{i-1}} \end{bmatrix} \odot (K_{i-1}^T q_{i-1})\right)
であるが、$K_i=[k_1,...,k_{i-1}, k_{i}], V_{i} = [v_1, ..., v_{i-1}, v_i]$より、
h_{i-1} = V_{i} \left(\begin{bmatrix} \frac{\gamma_{i-1}}{\gamma_1} \\ \frac{\gamma_{i-1}}{\gamma_2} \\ \vdots \\ \frac{\gamma_{i-1}}{\gamma_{i-2}} \\ \frac{\gamma_{i-1}}{\gamma_{i-1}} \\ 0 \end{bmatrix} \odot (K_{i}^T q_{i-1})\right)
と書くこともできる。したがって、
\Gamma = \begin{bmatrix}
\frac{\gamma_1}{\gamma_1} & \frac{\gamma_2}{\gamma_1} & \frac{\gamma_3}{\gamma_1} & \cdots & \frac{\gamma_{L-1}}{\gamma_1} & \frac{\gamma_L}{\gamma_1} \\
0 & \frac{\gamma_2}{\gamma_2} & \frac{\gamma_3}{\gamma_2} & \cdots & \frac{\gamma_{L-1}}{\gamma_2} & \frac{\gamma_{L}}{\gamma_2} \\
\vdots & \vdots & \vdots & \ddots & \vdots & \vdots \\
0 & 0 & 0 & \cdots & \frac{\gamma_{L-1}}{\gamma_{L-1}} & \frac{\gamma_{L}}{\gamma_{L-1}} \\
0 & 0 & 0 & \cdots & 0 & \frac{\gamma_{L}}{\gamma_{L}} \\
\end{bmatrix}
のような上三角行列$\Gamma$を
\Gamma_{i, j} = \begin{cases}
\frac{\gamma_{j}}{\gamma_{i}} & (i \le j) \\
0 & (i > j)
\end{cases}
と定義すれば、$H = [S_1 q_1, S_2 q_2,...,S_L q_L]$を
H = V(\Gamma \odot (K^T Q))
と行列演算で一気に計算できることになる。
$S_0 = O$ではない場合は、
H = S_0 Q D_{\gamma} + V(\Gamma \odot (K^T Q))
である。ただし、$D_{\gamma}$は$\gamma_1,...,\gamma_L$の対角行列
D_{\gamma} = \begin{bmatrix} \gamma_1 & 0 & \cdots & 0 \\ 0 & \gamma_2 & \cdots & 0 \\ 0 & 0 & \cdots & \gamma_L \end{bmatrix}
Gated-DeltaNet
Alibabaのopen-source modelのQwenの3.5以降で使われている。
Gated-DeltaNetはMamba2をさらに拡張した
\begin{eqnarray}
h_i &=& S_i q_i \\
S_i &=& \alpha_i S_{i-1} (I_d - \beta_i k_i k_i^T) + \beta_i v_i k_i^T \\
\alpha_i &=& \sigma (f_{\theta} (x_i)) \\
\beta_i &=& \sigma (g_{\theta} (x_i))
\end{eqnarray}
という更新式を使う。$\beta_i$という$x_i$に依存する係数が増えている。なお、gatedでないDeltaNetでは$\alpha_i=1$固定(Gated-DeltaNetがDeltaNetにgate $\alpha_i$を導入して拡張したもの)。
$\alpha_i=1$のDeltaNetでは、
S_i = S_{i-1} + \beta_i(v_i - S_{i-1}k_i)k_i^T
である。この両辺を$k_i$に対して左から積で作用させると、
S_i k_i = (1 - \beta_i \|k_i\|^2)S_{i-1} k_i + \beta_i \|k_i\|^2 v_k
となる。つまり、DeltaNetの$S_i$の更新は、$S_i k_i$が$S_{i-1} k_i$から$v_k$へ近づくように移動させている。これは、$k_i$で参照した時に$v_i$の方向が得られるように$S_i$を弱く上書きしているということになる。書き込みの強さは$\beta_i \|k_i\|^2$で決まる。
$S_i$から情報を取り出して$h_i$とする読み込みの操作は
h_i = S_i q_i
であり、$q_i$に含まれる$k_1,...,k_i$の成分によって、$S_i k_1,...,S_i k_i$から$v_1,...,v_i$の情報が重み付き和として得られていると考えられる。
Gated-DeltaNetは、この機構に対してさらに$S_{i-1}$を$\alpha_i S_{i-1}$に置き換えることで動的に忘却するようなgateを入れている。
Gated-DeltaNetも理論上は行列演算で配列全体を一気に計算できる。
\begin{eqnarray}
A_i &=& I_d - \beta_i k_i k_i^T \\
\gamma_i &=& \prod_{t=1}^i \alpha_t
\end{eqnarray}
とすれば、
\begin{eqnarray}
S_i &=& \alpha_i S_{i-1} A_i + \beta_i v_i k_i^T \\
&=& \alpha_i (\alpha_{i-1} S_{i-2} A_{i-1} + \beta_{i-1} v_{i-1} k_{i-1}^T) A_i + \beta_i v_i k_i^T \\
&=& \alpha_i(\alpha_{i-1}(\alpha_{i-2} S_{i-3} A_{i-2} + \beta_{i-2} v_{i-2} k_{i-2}^T) A_{i-1} + \beta_{i-1} v_{i-1} k_{i-1}^T) A_i + \beta_i v_i k_i^T \\
&\vdots& \\
&=& \gamma_i S_0 A_1 A_2...A_i + \sum_{s=1}^{i} \beta_s \frac{\gamma_i}{\gamma_s} v_s k_s^T A_{s+1}...A_{i}
\end{eqnarray}
となる。ただし$s=i$の$A_{s+1}...A_{i}$は$I_d$を意味するとする。この状態では$i$回の行列積計算がある。
ここで、
\begin{eqnarray}
w_1 &=& \beta_1 k_1 \\
w_i &=& \beta_i \left( k_i - \sum_{t=1}^{i-1} (k_i^T k_t) w_t \right)
\end{eqnarray}
という逐次計算で定義されたベクトル系列$w_i$を考える。これを使うと、
I_d - \sum_{t=1}^{i} w_t k_t^T = A_1 A_2 ... A_i
となる。これは帰納法で示せる。
$i = 1$は$I_d - w_1 k_1 = I_d - \beta_1 k_1 k_1^T = A_1$で成立。
$i = r-1$で成立するとすると、
\begin{eqnarray}
I_d - \sum_{t=1}^{r} w_t k_t^T &=& I_d - \sum_{t=1}^{r-1} w_t k_t^T - w_r k_r^T \\
&=& A_1 A_2 ... A_{r-1} - w_r k_r^T \\
&=& A_1 A_2 ... A_{r-1} - \beta_r \left( k_r - \sum_{t=1}^{r-1} (k_r^T k_t) w_t \right) k_r^T \\
&=& A_1 A_2 ... A_{r-1} - \beta_r \left( I_d - \sum_{t=1}^{r-1} w_t k_t^T \right) k_r k_r^T \\
&=& A_1 A_2 ... A_{r-1} - \beta_r A_1 A_2 ... A_{r-1} k_r k_r^T \\
&=& A_1 A_2 ... A_{r-1}(I_d - \beta_r k_r k_r^T) \\
&=& A_1 A_2 ... A_{r-1} A_{r}
\end{eqnarray}
となり、$i=r$でも成立する。
$w_i$を並べた$W_i = [w_1,...,w_i]$と$K_i = [k_1, ...,k_i]$によって、
I_d - \sum_{t=1}^{i} w_t k_t^T = I_d - W_i K_i^T
と計算できるので、$W_i$を求められれば、$i$個の行列の積$A_1 A_2 ... A_{i}$を一度に計算できる
そして、$W_i$は逐次計算でなく行列演算で求められる。
$w_i$の逐次計算の式
w_i = \beta_i \left( k_i - \sum_{t=1}^{i-1} (k_i^T k_t) w_t \right)
は、$\beta_1,...,\beta_i$の対角行列
D_{\beta_i} = \begin{bmatrix}
\beta_1 & 0 & \cdots 0 \\
0 & \beta_2 & \cdots 0 \\
\vdots & \vdots & \ddots \vdots \\
0 & 0 & \cdots \beta_i \\
\end{bmatrix}
および、狭義上三角行列のマスク
M^u_i = \begin{bmatrix}
0 & 1 & \cdots & 1 & 1 \\
0 & 0 & \cdots & 1 & 1 \\
\vdots & \vdots & \ddots & \vdots & \vdots \\
0 & 0 & \cdots & 0 & 1\\
0 & 0 & \cdots & 0 & 0\\
\end{bmatrix}
を定義すれば、
W_i = (K - W_i (U_i \odot (K_i^T K_i)) D_{\beta_i}
と表せる。$w_1 = \beta_1 k_1$も正しく表現される。$M^u_i \odot (K_i^T K_i)$は$K_i^T K_i \in \mathbb{R}^{i \times i}$を要素とする狭義上三角行列である。これを整理して逆行列により$W_i$について解くと、
W_i = K D_{\beta} (I_i + (M^u_i \odot (K_i^T K_i))D_{\beta_i})^{-1}
となる。$I_i + (M^u_i \odot (K_i^T K_i))D_{\beta_i}$は対角成分が全て$1$の上三角行列なので正則であり、逆行列もまた上三角行列である。
配列全体の$i=L$について$W = [w_1,...,w_L]$を一度計算することで、その一部の列から$W_i = [w_1,...,w_i]$が得られる。
これにより、第一項は
\begin{eqnarray}
&& [\gamma_1 S_0 A_1 q_1, \gamma_2 S_0 A_1 A_2 q_2, ..., \gamma_L S_0 A_1 A_2 \cdots A_L q_L] \\
&=& S_0 [(I_d - w_1 k_1^T)q_1, (I_d - w_1 k_1^T- w_2 k_2^T)q_2,...,(I_d - w_1 k_1^T - ... - w_L k_L^T)q_L] D_{\gamma} \\
&=& S_0 (Q - W (M^T \odot (K^T Q))) D_{\gamma}
\end{eqnarray}
と配列全体を計算できる。ただし、$M$はcausal maskであり、$D_{\gamma}$は$\gamma_1,...,\gamma_L$の対角行列。
次に、第二項の
\sum_{s=1}^{i} \beta_s \frac{\gamma_i}{\gamma_s} v_s k_s^T A_{s+1}...A_{i}
を行列演算で一度に計算できる形で表す。これを$G_i$とおく。
逐次計算で定義されるベクトル系列$u_1,...,u_i$を
\begin{eqnarray}
u_1 &=& \beta_1 v_1 \\
u_i &=& \beta_i \left(v_i - \sum_{s=1}^{i-1} \frac{\gamma_i}{\gamma_s} (k_s^T k_i) u_s \right)
\end{eqnarray}
と定義する。これを使うと、
G_i = \sum_{s=1}^{i} \frac{\gamma_i}{\gamma_s} u_s k_s^T
となる。これも帰納法で示すことができる。
定義から$G_1 = \beta_1 v_1 k_1$であり、$u_1 = \beta_1 v_1$なので$i=1$では成立。
$i=r-1$で成立するとした場合、
\begin{eqnarray}
G_r &=& \sum_{s=1}^{i} \beta_s \frac{\gamma_i}{\gamma_s} v_s k_s^T A_{s+1}...A_{i} \\
&=& \sum_{s=1}^{i-1} \beta_s \frac{\gamma_i}{\gamma_s} v_s k_s^T A_{s+1}...A_{i} + \beta_i v_i k_i^T \\
&=& \left(\sum_{s=1}^{i-1} \beta_s \frac{\gamma_{i-1}}{\gamma_s} v_s k_s^T A_{s+1}...A_{i-1}\right)\alpha_i A_i + \beta_i v_i k_i^T \\
&=& \left( \sum_{s=1}^{i-1} \frac{\gamma_{i-1}}{\gamma_s} u_s k_s^T \right) \alpha_i \left(I_d - \beta_i k_i k_i^T\right) + \beta_i v_i k_i^T \\
&=& \left( \sum_{s=1}^{i-1} \frac{\gamma_{i}}{\gamma_s} u_s k_s^T \right) - \beta_i \left( \sum_{s=1}^{i-1} \frac{\gamma_{i-1}}{\gamma_s} u_s k_s^T \right)k_i k_i^T + \beta_i v_i k_i^T \\
&=& \left( \sum_{s=1}^{i-1} \frac{\gamma_{i}}{\gamma_s} u_s k_s^T \right) + \beta_i \left(v_i - \sum_{s=1}^{i-1} \frac{\gamma_i}{\gamma_s} (k_s^T k_i) u_s \right) k_i^T \\
&=& \sum_{s=1}^{i} \frac{\gamma_{i}}{\gamma_s} u_s k_s^T
\end{eqnarray}
となるので、$i=r$でも成立する。
$U = [u_1,...,u_L]$とすれば、Mamba2で使用した上三角行列
\Gamma_{i, j} = \begin{cases}
\frac{\gamma_{j}}{\gamma_{i}} & (i \le j) \\
0 & (i > j)
\end{cases}
を使用して、$u_i$の更新式
u_i = \beta_i \left(v_i - \sum_{s=1}^{i-1} \frac{\gamma_i}{\gamma_s} (k_s^T k_i) u_s \right)
は行列演算で
U = (V - U (\Gamma \odot (K^T K))) D_{\beta}
と表せる。$u_1 = \beta_1 v_1$も正しく表現される。$\Gamma \odot (K^T K)$は$K^T K \in \mathbb{R}^{L \times L}$を要素とする狭義上三角行列である。これを整理して逆行列により$U$について解くと、
U = V D_{\beta} (I_L + (\Gamma \odot (K^T K)))D_{\beta})^{-1}
となる。$I_L + (\Gamma \odot (K^T K)))D_{\beta_i}$は対角成分が全て$1$の上三角行列なので正則であり、逆行列もまた上三角行列である。
この$U$を使って、第二項も第一項と同じように、
\begin{eqnarray}
&& [G_1 q1, G_2 q_2,...G_L q_L] \\
&=& \left[ u_1 k_1^T q_1, \left(\frac{\gamma_2}{\gamma_1} u_1 k_1^T + u_2 k_2^T\right)q_2,..., \left(\sum_{s=1}^{L} \frac{\gamma_L}{\gamma_s} u_s k_s^T \right) q_L \right] \\
&=& U (\Gamma \odot(K^T Q))
\end{eqnarray}
と行列演算で一気に計算できる。
したがって、Gated-DeltaNetは$H = [S_1 q_1, S_2 q_2,...,S_L q_L]$を
\begin{eqnarray}
H &=& S_0 (Q - W (M^T \odot (K^T Q))) D_{\gamma} + U (\Gamma \odot(K^T Q)) \\
W &=& K D_{\beta} (I_L + (M^u \odot (K^T K))D_{\beta})^{-1} \\
U &=& V D_{\beta} (I_L + (\Gamma \odot (K^T K)))D_{\beta})^{-1}
\end{eqnarray}
と計算できる。
$S_0 = O$では
H = U (\Gamma \odot(K^T Q))
であり、式としてはMamba2によく似ているが、$V$でなく
U = V D_{\beta} (I_L + (\Gamma \odot (K^T K)))D_{\beta})^{-1}
を使う。
逐次計算のない行列演算の形式で記述できるが、実際は大きな行列の逆行列の計算が非常に高コストなため、Gated-DeltaNetで長さ$L$の配列全体について直接この形式では計算しない。配列をchunkに分けて、chunk内の並列計算のためにこの方法を使う。
Chunk Computation
inputの配列が揃っている学習時において、逐次計算なしで配列全体のoutputを並列に計算できるのは非常に効率的で、RNN/LSTMよりもTransformerが優れている点でもある。
しかし、配列全体を並列に一気に計算する場合、ここまでのいずれの手法でも$K^T Q$や$K^T K$といった$L \times L$の配列長の二乗のサイズの行列計算が必要になり$O(L^2)$の計算が要求される。Gated-DeltaNetでは対角要素が全て$1$の上三角行列で逆行列が求めやすい形とはいえ$L \times L$の行列の逆行列が必要である。
Mamba2, Gated-DeltaNetは、配列全体をある程度のサイズのchunkに分けることで、並列計算と逐次計算を組み合わせたChunk内並列とChunk間逐次処理により、ハードウェアで可能な範囲まで並列して効率化しつつ、逐次処理により非常に長い配列も学習可能にしている。
chunkの考え方は単純で、$[x_1, ..., x_L]$のような長い配列を
[x_1,...,x_{C_1}], [x_{C_1+1},...,x_{C_2}], ..., [x_{C_{m-1}+1},...,x_{L}]
のように小さな配列に切り分け、
\begin{eqnarray}
X^{C}_1 &=& [x_1,...,x_{C_1}] \\
X^{C}_2 &=& [x_{C_1+1},...,x_{C_2}] \\
\vdots \\
X^{C}_m &=& [x_{C_{m-1}+1},...,x_{L}] \\
\end{eqnarray}
のようにして、$Q, K, V$もchunkごとに
\begin{eqnarray}
Q^C_1 &=& W_Q X^C_1, K^C_1 = W_K X^C_1, V^C_1 = W_V X^C_1 \\
Q^C_2 &=& W_Q X^C_2, K^C_2 = W_K X^C_2, V^C_1 = W_V X^C_2 \\
\vdots \\
Q^C_m &=& W_Q X^C_m, K^C_m = W_K X^C_m, V^C_m = W_V X^C_m \\
\end{eqnarray}
のように計算する。各chunkの配列長を$L^C_1, L_C^2,...,L^C_m$として、
\begin{eqnarray}
X^{C}_1 &=& [x^{C_1}_1,...,x^{C_1}_{L^C_1}] \\
X^{C}_2 &=& [x^{C_2}_1,...,x^{C_2}_{L^C_2}] \\
\vdots \\
X^{C}_m &=& [x^{C_m}_1,...,x^{C_m}_{L^C_m}] \\
\end{eqnarray}
とchunk内のindexで表せるとする。
chunk $C_j$の計算は、1つ前のchunk $C_{j-1}$の最後のtokenの$S_{C_{j-1}}$を$S_0$としたようなchunk内のtokenの$H^C_{j}$の計算と、$S_{C_{j-1}}$から$S_{C_{j}}$への更新である。chunk内では逐次計算を行わず、行列演算で一気に計算する。そして、次のchunkへ$S_{C_{j}}$を渡して次のchunkはそれを$S_0$とした計算を行う。つまり、chunk間は逐次処理である。
Mamba2であればこの計算は以下のようになる。
\begin{eqnarray}
H^C_j &=& S_{C_{j-1}} Q^C_j D_{\gamma^{C_j}} + V^C_j(\Gamma^{C_j} \odot ((K^C_j)^T Q^C_j)) \\
S_{C_{j}} &=& \gamma^{C_j}_{L^C_j} S_{C_{j-1}} + V^C_j \begin{bmatrix}\frac{\gamma^{C_j}_{L^C_j}}{\gamma^{C_j}_{1}} & 0 & \cdots & 0 & 0 \\ 0 & \frac{\gamma^{C_j}_{L^C_j}}{\gamma^{C_j}_{2}} & \cdots & 0 & 0 \\
\vdots & \vdots & \ddots & \vdots & \vdots \\ 0 & 0 & \cdots & \frac{\gamma^{C_j}_{L^C_j}}{\gamma^{C_j}_{L^C_j - 1}} & 0 \\ 0 & 0 & \cdots & 0 & \frac{\gamma^{C_j}_{L^C_j}}{\gamma^{C_j}_{L^C_j}} \end{bmatrix} (K^C_j)^T \\
\end{eqnarray}
ただし、$\gamma^{C_j}_s$はchunk $C_j$内での減衰係数$\alpha$の積であり、
\begin{eqnarray}
\alpha^{C_j}_t &=& \sigma (f_{\theta} (x^{C_j}_t)) \\
\gamma^{C_j}_s &=& \prod_{t=1}^{s} \alpha^{C_j}_t
\end{eqnarray}
のように定義し、$D_{\gamma^{C_j}}$や$\Gamma^{C_j}$はこのchunk内の減衰係数$\gamma^{C_j}_s$によって計算される行列である。
Gated-DeltaNetでも同様にchunkで計算することが可能で、逆行列の計算が必要な上三角行列のサイズはchunkの配列長まで小さくなり、現実的な計算コストでchunk内の並列計算が可能になる。